From abb496b6742f3a451d0243ff0ae0bbf4f6323482 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:46:11 -0400 Subject: [PATCH 01/20] Update .gitignore to exclude data directory and VSCode settings --- .gitignore | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.gitignore b/.gitignore index aa3bad6..78e4ae9 100755 --- a/.gitignore +++ b/.gitignore @@ -157,3 +157,5 @@ explorations/ __marimo__/ uv.lock +data/ +.vscode/settings.json From 716c191a5c4537b3e646570fe5163d3e78897ccf Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:46:26 -0400 Subject: [PATCH 02/20] Add Streamlit dependency and update network inspectors UI script --- pyproject.toml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 7277bf6..ae7187b 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,6 +57,7 @@ mlflow = ["mlflow>=3.14.0"] hf = ["huggingface-hub>=0.20.0"] sbi = ["sbi>=0.26", "nflows>=0.14"] bayesflow = ["bayesflow>=2.0.8", "keras>=3.12"] +ui = ["streamlit>=1.40.0"] all = [ "mlflow>=3.14.0", "huggingface-hub>=0.20.0", @@ -64,6 +65,7 @@ all = [ "nflows>=0.14", "bayesflow>=2.0.8", "keras>=3.12", + "streamlit>=1.40.0", ] [dependency-groups] @@ -156,6 +158,8 @@ torchtrain = "lanfactory.cli.torch_train:app" transform-onnx = "lanfactory.onnx.transform_onnx:app" upload-hf = "lanfactory.cli.upload_hf:app" download-hf = "lanfactory.cli.download_hf:app" +network-inspectors-ui = "lanfactory.cli.network_inspectors_ui:app" [tool.setuptools.package-data] "lanfactory.cli" = ["config_network_training_lan.yaml"] +"lanfactory.network_inspectors" = ["styles.css"] From 7e1682ee90c09d1000ba599a9531494eb14c88c8 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:46:34 -0400 Subject: [PATCH 03/20] Add styles for network inspectors UI components --- src/lanfactory/network_inspectors/styles.css | 186 +++++++++++++++++++ 1 file changed, 186 insertions(+) create mode 100644 src/lanfactory/network_inspectors/styles.css diff --git a/src/lanfactory/network_inspectors/styles.css b/src/lanfactory/network_inspectors/styles.css new file mode 100644 index 0000000..7787776 --- /dev/null +++ b/src/lanfactory/network_inspectors/styles.css @@ -0,0 +1,186 @@ +:root { + --lan-focus: #0a66c2; +} + +div[role="radiogroup"] { + gap: 0.75rem; +} + +div[role="radiogroup"] > label { + border: 2px solid #1f6feb; + border-radius: 10px; + padding: 0.35rem 0.8rem; + background: #f4f8ff; + color: #0b1f3a; + font-weight: 600; +} + +div[role="radiogroup"] > label * { + color: #0b1f3a !important; +} + +div[role="radiogroup"] > label:has(input:checked) { + background: #dbeafe; + color: #052249; + border-color: #0a66c2; +} + +div[role="radiogroup"] > label:has(input:checked) * { + color: #052249 !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] { + border: 2px solid #1f6feb !important; + border-radius: 10px !important; + background: #f4f8ff !important; + overflow: hidden !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] details { + border: 0 !important; + border-radius: 10px !important; + background: transparent !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] summary { + background: #f4f8ff !important; + color: #0b1f3a !important; + font-weight: 600 !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] summary * { + color: #0b1f3a !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] summary:hover { + background: #eaf3ff !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open], +div[data-testid="stSidebar"] div[data-testid="stExpander"]:has(details[open]) { + background: #dbeafe !important; + border-color: #0a66c2 !important; +} + +div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open] summary, +div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open] summary * { + color: #052249 !important; +} + +div.stButton > button { + border: 2px solid #1f6feb; + border-radius: 10px; + background: #f4f8ff; + color: #0b1f3a; + font-weight: 600; +} + +div.stButton > button:hover { + background: #eaf3ff; + border-color: #0a66c2; + color: #052249; +} + +div.stButton > button:active { + background: #dbeafe; + border-color: #0a66c2; + color: #052249; +} + +div[role="radiogroup"] > label:has(input:focus-visible), +button:focus-visible, +input:focus-visible, +textarea:focus-visible, +select:focus-visible { + outline: 3px solid var(--lan-focus) !important; + outline-offset: 2px !important; +} + +div[data-testid="stDataFrame"] thead tr th, +div[data-testid="stDataEditor"] thead tr th { + background-color: #0b3a75; + color: #ffffff; + font-weight: 700; + border-bottom: 2px solid #072b57; +} + +@media (prefers-reduced-motion: reduce) { + *, + *::before, + *::after { + animation-duration: 0.01ms !important; + animation-iteration-count: 1 !important; + transition-duration: 0.01ms !important; + scroll-behavior: auto !important; + } +} + +@media (prefers-color-scheme: dark) { + div[role="radiogroup"] > label { + background: #14365f; + color: #f5f9ff; + border-color: #66a3ff; + } + + div[role="radiogroup"] > label * { + color: #f5f9ff !important; + } + + div[role="radiogroup"] > label:has(input:checked) { + background: #1f4f86; + color: #ffffff; + border-color: #9dc4ff; + } + + div[role="radiogroup"] > label:has(input:checked) * { + color: #ffffff !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] { + background: #14365f !important; + border-color: #66a3ff !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] details { + background: transparent !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] summary, + div[data-testid="stSidebar"] div[data-testid="stExpander"] summary * { + background: #14365f !important; + color: #f5f9ff !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] summary:hover { + background: #1b4473 !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open], + div[data-testid="stSidebar"] div[data-testid="stExpander"]:has(details[open]) { + background: #1f4f86 !important; + border-color: #9dc4ff !important; + } + + div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open] summary, + div[data-testid="stSidebar"] div[data-testid="stExpander"] details[open] summary * { + color: #ffffff !important; + } + + div.stButton > button { + background: #14365f; + color: #f5f9ff; + border-color: #66a3ff; + } + + div.stButton > button:hover { + background: #1b4473; + color: #ffffff; + border-color: #9dc4ff; + } + + div.stButton > button:active { + background: #1f4f86; + color: #ffffff; + border-color: #9dc4ff; + } +} From b0891ddf3a4a60ec2a0e5d8b8579fc594cbccb3e Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:46:56 -0400 Subject: [PATCH 04/20] Refactor KDE-vs-LAN plotting functions for improved data handling and structure --- src/lanfactory/network_inspectors/plotting.py | 80 +++++++++++-------- 1 file changed, 48 insertions(+), 32 deletions(-) diff --git a/src/lanfactory/network_inspectors/plotting.py b/src/lanfactory/network_inspectors/plotting.py index 422b775..f87df6c 100644 --- a/src/lanfactory/network_inspectors/plotting.py +++ b/src/lanfactory/network_inspectors/plotting.py @@ -4,44 +4,38 @@ import logging import os -from typing import TYPE_CHECKING, TypedDict import matplotlib.pyplot as plt import numpy as np import plotly.graph_objects as go import seaborn as sns +from matplotlib.figure import Figure from numpy.typing import NDArray from .config import ModelSpec, PlotConfig - -if TYPE_CHECKING: - import pandas as pd +from .contracts import LikelihoodComparison, LikelihoodRow, ManifoldComputation logger = logging.getLogger(__name__) -class LikelihoodResult(TypedDict): - """Likelihood arrays for one parameter vector.""" - - lan: NDArray[np.float64] - kdes: list[NDArray[np.float64]] +def _save_figure(fig: Figure, filename: str, cfg: PlotConfig) -> None: + os.makedirs(cfg.save_dir, exist_ok=True) + fig.savefig(os.path.join(cfg.save_dir, filename), format="png", transparent=False) -def _save_figure(filename: str, cfg: PlotConfig) -> None: - os.makedirs(cfg.save_dir, exist_ok=True) - plt.savefig(os.path.join(cfg.save_dir, filename), format="png", transparent=False) +def _build_plot_data( + comparison: LikelihoodComparison, +) -> tuple[NDArray[np.float64], list[LikelihoodRow], ModelSpec]: + return comparison.grid, comparison.rows, comparison.spec -def plot_kde_vs_lan( - grid: NDArray[np.float64], - results: list[LikelihoodResult], - spec: ModelSpec, +def build_kde_vs_lan_figure( + comparison: LikelihoodComparison, cfg: PlotConfig, -) -> None: - """Render the KDE-vs-LAN comparison from precomputed likelihoods. +) -> Figure: + """Build and return a matplotlib figure for KDE-vs-LAN likelihoods.""" + grid, results, spec = _build_plot_data(comparison) - results: list of {"lan": array, "kdes": [arrays]}, one per parameter vector. - """ rows = int(np.ceil(len(results) / cfg.cols)) per_choice = grid.shape[0] // spec.n_choices sns.set(style="white", palette="muted", color_codes=True, font_scale=cfg.font_scale) @@ -62,7 +56,7 @@ def plot_kde_vs_lan( row_tmp = i // cfg.cols col_tmp = i - (cfg.cols * row_tmp) - for j, kde_like in enumerate(res["kdes"]): + for j, kde_like in enumerate(res.kdes): if j == 0: label = "KDE" else: @@ -90,11 +84,10 @@ def plot_kde_vs_lan( ax=ax[row_tmp, col_tmp], ) - lan_like = res["lan"] if spec.n_choices == 2: sns.lineplot( x=grid[:, 0] * grid[:, 1], - y=lan_like, + y=res.lan, color="green", label="MLP", alpha=1, @@ -109,7 +102,7 @@ def plot_kde_vs_lan( sns.lineplot( x=grid[per_choice * k : per_choice * (k + 1), 0], - y=lan_like[per_choice * k : per_choice * (k + 1)], + y=res.lan[per_choice * k : per_choice * (k + 1)], color="green", label=label, alpha=1, @@ -140,22 +133,37 @@ def plot_kde_vs_lan( col_tmp = i - (cfg.cols * row_tmp) ax[row_tmp, col_tmp].axis("off") - plt.subplots_adjust(top=0.9) - plt.subplots_adjust(hspace=0.3, wspace=0.3) + fig.subplots_adjust(top=0.9) + fig.subplots_adjust(hspace=0.3, wspace=0.3) + + return fig + + +def plot_kde_vs_lan( + comparison: LikelihoodComparison, + cfg: PlotConfig, +) -> None: + """Render the KDE-vs-LAN comparison from precomputed likelihoods. + + comparison: computed likelihood payload from compute_kde_vs_lan_likelihoods. + """ + fig = build_kde_vs_lan_figure(comparison, cfg) if cfg.save: - _save_figure("kde_vs_mlp_plot.png", cfg) + _save_figure(fig, "kde_vs_mlp_plot.png", cfg) if cfg.show: plt.show() - plt.close() + plt.close(fig) -def plot_manifold( - manifold: pd.DataFrame, spec: ModelSpec, vary_name: str, cfg: PlotConfig -) -> go.Figure: - """Render an interactive 3D LAN likelihood manifold.""" +def build_manifold_figure(computation: ManifoldComputation, cfg: PlotConfig) -> go.Figure: + """Build and return an interactive Plotly manifold figure.""" + manifold = computation.manifold + spec = computation.spec + vary_name = computation.vary_name + plot_data = manifold.assign(signed_rt=manifold["rt"] * manifold["choice"]) surface = ( plot_data.pivot(index="vary", columns="signed_rt", values="likelihood") @@ -185,6 +193,14 @@ def plot_manifold( }, ) + return fig + + +def plot_manifold(computation: ManifoldComputation, cfg: PlotConfig) -> go.Figure: + """Render an interactive 3D LAN likelihood manifold.""" + fig = build_manifold_figure(computation, cfg) + spec = computation.spec + if cfg.save: os.makedirs(cfg.save_dir, exist_ok=True) fig.write_html( From 66184357594b7a5292c9bd7ea5b6976ef0378c56 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:47:06 -0400 Subject: [PATCH 05/20] Add shared result contracts for network inspector compute and plotting layers --- .../network_inspectors/contracts.py | 40 +++++++++++++++++++ 1 file changed, 40 insertions(+) create mode 100644 src/lanfactory/network_inspectors/contracts.py diff --git a/src/lanfactory/network_inspectors/contracts.py b/src/lanfactory/network_inspectors/contracts.py new file mode 100644 index 0000000..4fa2be9 --- /dev/null +++ b/src/lanfactory/network_inspectors/contracts.py @@ -0,0 +1,40 @@ +"""Shared result contracts for network inspector compute and plotting layers.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import pandas as pd +from numpy.typing import NDArray + +from .config import ModelSpec + + +@dataclass +class LikelihoodRow: + """Likelihood arrays for one parameter vector.""" + + params: NDArray[np.float32] + lan: NDArray[np.float64] + kdes: list[NDArray[np.float64]] + + +@dataclass +class LikelihoodComparison: + """Computed LAN/KDE likelihoods over a shared reaction-time grid.""" + + spec: ModelSpec + grid: NDArray[np.float64] + rows: list[LikelihoodRow] + + +@dataclass +class ManifoldComputation: + """Computed LAN manifold payload and metadata for plotting or UI display.""" + + spec: ModelSpec + vary_name: str + vary_values: NDArray[np.float64] + grid: NDArray[np.float64] + manifold: pd.DataFrame From 42a622b8c3bbe7a927e50c9d6d7692db649ae16a Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:47:15 -0400 Subject: [PATCH 06/20] Refactor imports in network inspectors module for improved organization and clarity --- src/lanfactory/network_inspectors/__init__.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/src/lanfactory/network_inspectors/__init__.py b/src/lanfactory/network_inspectors/__init__.py index 6f8e2ab..520410b 100644 --- a/src/lanfactory/network_inspectors/__init__.py +++ b/src/lanfactory/network_inspectors/__init__.py @@ -10,14 +10,25 @@ from __future__ import annotations -from .api import kde_vs_lan_likelihoods, lan_manifold +from .api import ( + compute_kde_vs_lan_likelihoods, + compute_lan_manifold, + kde_vs_lan_likelihoods, + lan_manifold, +) from .config import GridSpec, ModelSpec, PlotConfig +from .contracts import LikelihoodComparison, LikelihoodRow, ManifoldComputation from .loaders import get_torch_mlp __all__ = [ "get_torch_mlp", + "compute_kde_vs_lan_likelihoods", + "compute_lan_manifold", "kde_vs_lan_likelihoods", "lan_manifold", + "LikelihoodComparison", + "LikelihoodRow", + "ManifoldComputation", "ModelSpec", "PlotConfig", "GridSpec", From 6f7db6e20d1e049e4133c488c80dccd8c8c14ed1 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:47:23 -0400 Subject: [PATCH 07/20] Add launcher for network inspectors Streamlit app --- src/lanfactory/cli/network_inspectors_ui.py | 25 +++++++++++++++++++++ 1 file changed, 25 insertions(+) create mode 100644 src/lanfactory/cli/network_inspectors_ui.py diff --git a/src/lanfactory/cli/network_inspectors_ui.py b/src/lanfactory/cli/network_inspectors_ui.py new file mode 100644 index 0000000..c6b7a4a --- /dev/null +++ b/src/lanfactory/cli/network_inspectors_ui.py @@ -0,0 +1,25 @@ +"""Launcher for the network inspectors Streamlit app.""" + +from __future__ import annotations + +from pathlib import Path +import sys + + +def app() -> None: + """Start the LANfactory network inspectors UI via Streamlit.""" + try: + from streamlit.web import cli as stcli + except ImportError as exc: # pragma: no cover + raise ImportError( + "Streamlit is required for the network-inspectors-ui command. " + "Install optional dependencies with: uv sync --extra ui" + ) from exc + + app_path = ( + Path(__file__).resolve().parents[1] + / "network_inspectors" + / "streamlit_app.py" + ) + sys.argv = ["streamlit", "run", str(app_path)] + raise SystemExit(stcli.main()) From cc2b7acecf4b6d20955175542efbc5ff9c0137a4 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:47:35 -0400 Subject: [PATCH 08/20] Refactor likelihood computation functions for improved validation and structure --- src/lanfactory/network_inspectors/api.py | 166 +++++++++++++++-------- 1 file changed, 108 insertions(+), 58 deletions(-) diff --git a/src/lanfactory/network_inspectors/api.py b/src/lanfactory/network_inspectors/api.py index 0635786..48ee7f8 100644 --- a/src/lanfactory/network_inspectors/api.py +++ b/src/lanfactory/network_inspectors/api.py @@ -19,7 +19,8 @@ simulate_ground_truth, ) from .config import GridSpec, ModelSpec, PlotConfig -from .plotting import LikelihoodResult, plot_kde_vs_lan, plot_manifold +from .contracts import LikelihoodComparison, LikelihoodRow, ManifoldComputation +from .plotting import plot_kde_vs_lan, plot_manifold logger = logging.getLogger(__name__) @@ -27,43 +28,60 @@ import plotly.graph_objects as go -def kde_vs_lan_likelihoods( +def _validate_parameter_df(parameter_df: pd.DataFrame, spec: ModelSpec) -> None: + if not isinstance(parameter_df, pd.DataFrame): + raise TypeError("parameter_df must be a pandas.DataFrame.") + if parameter_df.empty: + raise ValueError("parameter_df must contain at least one parameter vector.") + + missing_params = [param for param in spec.params if param not in parameter_df] + if missing_params: + raise ValueError( + "parameter_df is missing model parameter columns: " + + ", ".join(missing_params) + ) + + +def _normalize_parameter_vector( + parameter_df: pd.DataFrame | np.ndarray, + spec: ModelSpec, +) -> NDArray[np.float32]: + if isinstance(parameter_df, pd.DataFrame): + _validate_parameter_df(parameter_df, spec) + if parameter_df.shape[0] > 1: + logger.info("Using only the first row of the supplied parameter array.") + return parameter_df.iloc[0, :][spec.params].values.astype(np.float32) + + parameters = np.asarray(parameter_df, dtype=np.float32) + if parameters.ndim == 2: + if parameters.shape[0] == 0: + raise ValueError("parameter_df must contain at least one row.") + parameters = parameters[0] + if parameters.ndim != 1 or parameters.size != spec.n_params: + raise ValueError(f"Expected one parameter vector with {spec.n_params} values.") + return parameters + + +def compute_kde_vs_lan_likelihoods( parameter_df: pd.DataFrame, model: str, torch_mlp_predict: Callable[[NDArray[np.float32]], Any], n_samples: int = 10, n_reps: int = 10, grid: GridSpec | None = None, - plot: PlotConfig | None = None, -) -> None: - """Compare kernel density estimates from simulation data with LAN output. - - parameter_df: one model-compatible parameter vector per row. - model: model name. torch_mlp_predict: predict_on_batch from get_torch_mlp. - n_samples/n_reps: samples per KDE / KDEs per subplot. - grid: optional GridSpec. plot: optional PlotConfig. - """ +) -> LikelihoodComparison: + """Compute LAN/KDE likelihood arrays for each row in ``parameter_df``.""" if parameter_df is None or model is None or torch_mlp_predict is None: raise ValueError( "parameter_df, model, and torch_mlp_predict are required; build the" " predictor with get_torch_mlp()." ) - if not isinstance(parameter_df, pd.DataFrame): - raise TypeError("parameter_df must be a pandas.DataFrame.") - if parameter_df.empty: - raise ValueError("parameter_df must contain at least one parameter vector.") spec = ModelSpec.from_model(model, predictor=torch_mlp_predict) - missing_params = [param for param in spec.params if param not in parameter_df] - if missing_params: - raise ValueError( - "parameter_df is missing model parameter columns: " - + ", ".join(missing_params) - ) - cfg = plot or PlotConfig() + _validate_parameter_df(parameter_df, spec) grid_arr = make_rt_choice_grid(spec, grid) - results: list[LikelihoodResult] = [] + rows: list[LikelihoodRow] = [] for i in range(parameter_df.shape[0]): params = parameter_df.iloc[i, :][spec.params].values.astype(np.float32) lan_like = np.exp(evaluate_network(spec, params, grid_arr)) @@ -73,25 +91,19 @@ def kde_vs_lan_likelihoods( ) for _ in range(n_reps) ] - results.append({"lan": lan_like, "kdes": kdes}) + rows.append(LikelihoodRow(params=params, lan=lan_like, kdes=kdes)) - return plot_kde_vs_lan(grid_arr, results, spec, cfg) + return LikelihoodComparison(spec=spec, grid=grid_arr, rows=rows) -def lan_manifold( +def compute_lan_manifold( parameter_df: pd.DataFrame | np.ndarray | None = None, vary_dict: dict[str, Any] | None = None, model: str = "ddm", torch_mlp_predict: Callable[[NDArray[np.float32]], Any] | None = None, grid: GridSpec | None = None, - plot: PlotConfig | None = None, -) -> go.Figure: - """Plot LAN likelihoods as a 3D manifold while sweeping one parameter. - - parameter_df: parameter vector (first row used). vary_dict: {param: values}. - model: model name. torch_mlp_predict: predict_on_batch from get_torch_mlp. - grid: optional GridSpec. plot: optional PlotConfig. Returns a Plotly Figure. - """ +) -> ManifoldComputation: + """Compute manifold payload for a one-parameter sweep in 2-choice models.""" if parameter_df is None or torch_mlp_predict is None: raise ValueError( "parameter_df and torch_mlp_predict are required; build the predictor" @@ -115,30 +127,68 @@ def lan_manifold( f"got {spec.n_choices} choices." ) - if isinstance(parameter_df, pd.DataFrame): - if parameter_df.empty: - raise ValueError("parameter_df must contain at least one parameter vector.") - missing_params = [param for param in spec.params if param not in parameter_df] - if missing_params: - raise ValueError( - "parameter_df is missing model parameter columns: " - + ", ".join(missing_params) - ) - if parameter_df.shape[0] > 1: - logger.info("Using only the first row of the supplied parameter array.") - parameters = parameter_df.iloc[0, :][spec.params].values.astype(np.float32) - else: - parameters = np.asarray(parameter_df, dtype=np.float32) - if parameters.ndim == 2: - if parameters.shape[0] == 0: - raise ValueError("parameter_df must contain at least one row.") - parameters = parameters[0] - if parameters.ndim != 1 or parameters.size != spec.n_params: - raise ValueError( - f"Expected one parameter vector with {spec.n_params} values." - ) - + parameters = _normalize_parameter_vector(parameter_df, spec) grid_arr = make_manifold_grid(grid or GridSpec()) manifold = build_manifold(spec, parameters, vary_name, vary_values, grid_arr) - return plot_manifold(manifold, spec, vary_name, plot or PlotConfig()) + return ManifoldComputation( + spec=spec, + vary_name=vary_name, + vary_values=vary_values, + grid=grid_arr, + manifold=manifold, + ) + + +def kde_vs_lan_likelihoods( + parameter_df: pd.DataFrame, + model: str, + torch_mlp_predict: Callable[[NDArray[np.float32]], Any], + n_samples: int = 10, + n_reps: int = 10, + grid: GridSpec | None = None, + plot: PlotConfig | None = None, +) -> None: + """Compare kernel density estimates from simulation data with LAN output. + + parameter_df: one model-compatible parameter vector per row. + model: model name. torch_mlp_predict: predict_on_batch from get_torch_mlp. + n_samples/n_reps: samples per KDE / KDEs per subplot. + grid: optional GridSpec. plot: optional PlotConfig. + """ + computed = compute_kde_vs_lan_likelihoods( + parameter_df=parameter_df, + model=model, + torch_mlp_predict=torch_mlp_predict, + n_samples=n_samples, + n_reps=n_reps, + grid=grid, + ) + cfg = plot or PlotConfig() + + return plot_kde_vs_lan(computed, cfg) + + +def lan_manifold( + parameter_df: pd.DataFrame | np.ndarray | None = None, + vary_dict: dict[str, Any] | None = None, + model: str = "ddm", + torch_mlp_predict: Callable[[NDArray[np.float32]], Any] | None = None, + grid: GridSpec | None = None, + plot: PlotConfig | None = None, +) -> go.Figure: + """Plot LAN likelihoods as a 3D manifold while sweeping one parameter. + + parameter_df: parameter vector (first row used). vary_dict: {param: values}. + model: model name. torch_mlp_predict: predict_on_batch from get_torch_mlp. + grid: optional GridSpec. plot: optional PlotConfig. Returns a Plotly Figure. + """ + computed = compute_lan_manifold( + parameter_df=parameter_df, + vary_dict=vary_dict, + model=model, + torch_mlp_predict=torch_mlp_predict, + grid=grid, + ) + + return plot_manifold(computed, plot or PlotConfig()) From 0f66ceee2ad0aafe1e0d846b1560ff1d68239d1f Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:47:50 -0400 Subject: [PATCH 09/20] Add Streamlit UI for LAN network inspector workflows --- .../network_inspectors/streamlit_app.py | 433 ++++++++++++++++++ 1 file changed, 433 insertions(+) create mode 100644 src/lanfactory/network_inspectors/streamlit_app.py diff --git a/src/lanfactory/network_inspectors/streamlit_app.py b/src/lanfactory/network_inspectors/streamlit_app.py new file mode 100644 index 0000000..40c6fc4 --- /dev/null +++ b/src/lanfactory/network_inspectors/streamlit_app.py @@ -0,0 +1,433 @@ +"""Streamlit UI for LAN network inspector workflows.""" + +from __future__ import annotations + +from importlib.resources import files +from pathlib import Path +import sys + +import numpy as np +import pandas as pd +import ssms +import streamlit as st + +if __package__ in (None, ""): + # Support direct execution with: + # streamlit run src/lanfactory/network_inspectors/streamlit_app.py + src_root = Path(__file__).resolve().parents[2] + src_root_str = str(src_root) + if src_root_str not in sys.path: + sys.path.insert(0, src_root_str) + +from lanfactory.network_inspectors.api import ( # noqa: E402 + compute_kde_vs_lan_likelihoods, + compute_lan_manifold, +) +from lanfactory.network_inspectors.config import GridSpec, ModelSpec, PlotConfig # noqa: E402 +from lanfactory.network_inspectors.loaders import get_torch_mlp # noqa: E402 +from lanfactory.network_inspectors.plotting import ( # noqa: E402 + build_kde_vs_lan_figure, + build_manifold_figure, +) + + +def _load_stylesheet() -> str: + """Load Streamlit CSS from package data with local-file fallback.""" + try: + return ( + files("lanfactory.network_inspectors") + .joinpath("styles.css") + .read_text(encoding="utf-8") + ) + except (FileNotFoundError, ModuleNotFoundError, OSError): + return Path(__file__).resolve().with_name("styles.css").read_text( + encoding="utf-8" + ) + + +def _available_models(base_dir: str) -> list[str]: + model_root = Path(base_dir).expanduser() + if not model_root.exists() or not model_root.is_dir(): + return [] + + valid_ssms_models = set(ssms.config.model_config.keys()) + available = [] + for model_dir in sorted(path for path in model_root.iterdir() if path.is_dir()): + if model_dir.name not in valid_ssms_models: + continue + has_state = any(model_dir.glob("*state_dict*")) + has_cfg = any(model_dir.glob("*network_config*")) + if has_state and has_cfg: + available.append(model_dir.name) + + return available + + +def _resolve_model_paths(base_dir: str, model: str) -> tuple[Path, Path]: + model_dir = Path(base_dir).expanduser() / model + if not model_dir.exists(): + available = _available_models(base_dir) + available_str = ", ".join(available) if available else "none detected" + raise FileNotFoundError( + f"Model directory not found: {model_dir}. " + "Set the base directory to your torch_models folder. " + f"Detected models in this base dir: {available_str}." + ) + + state_candidates = sorted(model_dir.glob("*state_dict*")) + config_candidates = sorted(model_dir.glob("*network_config*")) + + if not state_candidates: + raise FileNotFoundError(f"No state dict file found in {model_dir}.") + if not config_candidates: + raise FileNotFoundError(f"No network config file found in {model_dir}.") + + return state_candidates[0], config_candidates[0] + + +@st.cache_resource(show_spinner=False) +def _load_predictor(base_dir: str, model: str): + spec = ModelSpec.from_model(model) + state_dict_path, network_config_path = _resolve_model_paths(base_dir, model) + input_dim = len(spec.params) + 2 + return get_torch_mlp( + model_file_path=state_dict_path, + network_config=network_config_path, + input_dim=input_dim, + ) + + +def _default_base_dir() -> str: + candidates = [Path("data/torch_models"), Path("../data/torch_models")] + for candidate in candidates: + if candidate.exists(): + return str(candidate) + return "data/torch_models" + + +def _detected_torch_model_dirs() -> list[str]: + """Return likely torch model directories that contain at least one valid model.""" + candidate_paths = [ + Path("data/torch_models"), + Path("data/torch_models/lan"), + Path("data/torch_models/cpn"), + Path("data/torch_models/opn"), + Path("../data/torch_models"), + Path("../data/torch_models/lan"), + Path("../data/torch_models/cpn"), + Path("../data/torch_models/opn"), + Path.cwd() / "data" / "torch_models", + Path.cwd() / "data" / "torch_models" / "lan", + Path.cwd() / "data" / "torch_models" / "cpn", + Path.cwd() / "data" / "torch_models" / "opn", + ] + + data_dir = Path("data") + if data_dir.exists(): + candidate_paths.extend(data_dir.glob("**/torch_models")) + candidate_paths.extend(data_dir.glob("**/torch_models/lan")) + candidate_paths.extend(data_dir.glob("**/torch_models/cpn")) + candidate_paths.extend(data_dir.glob("**/torch_models/opn")) + + discovered: list[str] = [] + seen: set[str] = set() + for candidate in candidate_paths: + path_str = str(candidate) + if path_str in seen: + continue + seen.add(path_str) + if _available_models(path_str): + discovered.append(path_str) + + return discovered + + +def _make_parameter_df(model: str, n_rows: int, seed: int) -> pd.DataFrame: + params = ssms.config.model_config[model]["params"] + lb, ub = ssms.config.model_config[model]["param_bounds"] + rng = np.random.default_rng(seed) + return pd.DataFrame(rng.uniform(lb, ub, size=(n_rows, len(params))), columns=params) + + +def _kde_tab(model: str, predictor, grid_spec: GridSpec, plot_cfg: PlotConfig) -> None: + st.subheader("KDE vs LAN Likelihoods") + col_a, col_b, col_c = st.columns(3) + n_parameter_sets = col_a.slider( + "Parameter sets", + min_value=1, + max_value=20, + value=6, + help="Number of parameter vectors sampled from the model bounds.", + ) + n_samples = col_b.slider( + "Simulator samples", + min_value=100, + max_value=5000, + value=1000, + help="Number of samples used per simulated dataset.", + ) + n_reps = col_c.slider( + "KDE repetitions", + min_value=1, + max_value=30, + value=8, + help="How many KDE estimates are overlaid for each parameter set.", + ) + seed = st.number_input( + "Random seed", + min_value=0, + max_value=2_000_000, + value=123, + help="Controls reproducible generation of parameter vectors.", + ) + + parameter_df = _make_parameter_df(model=model, n_rows=n_parameter_sets, seed=seed) + st.caption("Generated parameter vectors") + st.dataframe(parameter_df, use_container_width=True) + + if st.button("Run KDE vs LAN", use_container_width=True): + with st.spinner("Computing likelihoods..."): + comparison = compute_kde_vs_lan_likelihoods( + parameter_df=parameter_df, + model=model, + torch_mlp_predict=predictor, + n_samples=n_samples, + n_reps=n_reps, + grid=grid_spec, + ) + fig = build_kde_vs_lan_figure(comparison, plot_cfg) + + st.caption( + "KDE vs LAN chart. X-axis is signed reaction time and y-axis is " + "likelihood. Black curves are KDE estimates and green curves are LAN outputs." + ) + st.pyplot(fig, clear_figure=True, use_container_width=True) + + +def _manifold_tab(model: str, predictor, grid_spec: GridSpec, plot_cfg: PlotConfig) -> None: + st.subheader("LAN Manifold") + params = ssms.config.model_config[model]["params"] + defaults = ssms.config.model_config[model]["default_params"] + lb, ub = ssms.config.model_config[model]["param_bounds"] + + base_parameter_df = pd.DataFrame([defaults], columns=params) + st.caption("Base parameter vector") + edited_df = st.data_editor( + base_parameter_df, + num_rows="fixed", + use_container_width=True, + ) + + vary_param = st.selectbox( + "Parameter to sweep", + options=params, + index=0, + help="Choose one parameter to vary while others remain fixed.", + ) + vary_idx = params.index(vary_param) + p_min = float(lb[vary_idx]) + p_max = float(ub[vary_idx]) + + col_a, col_b, col_c = st.columns(3) + sweep_min = col_a.number_input( + "Sweep min", + value=p_min, + help="Minimum value in the parameter sweep range.", + ) + sweep_max = col_b.number_input( + "Sweep max", + value=p_max, + help="Maximum value in the parameter sweep range.", + ) + sweep_steps = col_c.slider( + "Sweep steps", + min_value=5, + max_value=80, + value=20, + help="Number of points between sweep min and sweep max.", + ) + + if st.button("Run Manifold", use_container_width=True): + if sweep_max <= sweep_min: + st.error("Sweep max must be larger than sweep min.") + return + + vary_values = np.linspace(sweep_min, sweep_max, sweep_steps) + with st.spinner("Computing manifold..."): + computation = compute_lan_manifold( + parameter_df=edited_df, + vary_dict={vary_param: vary_values}, + model=model, + torch_mlp_predict=predictor, + grid=grid_spec, + ) + fig = build_manifold_figure(computation, plot_cfg) + + st.caption( + "3D manifold chart. X-axis is signed reaction time, y-axis is the swept " + "parameter value, and z-axis is likelihood." + ) + st.plotly_chart(fig, use_container_width=True) + st.caption("Computed manifold table") + st.dataframe(computation.manifold, use_container_width=True) + + +def run() -> None: + """Render the Streamlit app.""" + st.set_page_config( + page_title="LANfactory Network Inspectors", + page_icon="LAN", + layout="wide", + ) + st.title("LANfactory Network Inspectors") + st.write( + "Inspect trained LAN likelihood behavior with KDE comparisons and manifold plots." + ) + # st.caption( + # "Accessibility: all controls include visible labels, keyboard focus outlines, " + # "and descriptive helper text." + # ) + css = _load_stylesheet() + st.markdown(f"", unsafe_allow_html=True) + + all_models = sorted(ssms.config.model_config.keys()) + + with st.sidebar: + with st.expander("Model Setup", expanded=True): + if "torch_models_base_dir" not in st.session_state: + st.session_state["torch_models_base_dir"] = _default_base_dir() + + detected_dirs = _detected_torch_model_dirs() + if detected_dirs: + default_detected_idx = ( + detected_dirs.index(st.session_state["torch_models_base_dir"]) + if st.session_state["torch_models_base_dir"] in detected_dirs + else 0 + ) + selected_dir = st.selectbox( + "Detected torch model folders", + options=detected_dirs, + index=default_detected_idx, + help="Select a detected folder, or type a custom path below.", + ) + st.session_state["torch_models_base_dir"] = selected_dir + + base_dir = st.text_input( + "Torch models base directory", + key="torch_models_base_dir", + help="Folder containing model subfolders, each with state_dict and network_config files.", + ) + available_models = _available_models(base_dir) + + if available_models: + default_model = "ddm" if "ddm" in available_models else available_models[0] + model = st.selectbox( + "Model", + options=available_models, + index=available_models.index(default_model), + help="Choose a model that exists in the selected torch models directory.", + ) + st.caption( + "Models detected on disk: " + + ", ".join(available_models) + ) + else: + default_model = "ddm" if "ddm" in all_models else all_models[0] + model = st.selectbox( + "Model", + options=all_models, + index=all_models.index(default_model), + help="Choose a model name. You still need matching files on disk.", + ) + st.warning( + "No valid model folders found in the selected base directory. " + "Set a folder containing per-model subdirectories with both " + "*state_dict* and *network_config* files." + ) + + with st.expander("Grid", expanded=True): + n_rt_steps = st.slider( + "Manifold RT steps", + min_value=50, + max_value=800, + value=300, + help="Number of reaction-time points for manifold evaluation.", + ) + max_rt = st.slider( + "Manifold max RT", + min_value=1.0, + max_value=10.0, + value=5.0, + help="Maximum reaction time represented in manifold plots.", + ) + n_points_2c = st.slider( + "KDE grid points (2-choice)", + 200, + 5000, + 2000, + help="Number of reaction-time grid points used for KDE/LAN comparison.", + ) + rt_step_2c = st.number_input( + "KDE grid step (2-choice)", + value=0.0025, + help="Reaction-time spacing in the KDE/LAN evaluation grid.", + ) + + with st.expander("Plot", expanded=True): + cols = st.slider( + "KDE plot columns", + min_value=1, + max_value=6, + value=3, + help="Number of subplot columns for KDE/LAN comparison charts.", + ) + alpha = st.slider( + "KDE alpha", + min_value=0.01, + max_value=0.8, + value=0.1, + help="Opacity of KDE overlay lines.", + ) + font_scale = st.slider( + "Font scale", + min_value=0.8, + max_value=2.5, + value=1.3, + help="Scale factor for chart text.", + ) + + grid_spec = GridSpec( + n_rt_steps=n_rt_steps, + max_rt=max_rt, + n_points_2c=n_points_2c, + rt_step_2c=float(rt_step_2c), + ) + plot_cfg = PlotConfig(show=False, save=False, cols=cols, alpha=alpha, font_scale=font_scale) + + try: + predictor = _load_predictor(base_dir, model) + except Exception as exc: # pragma: no cover - Streamlit interactive path + st.error(str(exc)) + st.stop() + + st.markdown("## Choose Analysis View") + st.info( + "Select a view below. The 3D surface plot is in the '3D Manifold' view." + ) + view = st.radio( + "Analysis view", + options=["KDE vs LAN", "3D Manifold"], + horizontal=True, + help="Choose between 2D KDE/LAN comparison plots and the 3D manifold surface.", + ) + + if view == "KDE vs LAN": + st.markdown("### KDE vs LAN Likelihood Comparison") + _kde_tab(model, predictor, grid_spec, plot_cfg) + else: + st.markdown("### 3D LAN Manifold Surface") + _manifold_tab(model, predictor, grid_spec, plot_cfg) + + +if __name__ == "__main__": + run() From 434d19166c01482b9f5145c0f80272757b25d3f2 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:48:26 -0400 Subject: [PATCH 10/20] Add script for batch training of multiple torch LAN models --- scripts/train_torch_models_batch.sh | 294 ++++++++++++++++++++++++++++ 1 file changed, 294 insertions(+) create mode 100755 scripts/train_torch_models_batch.sh diff --git a/scripts/train_torch_models_batch.sh b/scripts/train_torch_models_batch.sh new file mode 100755 index 0000000..537b913 --- /dev/null +++ b/scripts/train_torch_models_batch.sh @@ -0,0 +1,294 @@ +#!/usr/bin/env bash +set -euo pipefail + +# Train multiple torch LAN models (for example angle and ddm) in one command. +# Uses uv + torchtrain and writes outputs into the chosen networks path. + +CONFIG_TEMPLATE="src/lanfactory/cli/config_network_training_lan.yaml" +TRAINING_DATA_FOLDER="" +TRAINING_DATA_BASE="" +NETWORKS_PATH_BASE="data/torch_models" +MODELS="angle,ddm" +NETWORK_IDS="0,1,2" +DL_WORKERS="1" +LOG_LEVEL="INFO" +DRY_RUN="0" + +usage() { + cat <<'EOF' +Usage: + scripts/train_torch_models_batch.sh [options] + +Options: + --config-template PATH YAML template with a MODEL field. + Default: src/lanfactory/cli/config_network_training_lan.yaml + --training-data-folder PATH Shared training data folder used for all models. + --training-data-base PATH Base folder where each model has its own subfolder. + Example: /angle and /ddm + --networks-path-base PATH Base output path for trained models. + Default: data/torch_models + --models CSV Comma-separated model names. + Default: angle,ddm + --network-ids CSV Comma-separated network IDs. + Default: 0,1,2 + --dl-workers N DataLoader workers passed to torchtrain. + Default: 1 + --log-level LEVEL Logging level for torchtrain. + Default: INFO + --dry-run Validate commands without training. + --help Show this help message. + +Examples: + scripts/train_torch_models_batch.sh \ + --training-data-base data/data \ + --networks-path-base data/torch_models \ + --models angle,ddm \ + --network-ids 0,1,2,3 + + scripts/train_torch_models_batch.sh \ + --training-data-folder data/data/angle \ + --models angle \ + --network-ids 0,1 +EOF +} + +while [[ $# -gt 0 ]]; do + case "$1" in + --config-template) + CONFIG_TEMPLATE="$2" + shift 2 + ;; + --training-data-folder) + TRAINING_DATA_FOLDER="$2" + shift 2 + ;; + --training-data-base) + TRAINING_DATA_BASE="$2" + shift 2 + ;; + --networks-path-base) + NETWORKS_PATH_BASE="$2" + shift 2 + ;; + --models) + MODELS="$2" + shift 2 + ;; + --network-ids) + NETWORK_IDS="$2" + shift 2 + ;; + --dl-workers) + DL_WORKERS="$2" + shift 2 + ;; + --log-level) + LOG_LEVEL="$2" + shift 2 + ;; + --dry-run) + DRY_RUN="1" + shift + ;; + --help) + usage + exit 0 + ;; + *) + echo "Unknown option: $1" >&2 + usage + exit 1 + ;; + esac +done + +if [[ ! -f "$CONFIG_TEMPLATE" ]]; then + echo "Config template not found: $CONFIG_TEMPLATE" >&2 + exit 1 +fi + +if [[ -z "$TRAINING_DATA_FOLDER" && -z "$TRAINING_DATA_BASE" ]]; then + echo "Provide either --training-data-folder or --training-data-base" >&2 + exit 1 +fi + +IFS=',' read -r -a MODELS_ARR <<< "$MODELS" +IFS=',' read -r -a IDS_ARR <<< "$NETWORK_IDS" + +TMP_DIR="$(mktemp -d)" +trap 'rm -rf "$TMP_DIR"' EXIT + +make_model_config() { + local template="$1" + local model="$2" + local out_file="$3" + + # Replace MODEL: ... in YAML while preserving all other lines. + awk -v model="$model" ' + BEGIN { replaced = 0 } + /^[[:space:]]*MODEL[[:space:]]*:/ && replaced == 0 { + print "MODEL: \"" model "\"" + replaced = 1 + next + } + { print } + END { + if (replaced == 0) { + print "MODEL: \"" model "\"" + } + } + ' "$template" > "$out_file" +} + +resolve_model_training_folder() { + local candidate_folder="$1" + local model="$2" + + if [[ ! -d "$candidate_folder" ]]; then + echo "" + return + fi + + # Preferred: pickle shards directly in the provided folder. + if find "$candidate_folder" -maxdepth 1 -type f -name '*.pickle' | grep -q .; then + echo "$candidate_folder" + return + fi + + # Common layout: one subfolder per model. + if [[ -d "$candidate_folder/$model" ]] && find "$candidate_folder/$model" -maxdepth 1 -type f -name '*.pickle' | grep -q .; then + echo "$candidate_folder/$model" + return + fi + + # Fallback: if exactly one immediate subfolder contains pickle shards, use it. + local subdir_count + subdir_count="$(find "$candidate_folder" -mindepth 1 -maxdepth 1 -type d | wc -l | tr -d ' ')" + if [[ "$subdir_count" == "1" ]]; then + local only_subdir + only_subdir="$(find "$candidate_folder" -mindepth 1 -maxdepth 1 -type d | head -n 1)" + if find "$only_subdir" -maxdepth 1 -type f -name '*.pickle' | grep -q .; then + echo "$only_subdir" + return + fi + fi + + echo "$candidate_folder" +} + +validate_training_folder() { + local folder="$1" + local model="$2" + + local n_pickles + n_pickles="$(find "$folder" -maxdepth 1 -type f -name '*.pickle' | wc -l | tr -d ' ')" + + if [[ "$n_pickles" == "0" ]]; then + echo "No .pickle files found for model '$model' in: $folder" >&2 + echo "Expected either:" >&2 + echo " 1) /*.pickle" >&2 + echo " 2) /$model/*.pickle" >&2 + exit 1 + fi + + local validation_out + if ! validation_out="$((uv run python - "$folder" <<'PY' +import glob +import os +import pickle +import sys + +folder = sys.argv[1] +files = sorted(glob.glob(os.path.join(folder, "*.pickle"))) +if not files: + print("ERR_NO_PICKLES") + raise SystemExit(2) + +with open(files[0], "rb") as f: + obj = pickle.load(f) + +if not isinstance(obj, dict): + print("ERR_NOT_DICT") + raise SystemExit(3) + +keys = set(obj.keys()) +required = {"lan_data", "lan_labels"} +if not required.issubset(keys): + print("ERR_BAD_KEYS") + print(",".join(sorted(keys))) + raise SystemExit(4) + +print("OK_KEYS") +PY +))"; then + if grep -q "ERR_BAD_KEYS" <<<"$validation_out"; then + echo "Training data format mismatch for model '$model' in: $folder" >&2 + echo "Expected pickle keys: lan_data and lan_labels" >&2 + echo "Found keys:" >&2 + echo "$(tail -n 1 <<<"$validation_out")" >&2 + exit 1 + fi + + if grep -q "ERR_NOT_DICT" <<<"$validation_out"; then + echo "Unexpected pickle format in: $folder" >&2 + echo "Expected each shard to be a dict with lan_data/lan_labels." >&2 + exit 1 + fi + + echo "Failed to validate training folder '$folder'." >&2 + echo "$validation_out" >&2 + exit 1 + fi +} + +for model in "${MODELS_ARR[@]}"; do + model="${model//[[:space:]]/}" + if [[ -z "$model" ]]; then + continue + fi + + if [[ -n "$TRAINING_DATA_FOLDER" ]]; then + model_training_folder="$(resolve_model_training_folder "$TRAINING_DATA_FOLDER" "$model")" + else + model_training_folder="$(resolve_model_training_folder "$TRAINING_DATA_BASE/$model" "$model")" + fi + + if [[ ! -d "$model_training_folder" ]]; then + echo "Training data folder not found for model '$model': $model_training_folder" >&2 + exit 1 + fi + + validate_training_folder "$model_training_folder" "$model" + + model_cfg="$TMP_DIR/config_${model}.yaml" + make_model_config "$CONFIG_TEMPLATE" "$model" "$model_cfg" + + for net_id in "${IDS_ARR[@]}"; do + net_id="${net_id//[[:space:]]/}" + if [[ -z "$net_id" ]]; then + continue + fi + + cmd=( + uv run torchtrain + --config-path "$model_cfg" + --training-data-folder "$model_training_folder" + --networks-path-base "$NETWORKS_PATH_BASE" + --network-id "$net_id" + --dl-workers "$DL_WORKERS" + --log-level "$LOG_LEVEL" + ) + + if [[ "$DRY_RUN" == "1" ]]; then + cmd+=(--dry-run) + fi + + echo "" + echo "=== Training model=$model network_id=$net_id ===" + echo "Training data: $model_training_folder" + "${cmd[@]}" + done +done + +echo "" +echo "Batch training complete." From ebdc934383ce4100ca7da9377d456f9e1f45d784 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:49:28 -0400 Subject: [PATCH 11/20] Add Streamlit UI section for network inspection workflows and batch training instructions --- README.md | 37 +++++++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/README.md b/README.md index 6277a0f..bb00aba 100755 --- a/README.md +++ b/README.md @@ -52,6 +52,23 @@ Necessary dependency should be installed automatically in the process. Check the basic tutorial [here](docs/basic_tutorial/basic_tutorial_lan_torch.ipynb). +### Network Inspectors UI + +LANfactory includes a Streamlit interface for interactive network inspection +workflows (KDE vs LAN likelihoods and LAN manifold plots). + +Install the UI dependencies: + +```bash +uv sync --extra ui +``` + +Launch the app: + +```bash +uv run network-inspectors-ui +``` + ### Command Line Interface LANfactory includes a command line interface with the commands `jaxtrain` and `torchtrain`, which train neural networks using `jax` and `torch` as backends, respectively. @@ -133,6 +150,26 @@ To make your own configuration file, you can copy the example above into a new ` If you are using `uv`, you can also use the `uv run` command to run `jaxtrain` or `torchtrain` from the command line +### Batch Training Multiple Torch Models + +To generate multiple torch models (for example, `angle` and `ddm`) in one command, +use the helper script: + +scripts/train_torch_models_batch.sh \ + --training-data-base data/data \ + --networks-path-base data/torch_models \ + --models angle,ddm \ + --network-ids 0,1,2 + +If all models should use the same training-data folder, use: + +scripts/train_torch_models_batch.sh \ + --training-data-folder data/data/angle \ + --models angle \ + --network-ids 0,1 + +You can validate setup without training via `--dry-run`. + ### TorchMLP to ONNX Converter Once you have trained your model, you can convert it to the ONNX format using the provided `transform-onnx` command. From 903de376c5fd3c1edc80583615ed46e2c3a50d60 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:49:40 -0400 Subject: [PATCH 12/20] Add tests for compute_kde_vs_lan and compute_lan_manifold functions --- tests/test_network_inspectors_api.py | 51 +++++++++++++++++++++++++--- 1 file changed, 47 insertions(+), 4 deletions(-) diff --git a/tests/test_network_inspectors_api.py b/tests/test_network_inspectors_api.py index 606eb95..f1333fd 100644 --- a/tests/test_network_inspectors_api.py +++ b/tests/test_network_inspectors_api.py @@ -31,7 +31,7 @@ def test_kde_vs_lan_uses_model_parameter_order(monkeypatch): api, "simulate_ground_truth", lambda spec, params, n_samples: {} ) monkeypatch.setattr(api, "evaluate_kde", lambda sim_out, grid: np.zeros(1)) - monkeypatch.setattr(api, "plot_kde_vs_lan", lambda grid, results, spec, cfg: None) + monkeypatch.setattr(api, "plot_kde_vs_lan", lambda comparison, cfg: None) parameter_df = pd.DataFrame( [{"extra": 99.0, "t": 0.3, "z": 0.5, "a": 1.5, "v": 0.2}] @@ -45,6 +45,24 @@ def test_kde_vs_lan_uses_model_parameter_order(monkeypatch): ) +def test_compute_kde_vs_lan_returns_contract(monkeypatch): + monkeypatch.setattr(api, "make_rt_choice_grid", lambda spec, grid: np.zeros((2, 2))) + monkeypatch.setattr(api, "evaluate_network", lambda spec, params, grid: np.zeros(2)) + monkeypatch.setattr(api, "simulate_ground_truth", lambda spec, params, n_samples: {}) + monkeypatch.setattr(api, "evaluate_kde", lambda sim_out, grid: np.zeros(2)) + + parameter_df = pd.DataFrame( + [{"extra": 99.0, "t": 0.3, "z": 0.5, "a": 1.5, "v": 0.2}] + ) + + out = api.compute_kde_vs_lan_likelihoods(parameter_df, "ddm", _predictor, n_reps=2) + + assert out.grid.shape == (2, 2) + assert len(out.rows) == 1 + assert out.rows[0].params.shape == (4,) + assert len(out.rows[0].kdes) == 2 + + def test_lan_manifold_validates_vary_dict(): parameter_df = pd.DataFrame([{"v": 0.2, "a": 1.5, "z": 0.5, "t": 0.3}]) @@ -69,9 +87,7 @@ def build_manifold_stub(spec, parameters, vary_name, vary_values, grid): monkeypatch.setattr(api, "make_manifold_grid", lambda grid: np.zeros((1, 2))) monkeypatch.setattr(api, "build_manifold", build_manifold_stub) - monkeypatch.setattr( - api, "plot_manifold", lambda manifold, spec, vary_name, cfg: None - ) + monkeypatch.setattr(api, "plot_manifold", lambda computation, cfg: None) api.lan_manifold( np.array([[0.2, 1.5, 0.5, 0.3]], dtype=np.float32), @@ -86,6 +102,33 @@ def build_manifold_stub(spec, parameters, vary_name, vary_values, grid): ) +def test_compute_lan_manifold_returns_contract(monkeypatch): + monkeypatch.setattr(api, "make_manifold_grid", lambda grid: np.zeros((3, 2))) + monkeypatch.setattr( + api, + "build_manifold", + lambda spec, parameters, vary_name, vary_values, grid: pd.DataFrame( + { + "rt": [1.0], + "choice": [1], + "vary": [0.2], + "likelihood": [0.5], + } + ), + ) + + out = api.compute_lan_manifold( + parameter_df=np.array([[0.2, 1.5, 0.5, 0.3]], dtype=np.float32), + vary_dict={"v": [0.2]}, + model="ddm", + torch_mlp_predict=_predictor, + ) + + assert out.vary_name == "v" + assert out.grid.shape == (3, 2) + assert list(out.manifold.columns) == ["rt", "choice", "vary", "likelihood"] + + def test_lan_manifold_rejects_wrong_parameter_shape(): with pytest.raises(ValueError, match="Expected one parameter vector"): api.lan_manifold( From 51a56597350e6e8649270c797e449b1a05d00900 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Wed, 22 Jul 2026 11:49:59 -0400 Subject: [PATCH 13/20] Add tests for network inspector plotting functions and update imports --- test_network_inspectors.py | 239 ++++++++++++++++++++++ tests/test_network_inspectors_plotting.py | 64 +++++- 2 files changed, 301 insertions(+), 2 deletions(-) create mode 100644 test_network_inspectors.py diff --git a/test_network_inspectors.py b/test_network_inspectors.py new file mode 100644 index 0000000..5db3c5f --- /dev/null +++ b/test_network_inspectors.py @@ -0,0 +1,239 @@ +import importlib +import sys + +import numpy as np +import pandas as pd +import pytest + + +class DummyLogKDE: + def __init__(self, out): + self.out = out + + def kde_eval(self, data): + return np.zeros(len(data["rts"]), dtype=np.float32) + + +@pytest.fixture +def two_choice_config(): + return {"choices": [-1, 1], "params": ["v", "a"]} + + +@pytest.fixture +def three_choice_config(): + return {"choices": [0, 1, 2], "params": ["v", "a"]} + + +@pytest.fixture +def single_parameter_df(): + return pd.DataFrame([[0.1, 1.0]], columns=["v", "a"]) + + +@pytest.fixture +def network_inspectors_module(): + return importlib.import_module("lanfactory.network_inspectors") + + +@pytest.mark.xfail( + reason="lanfactory eagerly imports network_inspectors from package __init__", + strict=True, +) +def test_import_lanfactory_does_not_eagerly_import_network_inspectors(monkeypatch): + monkeypatch.delitem(sys.modules, "lanfactory.network_inspectors", raising=False) + monkeypatch.delitem(sys.modules, "lanfactory", raising=False) + + importlib.import_module("lanfactory") + + assert "lanfactory.network_inspectors" not in sys.modules + + +@pytest.mark.xfail( + reason="get_torch_mlp falls through to a late TypeError when LoadTorchMLPInfer is unavailable", + strict=True, +) +def test_get_torch_mlp_raises_clear_error_when_loader_is_unavailable( + monkeypatch, network_inspectors_module +): + monkeypatch.setattr(network_inspectors_module, "LoadTorchMLPInfer", None) + + with pytest.raises(ImportError, match="LoadTorchMLPInfer"): + network_inspectors_module.get_torch_mlp( + model_file_path="model.pt", + network_config={"network_type": "lan"}, + input_dim=4, + ) + + +@pytest.mark.xfail( + reason="kde_vs_lan_likelihoods assumes plt.subplots always returns a 2D axes array", + strict=True, +) +def test_kde_vs_lan_likelihoods_handles_single_subplot_layout( + monkeypatch, two_choice_config, single_parameter_df, network_inspectors_module +): + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: two_choice_config, + ) + monkeypatch.setattr( + network_inspectors_module, "simulator", lambda **kwargs: {"x": 1} + ) + monkeypatch.setattr(network_inspectors_module, "LogKDE", DummyLogKDE) + monkeypatch.setattr( + network_inspectors_module.sns, "lineplot", lambda *args, **kwargs: None + ) + + network_inspectors_module.kde_vs_lan_likelihoods( + parameter_df=single_parameter_df, + model="ddm", + torch_mlp_predict=lambda batch: np.zeros((batch.shape[0], 1), dtype=np.float32), + n_reps=1, + cols=1, + show=False, + ) + + +@pytest.mark.xfail( + reason="kde_vs_lan_likelihoods hardcodes a 4000-row input batch for non-binary choice models", + strict=True, +) +def test_kde_vs_lan_likelihoods_sizes_input_batch_from_choice_count( + monkeypatch, + three_choice_config, + single_parameter_df, + network_inspectors_module, +): + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: three_choice_config, + ) + monkeypatch.setattr( + network_inspectors_module, "simulator", lambda **kwargs: {"x": 1} + ) + monkeypatch.setattr(network_inspectors_module, "LogKDE", DummyLogKDE) + monkeypatch.setattr( + network_inspectors_module.sns, "lineplot", lambda *args, **kwargs: None + ) + + network_inspectors_module.kde_vs_lan_likelihoods( + parameter_df=single_parameter_df, + model="lca_3", + torch_mlp_predict=lambda batch: np.zeros((batch.shape[0], 1), dtype=np.float32), + n_reps=1, + cols=1, + show=False, + ) + + +@pytest.mark.xfail( + reason="kde_vs_lan_likelihoods ignores the caller-provided font_scale value", + strict=True, +) +def test_kde_vs_lan_likelihoods_passes_font_scale_argument( + monkeypatch, + two_choice_config, + single_parameter_df, + network_inspectors_module, +): + seen = {} + multi_parameter_df = pd.concat([single_parameter_df] * 4, ignore_index=True) + + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: two_choice_config, + ) + monkeypatch.setattr( + network_inspectors_module, "simulator", lambda **kwargs: {"x": 1} + ) + monkeypatch.setattr(network_inspectors_module, "LogKDE", DummyLogKDE) + monkeypatch.setattr( + network_inspectors_module.sns, "lineplot", lambda *args, **kwargs: None + ) + + def fake_set(**kwargs): + seen.update(kwargs) + + monkeypatch.setattr(network_inspectors_module.sns, "set", fake_set) + + network_inspectors_module.kde_vs_lan_likelihoods( + parameter_df=multi_parameter_df, + model="ddm", + torch_mlp_predict=lambda batch: np.zeros((batch.shape[0], 1), dtype=np.float32), + n_reps=1, + cols=2, + show=False, + font_scale=3.25, + ) + + assert seen["font_scale"] == 3.25 + + +@pytest.mark.xfail( + reason="kde_vs_lan_likelihoods does not validate parameter_df before dereferencing it", + strict=True, +) +def test_kde_vs_lan_likelihoods_rejects_missing_parameter_df( + monkeypatch, two_choice_config, network_inspectors_module +): + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: two_choice_config, + ) + + with pytest.raises(ValueError, match="parameter_df"): + network_inspectors_module.kde_vs_lan_likelihoods( + parameter_df=None, + model="ddm", + torch_mlp_predict=lambda batch: np.zeros( + (batch.shape[0], 1), dtype=np.float32 + ), + show=False, + ) + + +@pytest.mark.xfail( + reason="lan_manifold default vary_dict uses a list but the implementation expects an array with .shape", + strict=True, +) +def test_lan_manifold_accepts_default_vary_dict( + monkeypatch, single_parameter_df, network_inspectors_module +): + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: {"choices": [-1, 1], "params": ["v", "a"]}, + ) + + network_inspectors_module.lan_manifold( + parameter_df=single_parameter_df, + model="ddm", + torch_mlp_predict=lambda batch: np.zeros((batch.shape[0], 1), dtype=np.float32), + show=False, + ) + + +@pytest.mark.xfail( + reason="lan_manifold does not validate torch_mlp_predict before calling it", + strict=True, +) +def test_lan_manifold_rejects_missing_predictor( + monkeypatch, single_parameter_df, network_inspectors_module +): + monkeypatch.setattr( + network_inspectors_module.ModelConfigBuilder, + "from_model", + lambda model: {"choices": [-1, 1], "params": ["v", "a"]}, + ) + + with pytest.raises(ValueError, match="torch_mlp_predict"): + network_inspectors_module.lan_manifold( + parameter_df=single_parameter_df, + vary_dict={"v": np.array([-1.0, 0.0, 1.0], dtype=np.float32)}, + model="ddm", + torch_mlp_predict=None, + show=False, + ) diff --git a/tests/test_network_inspectors_plotting.py b/tests/test_network_inspectors_plotting.py index 4498ad0..4ddc0cd 100644 --- a/tests/test_network_inspectors_plotting.py +++ b/tests/test_network_inspectors_plotting.py @@ -5,9 +5,39 @@ import numpy as np import pandas as pd import plotly.graph_objects as go +from matplotlib.figure import Figure from lanfactory.network_inspectors.config import ModelSpec, PlotConfig -from lanfactory.network_inspectors.plotting import plot_manifold +from lanfactory.network_inspectors.contracts import ( + LikelihoodComparison, + LikelihoodRow, + ManifoldComputation, +) +from lanfactory.network_inspectors.plotting import ( + build_kde_vs_lan_figure, + build_manifold_figure, + plot_manifold, +) + + +def test_build_kde_vs_lan_figure_returns_matplotlib_figure(): + grid = np.array([[-1.0, -1.0], [1.0, 1.0]]) + rows = [ + LikelihoodRow( + params=np.array([0.2, 1.5, 0.5, 0.3], dtype=np.float32), + lan=np.array([0.4, 0.6]), + kdes=[np.array([0.3, 0.7])], + ) + ] + comparison = LikelihoodComparison( + spec=ModelSpec(name="ddm", params=["v", "a", "z", "t"], choices=[-1, 1]), + grid=grid, + rows=rows, + ) + + fig = build_kde_vs_lan_figure(comparison, PlotConfig(show=False, save=False)) + + assert isinstance(fig, Figure) def test_plot_manifold_returns_interactive_plotly_figure(tmp_path): @@ -22,8 +52,15 @@ def test_plot_manifold_returns_interactive_plotly_figure(tmp_path): ) spec = ModelSpec(name="ddm", params=["v"], choices=[-1, 1]) cfg = PlotConfig(show=False, save=True, save_dir=str(tmp_path)) + computation = ManifoldComputation( + spec=spec, + vary_name="v", + vary_values=np.array([0.1, 0.2]), + grid=np.array([[1.0, -1.0], [2.0, -1.0], [1.0, 1.0], [2.0, 1.0]]), + manifold=manifold, + ) - fig = plot_manifold(manifold, spec, "v", cfg) + fig = plot_manifold(computation, cfg) assert isinstance(fig, go.Figure) assert fig.data[0].type == "surface" @@ -34,3 +71,26 @@ def test_plot_manifold_returns_interactive_plotly_figure(tmp_path): [[0.2, 0.1, 0.3, 0.4], [0.3, 0.2, 0.4, 0.5]], ) assert (tmp_path / "mlp_manifold_ddm.html").exists() + + +def test_build_manifold_figure_returns_interactive_plotly_figure(): + manifold = pd.DataFrame( + { + "rt": [1.0, 2.0, 1.0, 2.0], + "choice": [-1, -1, 1, 1], + "vary": [0.1, 0.1, 0.1, 0.1], + "likelihood": [0.1, 0.2, 0.3, 0.4], + } + ) + computation = ManifoldComputation( + spec=ModelSpec(name="ddm", params=["v"], choices=[-1, 1]), + vary_name="v", + vary_values=np.array([0.1]), + grid=np.array([[1.0, -1.0], [2.0, -1.0], [1.0, 1.0], [2.0, 1.0]]), + manifold=manifold, + ) + + fig = build_manifold_figure(computation, PlotConfig(show=False, save=False)) + + assert isinstance(fig, go.Figure) + assert fig.data[0].type == "surface" From 2b12e863f10d695f1fc054b8409980d22610c53f Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:04:48 -0400 Subject: [PATCH 14/20] Reformat --- src/lanfactory/cli/network_inspectors_ui.py | 4 +--- src/lanfactory/network_inspectors/plotting.py | 4 +++- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/src/lanfactory/cli/network_inspectors_ui.py b/src/lanfactory/cli/network_inspectors_ui.py index c6b7a4a..77715b0 100644 --- a/src/lanfactory/cli/network_inspectors_ui.py +++ b/src/lanfactory/cli/network_inspectors_ui.py @@ -17,9 +17,7 @@ def app() -> None: ) from exc app_path = ( - Path(__file__).resolve().parents[1] - / "network_inspectors" - / "streamlit_app.py" + Path(__file__).resolve().parents[1] / "network_inspectors" / "streamlit_app.py" ) sys.argv = ["streamlit", "run", str(app_path)] raise SystemExit(stcli.main()) diff --git a/src/lanfactory/network_inspectors/plotting.py b/src/lanfactory/network_inspectors/plotting.py index f87df6c..9448bc5 100644 --- a/src/lanfactory/network_inspectors/plotting.py +++ b/src/lanfactory/network_inspectors/plotting.py @@ -158,7 +158,9 @@ def plot_kde_vs_lan( plt.close(fig) -def build_manifold_figure(computation: ManifoldComputation, cfg: PlotConfig) -> go.Figure: +def build_manifold_figure( + computation: ManifoldComputation, cfg: PlotConfig +) -> go.Figure: """Build and return an interactive Plotly manifold figure.""" manifold = computation.manifold spec = computation.spec From faaa84e178be55a7de3c4189232d981450b86367 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:06:48 -0400 Subject: [PATCH 15/20] Reformat --- .../network_inspectors/streamlit_app.py | 26 ++++++++++--------- tests/test_network_inspectors_api.py | 4 ++- 2 files changed, 17 insertions(+), 13 deletions(-) diff --git a/src/lanfactory/network_inspectors/streamlit_app.py b/src/lanfactory/network_inspectors/streamlit_app.py index 40c6fc4..51ebdc4 100644 --- a/src/lanfactory/network_inspectors/streamlit_app.py +++ b/src/lanfactory/network_inspectors/streamlit_app.py @@ -40,8 +40,8 @@ def _load_stylesheet() -> str: .read_text(encoding="utf-8") ) except (FileNotFoundError, ModuleNotFoundError, OSError): - return Path(__file__).resolve().with_name("styles.css").read_text( - encoding="utf-8" + return ( + Path(__file__).resolve().with_name("styles.css").read_text(encoding="utf-8") ) @@ -204,7 +204,9 @@ def _kde_tab(model: str, predictor, grid_spec: GridSpec, plot_cfg: PlotConfig) - st.pyplot(fig, clear_figure=True, use_container_width=True) -def _manifold_tab(model: str, predictor, grid_spec: GridSpec, plot_cfg: PlotConfig) -> None: +def _manifold_tab( + model: str, predictor, grid_spec: GridSpec, plot_cfg: PlotConfig +) -> None: st.subheader("LAN Manifold") params = ssms.config.model_config[model]["params"] defaults = ssms.config.model_config[model]["default_params"] @@ -320,17 +322,16 @@ def run() -> None: available_models = _available_models(base_dir) if available_models: - default_model = "ddm" if "ddm" in available_models else available_models[0] + default_model = ( + "ddm" if "ddm" in available_models else available_models[0] + ) model = st.selectbox( "Model", options=available_models, index=available_models.index(default_model), help="Choose a model that exists in the selected torch models directory.", ) - st.caption( - "Models detected on disk: " - + ", ".join(available_models) - ) + st.caption("Models detected on disk: " + ", ".join(available_models)) else: default_model = "ddm" if "ddm" in all_models else all_models[0] model = st.selectbox( @@ -370,6 +371,7 @@ def run() -> None: rt_step_2c = st.number_input( "KDE grid step (2-choice)", value=0.0025, + min_value=0.0001, help="Reaction-time spacing in the KDE/LAN evaluation grid.", ) @@ -402,7 +404,9 @@ def run() -> None: n_points_2c=n_points_2c, rt_step_2c=float(rt_step_2c), ) - plot_cfg = PlotConfig(show=False, save=False, cols=cols, alpha=alpha, font_scale=font_scale) + plot_cfg = PlotConfig( + show=False, save=False, cols=cols, alpha=alpha, font_scale=font_scale + ) try: predictor = _load_predictor(base_dir, model) @@ -411,9 +415,7 @@ def run() -> None: st.stop() st.markdown("## Choose Analysis View") - st.info( - "Select a view below. The 3D surface plot is in the '3D Manifold' view." - ) + st.info("Select a view below. The 3D surface plot is in the '3D Manifold' view.") view = st.radio( "Analysis view", options=["KDE vs LAN", "3D Manifold"], diff --git a/tests/test_network_inspectors_api.py b/tests/test_network_inspectors_api.py index f1333fd..7fe251b 100644 --- a/tests/test_network_inspectors_api.py +++ b/tests/test_network_inspectors_api.py @@ -48,7 +48,9 @@ def test_kde_vs_lan_uses_model_parameter_order(monkeypatch): def test_compute_kde_vs_lan_returns_contract(monkeypatch): monkeypatch.setattr(api, "make_rt_choice_grid", lambda spec, grid: np.zeros((2, 2))) monkeypatch.setattr(api, "evaluate_network", lambda spec, params, grid: np.zeros(2)) - monkeypatch.setattr(api, "simulate_ground_truth", lambda spec, params, n_samples: {}) + monkeypatch.setattr( + api, "simulate_ground_truth", lambda spec, params, n_samples: {} + ) monkeypatch.setattr(api, "evaluate_kde", lambda sim_out, grid: np.zeros(2)) parameter_df = pd.DataFrame( From 7a9eefdf8a1ce46502bb1bdb987b6210c77dd9ad Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:07:48 -0400 Subject: [PATCH 16/20] Close figure --- tests/test_network_inspectors_plotting.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/tests/test_network_inspectors_plotting.py b/tests/test_network_inspectors_plotting.py index 4ddc0cd..b3523c4 100644 --- a/tests/test_network_inspectors_plotting.py +++ b/tests/test_network_inspectors_plotting.py @@ -5,8 +5,6 @@ import numpy as np import pandas as pd import plotly.graph_objects as go -from matplotlib.figure import Figure - from lanfactory.network_inspectors.config import ModelSpec, PlotConfig from lanfactory.network_inspectors.contracts import ( LikelihoodComparison, @@ -38,6 +36,7 @@ def test_build_kde_vs_lan_figure_returns_matplotlib_figure(): fig = build_kde_vs_lan_figure(comparison, PlotConfig(show=False, save=False)) assert isinstance(fig, Figure) + plt.close(fig) def test_plot_manifold_returns_interactive_plotly_figure(tmp_path): From 6d21e7f4230acccd6b60c8dc91db1714fde2efe6 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:12:46 -0400 Subject: [PATCH 17/20] Add error handling --- .../network_inspectors/streamlit_app.py | 22 +++++++++++-------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/src/lanfactory/network_inspectors/streamlit_app.py b/src/lanfactory/network_inspectors/streamlit_app.py index 51ebdc4..ac1b533 100644 --- a/src/lanfactory/network_inspectors/streamlit_app.py +++ b/src/lanfactory/network_inspectors/streamlit_app.py @@ -255,15 +255,19 @@ def _manifold_tab( return vary_values = np.linspace(sweep_min, sweep_max, sweep_steps) - with st.spinner("Computing manifold..."): - computation = compute_lan_manifold( - parameter_df=edited_df, - vary_dict={vary_param: vary_values}, - model=model, - torch_mlp_predict=predictor, - grid=grid_spec, - ) - fig = build_manifold_figure(computation, plot_cfg) + try: + with st.spinner("Computing manifold..."): + computation = compute_lan_manifold( + parameter_df=edited_df, + vary_dict={vary_param: vary_values}, + model=model, + torch_mlp_predict=predictor, + grid=grid_spec, + ) + fig = build_manifold_figure(computation, plot_cfg) + except ValueError as exc: + st.error(str(exc)) + return st.caption( "3D manifold chart. X-axis is signed reaction time, y-axis is the swept " From 3e1113cf5ec534305208c338ecaf8bc02d7d1e2a Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:28:02 -0400 Subject: [PATCH 18/20] Reformat --- tests/test_network_inspectors_plotting.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/test_network_inspectors_plotting.py b/tests/test_network_inspectors_plotting.py index b3523c4..050379c 100644 --- a/tests/test_network_inspectors_plotting.py +++ b/tests/test_network_inspectors_plotting.py @@ -2,6 +2,7 @@ from __future__ import annotations +import matplotlib.pyplot as plt import numpy as np import pandas as pd import plotly.graph_objects as go @@ -16,6 +17,7 @@ build_manifold_figure, plot_manifold, ) +from matplotlib.figure import Figure def test_build_kde_vs_lan_figure_returns_matplotlib_figure(): From f603c4c5fa48d5b2fe853f2acb814f64f2e93c32 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Fri, 31 Jul 2026 12:32:17 -0400 Subject: [PATCH 19/20] Update sh script --- README.md | 4 ++++ scripts/train_torch_models_batch.sh | 2 +- 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index bb00aba..10850c4 100755 --- a/README.md +++ b/README.md @@ -155,18 +155,22 @@ If you are using `uv`, you can also use the `uv run` command to run `jaxtrain` o To generate multiple torch models (for example, `angle` and `ddm`) in one command, use the helper script: +```sh scripts/train_torch_models_batch.sh \ --training-data-base data/data \ --networks-path-base data/torch_models \ --models angle,ddm \ --network-ids 0,1,2 +``` If all models should use the same training-data folder, use: +```sh scripts/train_torch_models_batch.sh \ --training-data-folder data/data/angle \ --models angle \ --network-ids 0,1 +``` You can validate setup without training via `--dry-run`. diff --git a/scripts/train_torch_models_batch.sh b/scripts/train_torch_models_batch.sh index 537b913..286fe57 100755 --- a/scripts/train_torch_models_batch.sh +++ b/scripts/train_torch_models_batch.sh @@ -145,7 +145,7 @@ resolve_model_training_folder() { local model="$2" if [[ ! -d "$candidate_folder" ]]; then - echo "" + echo "$candidate_folder" return fi From fe1d7d64d833e3c1445d46790e9bbefac9cd6af3 Mon Sep 17 00:00:00 2001 From: Carlos Paniagua Date: Mon, 3 Aug 2026 11:53:57 -0400 Subject: [PATCH 20/20] Fix linting --- src/lanfactory/hf/__init__.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/lanfactory/hf/__init__.py b/src/lanfactory/hf/__init__.py index f83e6c5..4ee80ed 100644 --- a/src/lanfactory/hf/__init__.py +++ b/src/lanfactory/hf/__init__.py @@ -4,9 +4,6 @@ downloading models from HuggingFace Hub. """ -DEFAULT_REPO_ID = "franklab/HSSM" -VALID_NETWORK_TYPES = ("lan", "cpn", "opn") - from lanfactory.hf.download import download_model from lanfactory.hf.model_card import ( ModelCardConfig, @@ -15,6 +12,9 @@ ) from lanfactory.hf.upload import upload_model +DEFAULT_REPO_ID = "franklab/HSSM" +VALID_NETWORK_TYPES = ("lan", "cpn", "opn") + __all__ = [ "DEFAULT_REPO_ID", "VALID_NETWORK_TYPES",