Skip to content

lanfactory.network_inspectors

The public inspection namespace loads a trained Torch LAN, compares its likelihood with simulation-based KDE estimates, and plots likelihood manifolds. The configuration dataclasses keep model metadata, evaluation grids, and plot defaults explicit.

lanfactory.network_inspectors.get_torch_mlp

get_torch_mlp(model_file_path: str | PathLike[str], network_config: str | PathLike[str] | dict[str, Any], input_dim: int) -> Callable[[NDArray[np.float32]], Any]

Return a predict_on_batch callable for the TORCH_MLP likelihood.

The returned function expects a 2d float32 array whose rows are a parameter vector trailed by a reaction time and a choice, and returns the per-row LAN log-likelihood.

lanfactory.network_inspectors.kde_vs_lan_likelihoods

kde_vs_lan_likelihoods(parameter_df: DataFrame, model: str, torch_mlp_predict: Callable[[NDArray[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.

lanfactory.network_inspectors.lan_manifold

lan_manifold(parameter_df: DataFrame | ndarray | None = None, vary_dict: dict[str, Any] | None = None, model: str = 'ddm', torch_mlp_predict: Callable[[NDArray[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.

lanfactory.network_inspectors.ModelSpec dataclass

ModelSpec(name: str, params: list[str], choices: list[int], predictor: Callable[[NDArray[float32]], Any] | None = None)

Model metadata (name, params, choices) plus a supplied LAN predictor.

Methods:

  • from_model

    Build a ModelSpec from an ssms model name and optional predictor.

Attributes:

lanfactory.network_inspectors.ModelSpec.n_choices property

n_choices: int

Number of possible choices.

lanfactory.network_inspectors.ModelSpec.n_params property

n_params: int

Number of model parameters.

lanfactory.network_inspectors.ModelSpec.from_model classmethod

from_model(model: str, predictor: Callable[[NDArray[float32]], Any] | None = None) -> ModelSpec

Build a ModelSpec from an ssms model name and optional predictor.

lanfactory.network_inspectors.PlotConfig dataclass

PlotConfig(font_scale: float = 1.5, figsize: tuple[int, int] = (10, 10), cols: int = 3, alpha: float = 0.1, save: bool = False, show: bool = True, save_dir: str = 'figures/', fig_scale: float = 1.0)

Plotting defaults shared by the entry points.

lanfactory.network_inspectors.GridSpec dataclass

GridSpec(n_rt_steps: int = 200, max_rt: float = 5.0, n_points_2c: int = 2000, rt_step_2c: float = 0.0025, n_points_mc: int = 1000, rt_step_mc: float = 0.01)

Reaction-time grid resolution for the inspection plots.