diff --git a/README.md b/README.md
index 6277a0f..10850c4 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,30 @@ 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:
+
+```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`.
+
### TorchMLP to ONNX Converter
Once you have trained your model, you can convert it to the ONNX format using the provided `transform-onnx` command.
diff --git a/pyproject.toml b/pyproject.toml
index 045b4f3..548098f 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]
@@ -164,6 +166,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"]
diff --git a/scripts/train_torch_models_batch.sh b/scripts/train_torch_models_batch.sh
new file mode 100755
index 0000000..286fe57
--- /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 "$candidate_folder"
+ 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."
diff --git a/src/lanfactory/cli/network_inspectors_ui.py b/src/lanfactory/cli/network_inspectors_ui.py
new file mode 100644
index 0000000..77715b0
--- /dev/null
+++ b/src/lanfactory/cli/network_inspectors_ui.py
@@ -0,0 +1,23 @@
+"""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())
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",
diff --git a/src/lanfactory/network_inspectors/__init__.py b/src/lanfactory/network_inspectors/__init__.py
index 6154967..72b7df6 100644
--- a/src/lanfactory/network_inspectors/__init__.py
+++ b/src/lanfactory/network_inspectors/__init__.py
@@ -10,8 +10,14 @@
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__ = [
@@ -19,6 +25,11 @@
"ModelSpec",
"PlotConfig",
"get_torch_mlp",
+ "compute_kde_vs_lan_likelihoods",
+ "compute_lan_manifold",
"kde_vs_lan_likelihoods",
"lan_manifold",
+ "LikelihoodComparison",
+ "LikelihoodRow",
+ "ManifoldComputation",
]
diff --git a/src/lanfactory/network_inspectors/api.py b/src/lanfactory/network_inspectors/api.py
index 570cdbe..3da0588 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())
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
diff --git a/src/lanfactory/network_inspectors/plotting.py b/src/lanfactory/network_inspectors/plotting.py
index 422b775..9448bc5 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,39 @@ 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
+def build_manifold_figure(
+ computation: ManifoldComputation, cfg: PlotConfig
) -> go.Figure:
- """Render an interactive 3D LAN likelihood manifold."""
+ """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 +195,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(
diff --git a/src/lanfactory/network_inspectors/streamlit_app.py b/src/lanfactory/network_inspectors/streamlit_app.py
new file mode 100644
index 0000000..ac1b533
--- /dev/null
+++ b/src/lanfactory/network_inspectors/streamlit_app.py
@@ -0,0 +1,439 @@
+"""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)
+ 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 "
+ "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,
+ min_value=0.0001,
+ 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()
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;
+ }
+}
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_api.py b/tests/test_network_inspectors_api.py
index 606eb95..7fe251b 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,26 @@ 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 +89,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 +104,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(
diff --git a/tests/test_network_inspectors_plotting.py b/tests/test_network_inspectors_plotting.py
index 4498ad0..050379c 100644
--- a/tests/test_network_inspectors_plotting.py
+++ b/tests/test_network_inspectors_plotting.py
@@ -2,12 +2,43 @@
from __future__ import annotations
+import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import plotly.graph_objects as go
-
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,
+)
+from matplotlib.figure import Figure
+
+
+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)
+ plt.close(fig)
def test_plot_manifold_returns_interactive_plotly_figure(tmp_path):
@@ -22,8 +53,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 +72,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"