Skip to content

config

ssms.config

Configuration module for SSM simulators.

This module provides access to model configurations, boundary and drift function configurations, and various generator configurations used throughout the SSMS package. It centralizes all configuration-related functionality to ensure consistent parameter settings across simulations.

Modules:

Classes:

Functions:

ssms.config.CopyOnAccessDict

Bases: dict

A dict that returns a deep copy of the value on lookup.

ssms.config.ModelConfigBuilder

Helper class for building custom model configurations.

This class provides static methods for creating model configurations in various ways: - Starting from an existing model and overriding specific values - Building a configuration from scratch for fully custom simulators - Creating minimal valid configurations - Validating configurations

Examples:

Start from existing model and override:

>>> config = ModelConfigBuilder.from_model("ddm",
...                                     param_bounds=[[-4, 0.3, 0.1, 0],
...                                                   [4, 3.0, 0.9, 2.0]])

Build from scratch:

>>> config = ModelConfigBuilder.from_scratch(
...     name="my_model",
...     params=["v", "a", "z", "t"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )

Create minimal configuration:

>>> config = ModelConfigBuilder.minimal_config(
...     params=["v", "a"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )

Methods:

add_boundary staticmethod

add_boundary(
    config: dict,
    boundary: str | Callable,
    boundary_params: list[str] | None = None,
) -> dict

Add or replace boundary function in configuration.

Parameters:

  • config (dict) –

    Configuration to modify

  • boundary (str or Callable) –

    Boundary function name or callable

  • boundary_params (list[str] or None, default: None ) –

    Parameter names for boundary function (required if boundary is callable)

Returns:

  • dict

    Modified configuration (note: modifies in place and returns)

Raises:

  • ValueError

    If boundary specification is invalid

Examples:

>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_boundary(config, "angle")

add_drift staticmethod

add_drift(
    config: dict, drift: str | Callable, drift_params: list[str] | None = None
) -> dict

Add or replace drift function in configuration.

Parameters:

  • config (dict) –

    Configuration to modify

  • drift (str or Callable) –

    Drift function name or callable

  • drift_params (list[str] or None, default: None ) –

    Parameter names for drift function (required if drift is callable)

Returns:

  • dict

    Modified configuration (note: modifies in place and returns)

Raises:

  • ValueError

    If drift specification is invalid

Examples:

>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_drift(config, "gamma_drift")

from_model staticmethod

from_model(model_name: str, **overrides) -> dict

Create configuration starting from an existing model.

This method automatically parses variant suffixes like "_deadline" from the model name and applies the appropriate transformations.

Parameters:

  • model_name (str) –

    Name of the model, optionally with variant suffixes. Examples: "ddm", "angle", "ddm_deadline", "angle_deadline"

  • **overrides

    Configuration fields to override. Common options: - params : list[str] - Parameter names - param_bounds : list - [[lower bounds], [upper bounds]] - boundary : Callable - Custom boundary function - boundary_name : str - Boundary name - boundary_params : list[str] - Boundary parameter names - drift : Callable - Custom drift function - drift_name : str - Drift name - drift_params : list[str] - Drift parameter names - simulator : Callable - Custom simulator function - nchoices : int - Number of choices - choices : list - Possible choice values

Returns:

  • dict

    Configuration dictionary with any variants applied

Raises:

  • ValueError

    If the base model name is not recognized

Examples:

>>> config = ModelConfigBuilder.from_model("ddm",
...                                     param_bounds=[[-4, 0.3, 0.1, 0],
...                                                   [4, 3.0, 0.9, 2.0]])
>>> # With deadline variant
>>> config = ModelConfigBuilder.from_model("ddm_deadline")
>>> "deadline" in config["params"]
True

from_scratch staticmethod

from_scratch(
    name: str,
    params: list[str],
    simulator_function: Callable,
    nchoices: int,
    **config
) -> dict

Build a complete configuration from scratch.

Use this method when creating a fully custom simulator that doesn't build on any existing model.

Parameters:

  • name (str) –

    Model name

  • params (list[str]) –

    List of parameter names

  • simulator_function (Callable) –

    Simulator function

  • nchoices (int) –

    Number of choices

  • **config

    Additional configuration fields: - param_bounds : list - [[lower bounds], [upper bounds]] - default_params : list - Default parameter values - choices : list - Possible choice values - n_particles : int - Number of particles (default 1) - boundary : Callable - Boundary function - boundary_name : str - Boundary name - boundary_params : list[str] - Boundary parameter names (including 'a') - drift : Callable - Drift function - drift_name : str - Drift name - drift_params : list[str] - Drift parameter names

Returns:

  • dict

    Complete configuration dictionary

Examples:

>>> def my_sim(v, a, **kwargs):
...     # Custom simulation logic
...     return {'rts': ..., 'choices': ..., 'metadata': ...}
>>>
>>> config = ModelConfigBuilder.from_scratch(
...     name="my_custom_model",
...     params=["v", "a"],
...     simulator_function=my_sim,
...     nchoices=2,
...     param_bounds=[[-2, 0.5], [2, 2.0]],
...     default_params=[0.0, 1.0]
... )

get_sampling_transforms staticmethod

get_sampling_transforms(config: dict) -> list

Get parameter sampling transforms from model config.

These transforms are applied during the parameter sampling stage of the training data generation workflow. They enforce parameter relationships (e.g., a > z) when generating synthetic training data for likelihood approximation networks.

Note: These are NOT directly relevant for basic Simulator usage, which uses simulation transforms via ParameterSimulatorAdapters instead.

Parameters:

  • config (dict) –

    Model configuration dictionary

Returns:

  • list

    List of transform/constraint instances (empty if none defined)

Examples:

>>> config = {
...     "parameter_transforms": {
...         "sampling": [SwapIfLessConstraint("a", "z")],
...     }
... }
>>> transforms = ModelConfigBuilder.get_sampling_transforms(config)

get_simulation_transforms staticmethod

get_simulation_transforms(config: dict) -> list

Get simulation transforms from model config.

These transforms are applied via ParameterSimulatorAdapters when running the basic Simulator. They prepare user-provided parameters for the low-level C/Cython simulators (e.g., stacking v0, v1, v2 into a single v array, expanding dimensions).

Parameters:

  • config (dict) –

    Model configuration dictionary

Returns:

  • list

    List of transform instances (empty if none defined)

Examples:

>>> config = {
...     "parameter_transforms": {
...         "simulation": [
...             ColumnStackParameters(["v0", "v1", "v2"], "v"),
...             ExpandDimension(["a", "z"]),
...         ],
...     }
... }
>>> transforms = ModelConfigBuilder.get_simulation_transforms(config)
>>> len(transforms)
2

get_transforms staticmethod

get_transforms(config: dict, phase: str) -> list

Get transforms for a specific phase from model config.

This method extracts parameter transforms from the unified parameter_transforms field in the model configuration.

Parameters:

  • config (dict) –

    Model configuration dictionary

  • phase (str) –

    Either 'sampling' or 'simulation'

Returns:

  • list

    List of transform instances (empty if none defined)

Examples:

>>> config = {
...     "name": "lba_angle_3",
...     "parameter_transforms": {
...         "sampling": [SwapIfLessConstraint("a", "z")],
...         "simulation": [ColumnStackParameters(["v0", "v1"], "v")],
...     }
... }
>>> sampling_transforms = ModelConfigBuilder.get_transforms(config, "sampling")
>>> len(sampling_transforms)
1

minimal_config staticmethod

minimal_config(
    params: list[str],
    simulator_function: Callable,
    nchoices: int = 2,
    name: str = "custom",
) -> dict

Create a minimal valid configuration.

This is the simplest way to create a configuration for a custom simulator. It includes only the required fields.

Parameters:

  • params (list[str]) –

    List of parameter names

  • simulator_function (Callable) –

    Simulator function

  • nchoices (int, default: 2 ) –

    Number of choices

  • name (str, default: "custom" ) –

    Model name

Returns:

  • dict

    Minimal configuration dictionary

Examples:

>>> config = ModelConfigBuilder.minimal_config(
...     params=["v", "a", "z", "t"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )

validate_config staticmethod

validate_config(config: dict, strict: bool = False) -> tuple[bool, list[str]]

Validate a configuration dictionary.

Parameters:

  • config (dict) –

    Configuration to validate

  • strict (bool, default: False ) –

    If True, also check for recommended optional fields

Returns:

  • is_valid ( bool ) –

    Whether the configuration is valid

  • errors ( list[str] ) –

    List of error messages (empty if valid)

Examples:

>>> config = {"params": ["v", "a"], "nchoices": 2}
>>> is_valid, errors = ModelConfigBuilder.validate_config(config)
>>> if not is_valid:
...     print("Errors:", errors)

with_deadline staticmethod

with_deadline(config: dict) -> dict

Add deadline parameter to a model configuration.

Creates a NEW configuration with the deadline parameter added. This is an immutable operation - the original config is not modified.

The deadline parameter allows models to incorporate response deadlines, where trials are terminated if no response is made within the deadline.

This method is idempotent - calling it on a config that already has the deadline parameter will return an equivalent config.

Parameters:

  • config (dict) –

    Model configuration to extend with deadline support

Returns:

  • dict

    New configuration with deadline parameter added. Includes: - "deadline" appended to params list - Updated param_bounds (both list and dict formats) - Updated default_params - "_deadline" suffix added to name - Incremented n_params - "deadline" metadata flag set to True

Examples:

>>> base_config = ModelConfigBuilder.from_model("ddm")
>>> deadline_config = ModelConfigBuilder.with_deadline(base_config)
>>> "deadline" in deadline_config["params"]
True
>>> deadline_config["name"]
'ddm_deadline'

ssms.config.boundary_config_to_function_params

boundary_config_to_function_params(config: dict) -> dict

Convert boundary configuration to function parameters.

Parameters:

  • config (dict) –

    Dictionary containing the boundary configuration

Returns:

  • dict

    Dictionary with adjusted key names so that they match function parameters names directly.

ssms.config.get_boundary_registry

get_boundary_registry() -> BoundaryRegistry

Get the global boundary registry.

Use this to access registry methods like list_boundaries() or is_registered().

Returns:

Examples:

>>> from ssms.config import get_boundary_registry
>>>
>>> # List all available boundaries
>>> registry = get_boundary_registry()
>>> print(registry.list_boundaries())
['angle', 'constant', 'weibull_cdf', ...]
>>>
>>> # Check if a boundary exists
>>> if registry.is_registered("my_boundary"):
...     config = registry.get("my_boundary")

ssms.config.get_drift_registry

get_drift_registry() -> DriftRegistry

Get the global drift registry.

Use this to access registry methods like list_drifts() or is_registered().

Returns:

Examples:

>>> from ssms.config import get_drift_registry
>>>
>>> # List all available drifts
>>> registry = get_drift_registry()
>>> print(registry.list_drifts())
['constant', 'gamma_drift', ...]
>>>
>>> # Check if a drift exists
>>> if registry.is_registered("my_drift"):
...     config = registry.get("my_drift")

ssms.config.get_model_registry

get_model_registry() -> ModelConfigRegistry

Get the global model registry.

Use this to access registry methods like list_models() or has_model().

Returns:

Examples:

>>> from ssms.config import get_model_registry
>>>
>>> # List all available models
>>> registry = get_model_registry()
>>> print(registry.list_models())
['ddm', 'angle', 'weibull_cdf', ...]
>>>
>>> # Check if a model exists
>>> if registry.has_model("my_model"):
...     config = registry.get("my_model")

ssms.config.register_boundary

register_boundary(name: str, function: Callable, params: list[str]) -> None

Register a boundary function globally.

This is the main entry point for registering custom boundary functions. Once registered, boundaries can be used with ModelConfigBuilder.add_boundary() just like built-in boundaries.

Parameters:

  • name (str) –

    Unique name for the boundary

  • function (Callable) –

    Boundary function with signature (t, **params) -> float or array

  • params (list[str]) –

    List of parameter names the function expects (must include 'a')

Raises:

Examples:

Register a custom exponential decay boundary:

>>> import numpy as np
>>> from ssms.config import register_boundary
>>>
>>> def exponential_decay(t, a=1.0, rate=0.1):
...     return a * np.exp(-rate * t)
>>>
>>> register_boundary(
...     name="exponential_decay",
...     function=exponential_decay,
...     params=["a", "rate"]
... )
>>>
>>> # Now use it with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_boundary(config, "exponential_decay")

Register a collapsing boundary:

>>> def linear_collapse(t, a=1.0, slope=-0.5):
...     return a + slope * t
>>>
>>> register_boundary("linear_collapse", linear_collapse, ["a", "slope"])

ssms.config.register_drift

register_drift(name: str, function: Callable, params: list[str]) -> None

Register a drift function globally.

This is the main entry point for registering custom drift functions. Once registered, drifts can be used with ModelConfigBuilder.add_drift() just like built-in drifts.

Parameters:

  • name (str) –

    Unique name for the drift

  • function (Callable) –

    Drift function with signature (t, **params) -> float or array

  • params (list[str]) –

    List of parameter names the function expects

Raises:

Examples:

Register a custom sinusoidal drift:

>>> import numpy as np
>>> from ssms.config import register_drift
>>>
>>> def sinusoidal_drift(t, frequency=1.0, amplitude=0.5, baseline=1.0):
...     return baseline + amplitude * np.sin(2 * np.pi * frequency * t)
>>>
>>> register_drift(
...     name="sinusoidal",
...     function=sinusoidal_drift,
...     params=["frequency", "amplitude", "baseline"]
... )
>>>
>>> # Now use it with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_drift(config, "sinusoidal")

Register a time-varying drift:

>>> def exponential_drift(t, rate=0.1, asymptote=2.0):
...     return asymptote * (1 - np.exp(-rate * t))
>>>
>>> register_drift("exponential", exponential_drift, ["rate", "asymptote"])

ssms.config.register_model_config

register_model_config(name: str, config: dict) -> None

Register a model configuration globally.

This is the main entry point for registering custom model configurations. Once registered, models can be used with ModelConfigBuilder.from_model() just like built-in models.

Parameters:

  • name (str) –

    Unique name for the model

  • config (dict) –

    Complete model configuration dictionary

Raises:

Examples:

Register a complete custom model:

>>> from ssms.config import register_model_config
>>>
>>> my_model = {
...     "name": "my_custom_ddm",
...     "params": ["v", "a", "z", "t"],
...     "param_bounds": [[-3, 0.3, 0.1, 0], [3, 3.0, 0.9, 2]],
...     "nchoices": 2,
...     "n_params": 4,
...     "default_params": [1.0, 1.5, 0.5, 0.3],
...     "simulator": my_simulator_function,
... }
>>>
>>> register_model_config("my_custom_ddm", my_model)
>>>
>>> # Now use it like any built-in model
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("my_custom_ddm")
>>>
>>> # Or with Simulator
>>> from ssms.basic_simulators import Simulator
>>> sim = Simulator(model="my_custom_ddm")

Register with custom boundary and drift:

>>> advanced_model = {
...     "name": "advanced_ddm",
...     "params": ["v", "a", "z", "t", "theta"],
...     "nchoices": 2,
...     "boundary": my_boundary_fn,
...     "boundary_name": "custom",
...     "boundary_params": ["theta"],
...     "drift": my_drift_fn,
...     "drift_name": "custom",
...     "drift_params": [],
...     "simulator": my_simulator,
... }
>>>
>>> register_model_config("advanced_ddm", advanced_model)

ssms.config.register_model_config_factory

register_model_config_factory(name: str, factory: Callable[[], dict]) -> None

Register a model config factory function globally.

Use this when you want lazy loading of model configurations, or when the config requires computation/processing at access time.

Parameters:

  • name (str) –

    Unique name for the model

  • factory (Callable[[], dict]) –

    Function that returns a complete model configuration dict

Raises:

Examples:

Register with lazy loading:

>>> from ssms.config import register_model_config_factory
>>>
>>> def get_my_model():
...     # This only runs when the model is first accessed
...     return {
...         "name": "my_model",
...         "params": ["v", "a", "z", "t"],
...         "nchoices": 2,
...         "simulator": create_simulator(),  # Expensive operation
...     }
>>>
>>> register_model_config_factory("my_model", get_my_model)
>>>
>>> # Factory is called only when accessing the model
>>> config = ModelConfigBuilder.from_model("my_model")

ssms.config.boundary_registry

Global registry for boundary functions.

This module provides a centralized registry for boundary functions used in sequential sampling models. It follows the same pattern as the parameter sampling constraint registry for consistency.

Examples:

Register a custom boundary:

>>> from ssms.config import register_boundary
>>>
>>> def my_boundary(t, a=1.0, decay=0.1):
...     return a * np.exp(-decay * t)
>>>
>>> register_boundary("exponential", my_boundary, ["a", "decay"])
>>>
>>> # Use with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_boundary(config, "exponential")

List available boundaries:

>>> from ssms.config import get_boundary_registry
>>> print(get_boundary_registry().list_boundaries())

Classes:

Functions:

BoundaryRegistry

BoundaryRegistry()

Global registry for boundary functions.

This registry maintains a mapping of boundary names to their configuration, including the function and parameters.

All boundary functions accept 'a' as an explicit parameter and return the final boundary value directly.

Methods:

  • get

    Get boundary configuration by name.

  • is_registered

    Check if boundary name is registered.

  • list_boundaries

    List all registered boundary names.

  • register

    Register a boundary function.

get
get(name: str) -> dict[str, Any]

Get boundary configuration by name.

Parameters:

  • name (str) –

    Name of the registered boundary

Returns:

  • dict

    Dictionary containing: - 'fun': The boundary function - 'params': List of parameter names (including 'a')

Raises:

  • KeyError

    If boundary name not registered

Examples:

>>> registry = BoundaryRegistry()
>>> config = registry.get("angle")
>>> print(config["params"])
['theta']
is_registered
is_registered(name: str) -> bool

Check if boundary name is registered.

Parameters:

  • name (str) –

    Boundary name to check

Returns:

  • bool

    True if boundary is registered, False otherwise

Examples:

>>> registry = BoundaryRegistry()
>>> registry.is_registered("angle")
True
>>> registry.is_registered("my_custom_boundary")
False
list_boundaries
list_boundaries() -> list[str]

List all registered boundary names.

Returns:

  • list[str]

    Sorted list of all registered boundary names

Examples:

>>> registry = BoundaryRegistry()
>>> boundaries = registry.list_boundaries()
>>> print(boundaries)
['angle', 'constant', 'weibull_cdf', ...]
register
register(name: str, function: Callable, params: list[str]) -> None

Register a boundary function.

Parameters:

  • name (str) –

    Unique name for the boundary (e.g., "angle", "weibull_cdf")

  • function (Callable) –

    Boundary function with signature (t, **params) -> float or array where t is time and params are boundary-specific parameters (including 'a')

  • params (list[str]) –

    List of parameter names the function expects (e.g., ["a", "theta"] for angle)

Raises:

Examples:

>>> def exponential_decay(t, a=1.0, rate=0.1):
...     return a * np.exp(-rate * t)
>>>
>>> registry = BoundaryRegistry()
>>> registry.register("exp_decay", exponential_decay, ["a", "rate"])

get_boundary_registry

get_boundary_registry() -> BoundaryRegistry

Get the global boundary registry.

Use this to access registry methods like list_boundaries() or is_registered().

Returns:

Examples:

>>> from ssms.config import get_boundary_registry
>>>
>>> # List all available boundaries
>>> registry = get_boundary_registry()
>>> print(registry.list_boundaries())
['angle', 'constant', 'weibull_cdf', ...]
>>>
>>> # Check if a boundary exists
>>> if registry.is_registered("my_boundary"):
...     config = registry.get("my_boundary")

register_boundary

register_boundary(name: str, function: Callable, params: list[str]) -> None

Register a boundary function globally.

This is the main entry point for registering custom boundary functions. Once registered, boundaries can be used with ModelConfigBuilder.add_boundary() just like built-in boundaries.

Parameters:

  • name (str) –

    Unique name for the boundary

  • function (Callable) –

    Boundary function with signature (t, **params) -> float or array

  • params (list[str]) –

    List of parameter names the function expects (must include 'a')

Raises:

Examples:

Register a custom exponential decay boundary:

>>> import numpy as np
>>> from ssms.config import register_boundary
>>>
>>> def exponential_decay(t, a=1.0, rate=0.1):
...     return a * np.exp(-rate * t)
>>>
>>> register_boundary(
...     name="exponential_decay",
...     function=exponential_decay,
...     params=["a", "rate"]
... )
>>>
>>> # Now use it with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_boundary(config, "exponential_decay")

Register a collapsing boundary:

>>> def linear_collapse(t, a=1.0, slope=-0.5):
...     return a + slope * t
>>>
>>> register_boundary("linear_collapse", linear_collapse, ["a", "slope"])

ssms.config.config_utils

Utilities for handling generator config with nested structure.

This module provides utilities for working with the nested generator_config structure.

Nested structure (REQUIRED): { "pipeline": {"n_parameter_sets": 100, "n_subruns": 10, ...}, "estimator": {"type": "kde", "bandwidth": 0.1, ...}, "training": {"mixture_probabilities": [0.8, 0.1, 0.1], ...}, "simulator": {"delta_t": 0.001, "max_t": 20.0, ...}, "output": {"folder": "...", "pickle_protocol": 4, ...}, }

Note: Only nested configs are supported. Flat configs are no longer accepted.

Functions:

get_nested_config

get_nested_config(
    config: dict, section: str, key: str, default: Any = None
) -> Any

Get a value from nested config structure.

Args: config: Generator configuration dictionary (must use nested structure) section: Nested section name ("pipeline", "estimator", "training", "simulator", "output") key: Key name within the section default: Default value if key not found

Returns: Value from nested structure if available, else default

Examples: >>> config = {"pipeline": {"n_parameter_sets": 100}} >>> get_nested_config(config, "pipeline", "n_parameter_sets") 100

>>> get_nested_config(config, "pipeline", "missing_key", default=42)
42

has_nested_structure

has_nested_structure(config: dict) -> bool

Check if config uses the required nested structure.

Args: config: Generator configuration dictionary

Returns: True if config has nested sections (required format), False otherwise

Note: Only nested configs are supported. This function validates that configs have the correct structure with at least one of the required sections.

ssms.config.drift_registry

Global registry for drift functions.

This module provides a centralized registry for drift functions used in sequential sampling models. It follows the same pattern as the parameter sampling constraint registry and boundary registry for consistency.

Examples:

Register a custom drift:

>>> from ssms.config import register_drift
>>>
>>> def sinusoidal_drift(t, frequency=1.0, amplitude=0.5, baseline=1.0):
...     return baseline + amplitude * np.sin(2 * np.pi * frequency * t)
>>>
>>> register_drift("sinusoidal", sinusoidal_drift, ["frequency", "amplitude", "baseline"])
>>>
>>> # Use with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_drift(config, "sinusoidal")

List available drifts:

>>> from ssms.config import get_drift_registry
>>> print(get_drift_registry().list_drifts())

Classes:

Functions:

DriftRegistry

DriftRegistry()

Global registry for drift functions.

This registry maintains a mapping of drift names to their configuration, including the function and its parameters.

Methods:

  • get

    Get drift configuration by name.

  • is_registered

    Check if drift name is registered.

  • list_drifts

    List all registered drift names.

  • register

    Register a drift function.

get
get(name: str) -> dict[str, Any]

Get drift configuration by name.

Parameters:

  • name (str) –

    Name of the registered drift

Returns:

  • dict

    Dictionary containing: - 'fun': The drift function - 'params': List of parameter names

Raises:

  • KeyError

    If drift name not registered

Examples:

>>> registry = DriftRegistry()
>>> config = registry.get("gamma_drift")
>>> print(config["params"])
['shape', 'scale', 'c']
is_registered
is_registered(name: str) -> bool

Check if drift name is registered.

Parameters:

  • name (str) –

    Drift name to check

Returns:

  • bool

    True if drift is registered, False otherwise

Examples:

>>> registry = DriftRegistry()
>>> registry.is_registered("gamma_drift")
True
>>> registry.is_registered("my_custom_drift")
False
list_drifts
list_drifts() -> list[str]

List all registered drift names.

Returns:

  • list[str]

    Sorted list of all registered drift names

Examples:

>>> registry = DriftRegistry()
>>> drifts = registry.list_drifts()
>>> print(drifts)
['constant', 'gamma_drift', ...]
register
register(name: str, function: Callable, params: list[str]) -> None

Register a drift function.

Parameters:

  • name (str) –

    Unique name for the drift (e.g., "gamma_drift", "constant")

  • function (Callable) –

    Drift function with signature (t, **params) -> float or array where t is time and params are drift-specific parameters

  • params (list[str]) –

    List of parameter names the function expects (e.g., ["shape", "scale", "c"])

Raises:

Examples:

>>> def linear_drift(t, slope=0.5, intercept=1.0):
...     return intercept + slope * t
>>>
>>> registry = DriftRegistry()
>>> registry.register("linear", linear_drift, ["slope", "intercept"])

get_drift_registry

get_drift_registry() -> DriftRegistry

Get the global drift registry.

Use this to access registry methods like list_drifts() or is_registered().

Returns:

Examples:

>>> from ssms.config import get_drift_registry
>>>
>>> # List all available drifts
>>> registry = get_drift_registry()
>>> print(registry.list_drifts())
['constant', 'gamma_drift', ...]
>>>
>>> # Check if a drift exists
>>> if registry.is_registered("my_drift"):
...     config = registry.get("my_drift")

register_drift

register_drift(name: str, function: Callable, params: list[str]) -> None

Register a drift function globally.

This is the main entry point for registering custom drift functions. Once registered, drifts can be used with ModelConfigBuilder.add_drift() just like built-in drifts.

Parameters:

  • name (str) –

    Unique name for the drift

  • function (Callable) –

    Drift function with signature (t, **params) -> float or array

  • params (list[str]) –

    List of parameter names the function expects

Raises:

Examples:

Register a custom sinusoidal drift:

>>> import numpy as np
>>> from ssms.config import register_drift
>>>
>>> def sinusoidal_drift(t, frequency=1.0, amplitude=0.5, baseline=1.0):
...     return baseline + amplitude * np.sin(2 * np.pi * frequency * t)
>>>
>>> register_drift(
...     name="sinusoidal",
...     function=sinusoidal_drift,
...     params=["frequency", "amplitude", "baseline"]
... )
>>>
>>> # Now use it with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_drift(config, "sinusoidal")

Register a time-varying drift:

>>> def exponential_drift(t, rate=0.1, asymptote=2.0):
...     return asymptote * (1 - np.exp(-rate * t))
>>>
>>> register_drift("exponential", exponential_drift, ["rate", "asymptote"])

ssms.config.model_config_builder

Utilities for building custom model configurations.

This module provides helper classes and functions for creating valid model configurations for use with the Simulator class.

Classes:

ModelConfigBuilder

Helper class for building custom model configurations.

This class provides static methods for creating model configurations in various ways: - Starting from an existing model and overriding specific values - Building a configuration from scratch for fully custom simulators - Creating minimal valid configurations - Validating configurations

Examples:

Start from existing model and override:

>>> config = ModelConfigBuilder.from_model("ddm",
...                                     param_bounds=[[-4, 0.3, 0.1, 0],
...                                                   [4, 3.0, 0.9, 2.0]])

Build from scratch:

>>> config = ModelConfigBuilder.from_scratch(
...     name="my_model",
...     params=["v", "a", "z", "t"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )

Create minimal configuration:

>>> config = ModelConfigBuilder.minimal_config(
...     params=["v", "a"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )

Methods:

add_boundary staticmethod
add_boundary(
    config: dict,
    boundary: str | Callable,
    boundary_params: list[str] | None = None,
) -> dict

Add or replace boundary function in configuration.

Parameters:

  • config (dict) –

    Configuration to modify

  • boundary (str or Callable) –

    Boundary function name or callable

  • boundary_params (list[str] or None, default: None ) –

    Parameter names for boundary function (required if boundary is callable)

Returns:

  • dict

    Modified configuration (note: modifies in place and returns)

Raises:

  • ValueError

    If boundary specification is invalid

Examples:

>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_boundary(config, "angle")
add_drift staticmethod
add_drift(
    config: dict, drift: str | Callable, drift_params: list[str] | None = None
) -> dict

Add or replace drift function in configuration.

Parameters:

  • config (dict) –

    Configuration to modify

  • drift (str or Callable) –

    Drift function name or callable

  • drift_params (list[str] or None, default: None ) –

    Parameter names for drift function (required if drift is callable)

Returns:

  • dict

    Modified configuration (note: modifies in place and returns)

Raises:

  • ValueError

    If drift specification is invalid

Examples:

>>> config = ModelConfigBuilder.from_model("ddm")
>>> config = ModelConfigBuilder.add_drift(config, "gamma_drift")
from_model staticmethod
from_model(model_name: str, **overrides) -> dict

Create configuration starting from an existing model.

This method automatically parses variant suffixes like "_deadline" from the model name and applies the appropriate transformations.

Parameters:

  • model_name (str) –

    Name of the model, optionally with variant suffixes. Examples: "ddm", "angle", "ddm_deadline", "angle_deadline"

  • **overrides

    Configuration fields to override. Common options: - params : list[str] - Parameter names - param_bounds : list - [[lower bounds], [upper bounds]] - boundary : Callable - Custom boundary function - boundary_name : str - Boundary name - boundary_params : list[str] - Boundary parameter names - drift : Callable - Custom drift function - drift_name : str - Drift name - drift_params : list[str] - Drift parameter names - simulator : Callable - Custom simulator function - nchoices : int - Number of choices - choices : list - Possible choice values

Returns:

  • dict

    Configuration dictionary with any variants applied

Raises:

  • ValueError

    If the base model name is not recognized

Examples:

>>> config = ModelConfigBuilder.from_model("ddm",
...                                     param_bounds=[[-4, 0.3, 0.1, 0],
...                                                   [4, 3.0, 0.9, 2.0]])
>>> # With deadline variant
>>> config = ModelConfigBuilder.from_model("ddm_deadline")
>>> "deadline" in config["params"]
True
from_scratch staticmethod
from_scratch(
    name: str,
    params: list[str],
    simulator_function: Callable,
    nchoices: int,
    **config
) -> dict

Build a complete configuration from scratch.

Use this method when creating a fully custom simulator that doesn't build on any existing model.

Parameters:

  • name (str) –

    Model name

  • params (list[str]) –

    List of parameter names

  • simulator_function (Callable) –

    Simulator function

  • nchoices (int) –

    Number of choices

  • **config

    Additional configuration fields: - param_bounds : list - [[lower bounds], [upper bounds]] - default_params : list - Default parameter values - choices : list - Possible choice values - n_particles : int - Number of particles (default 1) - boundary : Callable - Boundary function - boundary_name : str - Boundary name - boundary_params : list[str] - Boundary parameter names (including 'a') - drift : Callable - Drift function - drift_name : str - Drift name - drift_params : list[str] - Drift parameter names

Returns:

  • dict

    Complete configuration dictionary

Examples:

>>> def my_sim(v, a, **kwargs):
...     # Custom simulation logic
...     return {'rts': ..., 'choices': ..., 'metadata': ...}
>>>
>>> config = ModelConfigBuilder.from_scratch(
...     name="my_custom_model",
...     params=["v", "a"],
...     simulator_function=my_sim,
...     nchoices=2,
...     param_bounds=[[-2, 0.5], [2, 2.0]],
...     default_params=[0.0, 1.0]
... )
get_sampling_transforms staticmethod
get_sampling_transforms(config: dict) -> list

Get parameter sampling transforms from model config.

These transforms are applied during the parameter sampling stage of the training data generation workflow. They enforce parameter relationships (e.g., a > z) when generating synthetic training data for likelihood approximation networks.

Note: These are NOT directly relevant for basic Simulator usage, which uses simulation transforms via ParameterSimulatorAdapters instead.

Parameters:

  • config (dict) –

    Model configuration dictionary

Returns:

  • list

    List of transform/constraint instances (empty if none defined)

Examples:

>>> config = {
...     "parameter_transforms": {
...         "sampling": [SwapIfLessConstraint("a", "z")],
...     }
... }
>>> transforms = ModelConfigBuilder.get_sampling_transforms(config)
get_simulation_transforms staticmethod
get_simulation_transforms(config: dict) -> list

Get simulation transforms from model config.

These transforms are applied via ParameterSimulatorAdapters when running the basic Simulator. They prepare user-provided parameters for the low-level C/Cython simulators (e.g., stacking v0, v1, v2 into a single v array, expanding dimensions).

Parameters:

  • config (dict) –

    Model configuration dictionary

Returns:

  • list

    List of transform instances (empty if none defined)

Examples:

>>> config = {
...     "parameter_transforms": {
...         "simulation": [
...             ColumnStackParameters(["v0", "v1", "v2"], "v"),
...             ExpandDimension(["a", "z"]),
...         ],
...     }
... }
>>> transforms = ModelConfigBuilder.get_simulation_transforms(config)
>>> len(transforms)
2
get_transforms staticmethod
get_transforms(config: dict, phase: str) -> list

Get transforms for a specific phase from model config.

This method extracts parameter transforms from the unified parameter_transforms field in the model configuration.

Parameters:

  • config (dict) –

    Model configuration dictionary

  • phase (str) –

    Either 'sampling' or 'simulation'

Returns:

  • list

    List of transform instances (empty if none defined)

Examples:

>>> config = {
...     "name": "lba_angle_3",
...     "parameter_transforms": {
...         "sampling": [SwapIfLessConstraint("a", "z")],
...         "simulation": [ColumnStackParameters(["v0", "v1"], "v")],
...     }
... }
>>> sampling_transforms = ModelConfigBuilder.get_transforms(config, "sampling")
>>> len(sampling_transforms)
1
minimal_config staticmethod
minimal_config(
    params: list[str],
    simulator_function: Callable,
    nchoices: int = 2,
    name: str = "custom",
) -> dict

Create a minimal valid configuration.

This is the simplest way to create a configuration for a custom simulator. It includes only the required fields.

Parameters:

  • params (list[str]) –

    List of parameter names

  • simulator_function (Callable) –

    Simulator function

  • nchoices (int, default: 2 ) –

    Number of choices

  • name (str, default: "custom" ) –

    Model name

Returns:

  • dict

    Minimal configuration dictionary

Examples:

>>> config = ModelConfigBuilder.minimal_config(
...     params=["v", "a", "z", "t"],
...     simulator_function=my_sim_fn,
...     nchoices=2
... )
validate_config staticmethod
validate_config(config: dict, strict: bool = False) -> tuple[bool, list[str]]

Validate a configuration dictionary.

Parameters:

  • config (dict) –

    Configuration to validate

  • strict (bool, default: False ) –

    If True, also check for recommended optional fields

Returns:

  • is_valid ( bool ) –

    Whether the configuration is valid

  • errors ( list[str] ) –

    List of error messages (empty if valid)

Examples:

>>> config = {"params": ["v", "a"], "nchoices": 2}
>>> is_valid, errors = ModelConfigBuilder.validate_config(config)
>>> if not is_valid:
...     print("Errors:", errors)
with_deadline staticmethod
with_deadline(config: dict) -> dict

Add deadline parameter to a model configuration.

Creates a NEW configuration with the deadline parameter added. This is an immutable operation - the original config is not modified.

The deadline parameter allows models to incorporate response deadlines, where trials are terminated if no response is made within the deadline.

This method is idempotent - calling it on a config that already has the deadline parameter will return an equivalent config.

Parameters:

  • config (dict) –

    Model configuration to extend with deadline support

Returns:

  • dict

    New configuration with deadline parameter added. Includes: - "deadline" appended to params list - Updated param_bounds (both list and dict formats) - Updated default_params - "_deadline" suffix added to name - Incremented n_params - "deadline" metadata flag set to True

Examples:

>>> base_config = ModelConfigBuilder.from_model("ddm")
>>> deadline_config = ModelConfigBuilder.with_deadline(base_config)
>>> "deadline" in deadline_config["params"]
True
>>> deadline_config["name"]
'ddm_deadline'

ssms.config.model_registry

Global registry for model configurations.

This module provides a centralized registry for complete model configurations. It follows the same pattern as boundary and drift registries for consistency, with additional support for factory functions to enable lazy loading.

Examples:

Register a custom model configuration:

>>> from ssms.config import register_model_config
>>>
>>> my_config = {
...     "name": "my_ddm",
...     "params": ["v", "a", "z", "t"],
...     "nchoices": 2,
...     "simulator": my_sim_fn,
... }
>>>
>>> register_model_config("my_ddm", my_config)
>>>
>>> # Now use it with ModelConfigBuilder
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("my_ddm")

Register using a factory function:

>>> def get_my_model_config():
...     return {...}
>>>
>>> register_model_config_factory("my_model", get_my_model_config)

List available models:

>>> from ssms.config import get_model_registry
>>> print(get_model_registry().list_models())

Classes:

Functions:

ModelConfigRegistry

ModelConfigRegistry()

Global registry for complete model configurations.

This registry maintains a mapping of model names to their configurations. Supports both direct config registration and factory functions for lazy loading.

Methods:

get
get(name: str) -> dict

Get model configuration by name.

Returns a deep copy of the configuration to prevent accidental mutation of the registered config.

Parameters:

  • name (str) –

    Name of the registered model

Returns:

  • dict

    Complete model configuration dictionary (deep copy)

Raises:

  • KeyError

    If model name not registered

Examples:

>>> registry = ModelConfigRegistry()
>>> config = registry.get("ddm")
>>> print(config["params"])
['v', 'a', 'z', 't']
has_model
has_model(name: str) -> bool

Check if model name is registered.

Parameters:

  • name (str) –

    Model name to check

Returns:

  • bool

    True if model is registered, False otherwise

Examples:

>>> registry = ModelConfigRegistry()
>>> registry.has_model("ddm")
True
>>> registry.has_model("my_custom_model")
False
list_models
list_models() -> list[str]

List all registered model names.

Returns:

  • list[str]

    Sorted list of all registered model names

Examples:

>>> registry = ModelConfigRegistry()
>>> models = registry.list_models()
>>> print(models[:5])
['angle', 'ddm', 'ddm_par2', 'ddm_sdv', 'ddm_st']
register_config
register_config(name: str, config: dict) -> None

Register a model configuration directly.

Parameters:

  • name (str) –

    Unique name for the model (e.g., "ddm", "my_custom_model")

  • config (dict) –

    Complete model configuration dictionary containing at minimum: - 'name': Model name - 'params': List of parameter names - 'nchoices': Number of choices - 'simulator': Simulator function

Raises:

  • ValueError

    If name already registered (either as config or factory)

Examples:

>>> config = {
...     "name": "my_model",
...     "params": ["v", "a", "z", "t"],
...     "param_bounds": [[-3, 0.3, 0.1, 0], [3, 3.0, 0.9, 2]],
...     "nchoices": 2,
...     "simulator": my_sim_fn,
... }
>>>
>>> registry = ModelConfigRegistry()
>>> registry.register_config("my_model", config)
register_factory
register_factory(name: str, factory: Callable[[], dict]) -> None

Register a model config factory function.

Factory functions enable lazy loading - the config is only created when first accessed. This is useful for models with expensive initialization or to reduce memory footprint.

Parameters:

  • name (str) –

    Unique name for the model

  • factory (Callable[[], dict]) –

    Function that returns a complete model configuration dict

Raises:

  • ValueError

    If name already registered (either as config or factory)

Examples:

>>> def get_my_model_config():
...     # Expensive computation here
...     return {...}
>>>
>>> registry = ModelConfigRegistry()
>>> registry.register_factory("my_model", get_my_model_config)

get_model_registry

get_model_registry() -> ModelConfigRegistry

Get the global model registry.

Use this to access registry methods like list_models() or has_model().

Returns:

Examples:

>>> from ssms.config import get_model_registry
>>>
>>> # List all available models
>>> registry = get_model_registry()
>>> print(registry.list_models())
['ddm', 'angle', 'weibull_cdf', ...]
>>>
>>> # Check if a model exists
>>> if registry.has_model("my_model"):
...     config = registry.get("my_model")

register_model_config

register_model_config(name: str, config: dict) -> None

Register a model configuration globally.

This is the main entry point for registering custom model configurations. Once registered, models can be used with ModelConfigBuilder.from_model() just like built-in models.

Parameters:

  • name (str) –

    Unique name for the model

  • config (dict) –

    Complete model configuration dictionary

Raises:

Examples:

Register a complete custom model:

>>> from ssms.config import register_model_config
>>>
>>> my_model = {
...     "name": "my_custom_ddm",
...     "params": ["v", "a", "z", "t"],
...     "param_bounds": [[-3, 0.3, 0.1, 0], [3, 3.0, 0.9, 2]],
...     "nchoices": 2,
...     "n_params": 4,
...     "default_params": [1.0, 1.5, 0.5, 0.3],
...     "simulator": my_simulator_function,
... }
>>>
>>> register_model_config("my_custom_ddm", my_model)
>>>
>>> # Now use it like any built-in model
>>> from ssms.config import ModelConfigBuilder
>>> config = ModelConfigBuilder.from_model("my_custom_ddm")
>>>
>>> # Or with Simulator
>>> from ssms.basic_simulators import Simulator
>>> sim = Simulator(model="my_custom_ddm")

Register with custom boundary and drift:

>>> advanced_model = {
...     "name": "advanced_ddm",
...     "params": ["v", "a", "z", "t", "theta"],
...     "nchoices": 2,
...     "boundary": my_boundary_fn,
...     "boundary_name": "custom",
...     "boundary_params": ["theta"],
...     "drift": my_drift_fn,
...     "drift_name": "custom",
...     "drift_params": [],
...     "simulator": my_simulator,
... }
>>>
>>> register_model_config("advanced_ddm", advanced_model)

register_model_config_factory

register_model_config_factory(name: str, factory: Callable[[], dict]) -> None

Register a model config factory function globally.

Use this when you want lazy loading of model configurations, or when the config requires computation/processing at access time.

Parameters:

  • name (str) –

    Unique name for the model

  • factory (Callable[[], dict]) –

    Function that returns a complete model configuration dict

Raises:

Examples:

Register with lazy loading:

>>> from ssms.config import register_model_config_factory
>>>
>>> def get_my_model():
...     # This only runs when the model is first accessed
...     return {
...         "name": "my_model",
...         "params": ["v", "a", "z", "t"],
...         "nchoices": 2,
...         "simulator": create_simulator(),  # Expensive operation
...     }
>>>
>>> register_model_config_factory("my_model", get_my_model)
>>>
>>> # Factory is called only when accessing the model
>>> config = ModelConfigBuilder.from_model("my_model")