Skip to content

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.

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
... )

add_boundary staticmethod

add_boundary(config, boundary, boundary_params=None)

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)

  • multiplicative (bool, default: True ) –

    Whether boundary is multiplicative (True) or additive (False)

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, drift, drift_params=None)

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, **overrides)

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, params, simulator_function, nchoices, **config
)

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)

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)

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, phase)

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, simulator_function, nchoices=2, name="custom"
)

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, strict=False)

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)

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)

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.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())

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.

__repr__

__repr__()

String representation of registry.

get

get(name)

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)

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 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, function, params)

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:

  • ValueError

    If name already registered

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()

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, function, params)

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:

  • ValueError

    If name already registered

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.

get_nested_config

get_nested_config(config, section, key, default=None)

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)

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())

DriftRegistry

DriftRegistry()

Global registry for drift functions.

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

__repr__

__repr__()

String representation of registry.

get

get(name)

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)

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 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, function, params)

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:

  • ValueError

    If name already registered

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()

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, function, params)

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:

  • ValueError

    If name already registered

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.get_boundary_registry

get_boundary_registry()

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()

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()

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.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.

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
... )

add_boundary staticmethod

add_boundary(config, boundary, boundary_params=None)

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)

  • multiplicative (bool, default: True ) –

    Whether boundary is multiplicative (True) or additive (False)

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, drift, drift_params=None)

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, **overrides)

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, params, simulator_function, nchoices, **config
)

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)

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)

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, phase)

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, simulator_function, nchoices=2, name="custom"
)

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, strict=False)

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)

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())

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.

__repr__

__repr__()

String representation of registry.

get

get(name)

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)

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 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, config)

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, factory)

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()

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, config)

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:

  • ValueError

    If name already registered

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, factory)

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:

  • ValueError

    If name already registered

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.register_boundary

register_boundary(name, function, params)

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:

  • ValueError

    If name already registered

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, function, params)

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:

  • ValueError

    If name already registered

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, config)

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:

  • ValueError

    If name already registered

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, factory)

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:

  • ValueError

    If name already registered

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")