snnlab
API referencesnnlab.sim

snnlab.sim.tool

Complete declared API of the tool module, with signatures, data fields, validation and source.

Back to sim reference

PING network toolkit — CLI entrypoint.

Subcommands: sim, train.

Usage: uv run python -m snnlab.sim sim --out-dir temp/my-sim uv run python -m snnlab.sim sim --infer --load-weights weights.pth --out-dir temp/my-infer uv run python -m snnlab.sim train --epochs 10 --out-dir temp/my-train

The signatures, defaults, fields, docstrings and implementation excerpts below are generated from the Python source. Annotations are shown as declared; unannotated means the source supplies no type annotation. These pages document callable surfaces, including legacy support utilities, without promising backend support for every declaration.

SymbolKind
parse_argsfunction
configure_modelsfunction
save_run_artifactsfunction
mainfunction

parse_args

View source

def parse_args(argv=None)

Source docstring:

Parse command-line arguments with subparsers for sim/train.
ParameterAnnotationDefaultMeaning
argvunannotatedNoneDefined by the source contract and implementation below.

Return expressions (branch-dependent; names refer to the linked implementation):

args
Implementation
def parse_args(argv=None):
    """Parse command-line arguments with subparsers for sim/train."""
    import argparse

    argv = sys.argv[1:] if argv is None else list(argv)

    _examples = """\
Each subcommand has its own complete argument listing. The top-level help
above only shows the dispatcher; for the actual flags accepted by a mode,
run:

  python -m snnlab.sim sim    --help
  python -m snnlab.sim train  --help

The flags fall into the following groups (every group is documented in
each subcommand's --help):

  Network        --model, --n-hidden, --ei-strength, --ei-ratio,
                 --recurrent-initial-zero-fraction, --w-in-initial-zero-fraction, --dt, --t-ms, --seed
  Dynamics       --train-leak, --adaptive-threshold
  Readout        --readout {rate,mem-mean,spike-count,spike-rate,cumulative-potential},
                 --signed-readout, --readout-bias, --readout-w-init-mean,
                 --readout-w-init-std, --readout-w-out-scale,
                 --dales-law, --no-dales-law
  Input          --input, --input-rate, --dataset, --digit, --sample
  Weights        --w-in, --w-ee, --w-ei, --w-ie, --w-ii
  Gradient       --v-grad-dampen, --surrogate-slope
  Train (train)  --lr, --epochs, --batch-size, --max-samples,
                 --fr-reg-upper-target-hz, --fr-reg-upper-strength
  Sim (sim)      --infer, --load-config, --load-weights, --max-samples
  Output / exec  --out-dir, --wipe-dir

Examples:
  python -m cli                                    # sim (metrics only)
  python -m cli sim --input dataset --dataset mnist --digit 3
  python -m cli train --epochs 100
  python -m cli sim --infer --load-weights weights.pth --dt 0.5
  python -m cli sim --infer --load-config runs/foo/config.json --load-weights runs/foo/weights.pth

Models:
  ping        COBANet with E↔I coupling. With --ei-strength > 0 the
              recurrent inhibitory loop is wired up and frozen at init;
              feedforward weights train against this fixed substrate.
              (With --ei-strength 0 the I-loop is silenced — E cells only.)

Voltage-gradient damping scales the surrogate-gradient contribution to the
membrane-voltage update.
"""
    parser = argparse.ArgumentParser(
        prog="pinglab-cli",
        description="pinglab-cli — PING network toolkit",
        epilog=_examples,
        formatter_class=argparse.RawDescriptionHelpFormatter,
    )

    parent = _build_parent_parser()

    _build_subparsers(parser, parent)

    args = parser.parse_args(argv)
    if args.mode is None:
        parser.print_help()
        sys.exit(0)

    try:
        apply_bundle_to_args(args, argv)
    except BundleCompatibilityError as exc:
        parser.error(str(exc))

    # --load-config: load training params from config.json, fill unset values
    if args.mode in ("sim", "dump-weights") and getattr(args, "load_config", None):
        config_to_args, dest_to_flag = _build_config_mapping(parent)
        if args.mode == "sim":
            # A training config's max_samples caps its training partition. In
            # simulation, --max-samples instead caps evaluation and must be an
            # explicit choice; otherwise inference uses the full official test
            # partition.
            config_to_args.pop("max_samples", None)
        # --tau-gaba lives on the sim subparser, not the parent parser that
        # _build_config_mapping walks, so it isn't auto-mapped. Register it here so a
        # loaded cell replays at its TRAINED τ_GABA (config key tau_gaba_ms), not the
        # models.py module default. Without this a cell trained at a non-default
        # τ_GABA (e.g. the 6 ms canonical) loaded via --load-config would silently
        # evaluate at the module default — a train/eval mismatch.
        config_to_args["tau_gaba_ms"] = "tau_gaba"
        dest_to_flag.setdefault("tau_gaba", "--tau-gaba")
        _apply_load_config(args, argv, config_to_args, dest_to_flag)
        if getattr(args, "bundle", None) and not any(
            item == "--bundle" or item.startswith("--bundle=") for item in argv
        ):
            parser.error(
                "bundle paths are not inherited from --load-config; pass "
                "--bundle explicitly without --load-config"
            )
    if getattr(args, "infer", False) and not getattr(args, "load_weights", None):
        print("Error: sim --infer requires --load-weights")
        sys.exit(1)

    # Auto-detect: if user explicitly passed --dataset/--digit/--sample but
    # left --input at the default "synthetic-spikes", flip to "dataset". The
    # explicit dataset flags only make sense in dataset input mode, so this
    # avoids the silent footgun where "image --dataset mnist --digit 0" went
    # through the synthetic-spikes branch and ignored the digit.
    def _flag_in_argv(*names):
        for arg in argv:
            for n in names:
                if arg == n or arg.startswith(n + "="):
                    return True
        return False

    args._input_auto = False
    config_set_dataset = (
        getattr(args, "load_config", None) and getattr(args, "dataset", None) == "mnist"
    )
    # Only auto-flip when --input was LEFT AT DEFAULT. An explicit --input
    # synthetic-spikes (e.g. uniform-Poisson f–I on a trained cell, which also
    # passes --load-config → dataset) must be honoured, not overridden.
    input_explicit = _flag_in_argv("--input")
    if (
        not input_explicit
        and args.input == "synthetic-spikes"
        and (_flag_in_argv("--dataset", "--digit", "--sample") or config_set_dataset)
    ):
        args.input = "dataset"
        args._input_auto = True

    return args

configure_models

View source

def configure_models(args)

Source docstring:

Apply CLI overrides to models.py globals — the one sanctioned boundary.

Model globals (M.SURROGATE_SLOPE, M.tau_snn, M.max_rate_hz, the dt-derived
constants, …) are kept as module globals so torch.compile specializes the
graph on them as constants, so they cannot live on the Config dataclass
([[project_models_globals_are_torch_compile_choice]]). This is the single
place CLI arguments are written into them. The data-dependent globals
(M.N_IN / M.N_HID / M.T_steps) are set later by train.py / scan.py, where
the dataset shape is known.

Declarative table — {arg_name: (M attribute, cast)} — applied only when the
arg was given. Add a new model global by adding a row, not an if-branch.
ParameterAnnotationDefaultMeaning
argsunannotatedrequiredDefined by the source contract and implementation below.
Implementation
def configure_models(args):
    """Apply CLI overrides to models.py globals — the one sanctioned boundary.

    Model globals (M.SURROGATE_SLOPE, M.tau_snn, M.max_rate_hz, the dt-derived
    constants, …) are kept as module globals so torch.compile specializes the
    graph on them as constants, so they cannot live on the Config dataclass
    ([[project_models_globals_are_torch_compile_choice]]). This is the single
    place CLI arguments are written into them. The data-dependent globals
    (M.N_IN / M.N_HID / M.T_steps) are set later by train.py / scan.py, where
    the dataset shape is known.

    Declarative table — {arg_name: (M attribute, cast)} — applied only when the
    arg was given. Add a new model global by adding a row, not an if-branch.
    """
    arg_to_global = {
        "surrogate_slope": ("SURROGATE_SLOPE", float),
        "tau_gaba": ("tau_gaba", float),
    }
    for arg, (attr, cast) in arg_to_global.items():
        val = getattr(args, arg, None)
        if val is not None:
            setattr(M, attr, cast(val))
    # Special case: a bare flag.
    M.EXACT_K_INITIALIZATION = bool(getattr(args, "exact_k_initialization", False))
    # Input Poisson rate and trial duration — single source of truth. Every
    # code path reads M.max_rate_hz / M.T_ms, so setting them here once means
    # all dispatch branches (sim/train × all input types) respect
    # --input-rate / --t-ms. Subfunctions that change dt recalc M.T_steps locally.
    M.max_rate_hz = args.spike_rate
    M.T_ms = args.t_ms

save_run_artifacts

View source

def save_run_artifacts(out_dir, args, mode)

Source docstring:

Save config.json (with provenance), run.sh, set up logging, print intro.
ParameterAnnotationDefaultMeaning
out_dirunannotatedrequiredDefined by the source contract and implementation below.
argsunannotatedrequiredDefined by the source contract and implementation below.
modeunannotatedrequiredDefined by the source contract and implementation below.

Return expressions (branch-dependent; names refer to the linked implementation):

log
Implementation
def save_run_artifacts(out_dir, args, mode):
    """Save config.json (with provenance), run.sh, set up logging, print intro."""
    import json
    import logging

    from snnlab.sim import runlog

    out_dir = Path(out_dir)
    if args.wipe_dir and out_dir.exists():
        import shutil

        shutil.rmtree(out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)

    # config.json — with provenance metadata at top
    config = {"mode": mode}
    config.update(runlog.provenance())
    for k, v in vars(args).items():
        if v is not None:
            config[k] = v
    config.update(duration_metadata(args.t_ms, args.dt))
    config.update(
        refractory_metadata(
            getattr(args, "refractory_e_ms", 3.0),
            getattr(args, "refractory_i_ms", 1.5),
            args.dt,
            policy=getattr(args, "refractory_policy", "nearest"),
        )
    )
    with open(out_dir / "config.json", "w") as f:
        json.dump(config, f, indent=2, default=str)

    # run.sh
    with open(out_dir / "run.sh", "w") as f:
        f.write("#!/bin/bash\n")
        f.write(" ".join(sys.argv) + "\n")

    # output.log — file handler strips ANSI, stdout keeps it
    log = logging.getLogger("cli")
    log.setLevel(logging.DEBUG)
    # Close (not just drop) any handlers from a prior run in the same process —
    # list.clear() would leak the previous output.log file handle. Matters in
    # tests, which call this many times per process.
    for _h in log.handlers:
        _h.close()
    log.handlers.clear()

    class _StripAnsiFormatter(logging.Formatter):
        def format(self, record):
            msg = super().format(record)
            return runlog._strip_ansi(msg)

    fh = logging.FileHandler(out_dir / "output.log", mode="w")
    fh.setFormatter(_StripAnsiFormatter("%(message)s"))
    log.addHandler(fh)
    sh = logging.StreamHandler(sys.stdout)
    sh.setFormatter(logging.Formatter("%(message)s"))
    log.addHandler(sh)

    # run.jsonl — the canonical, machine-readable event spine for this run.
    runlog.init_events(out_dir)

    # Print structured intro
    _print_intro(log, config, args, mode)

    return log

main

View source

def main(argv=None)

Source docstring:

Parse args and run the requested mode. Returns a process exit code.
ParameterAnnotationDefaultMeaning
argvunannotatedNoneDefined by the source contract and implementation below.

Return expressions (branch-dependent; names refer to the linked implementation):

0

Explicit exceptions in this implementation; called helpers may raise additional errors:

Explicit exception expression
SystemExit(f'{mode} writes run artifacts and requires an explicit --out-dir')
SystemExit('--load-runtime-state/--save-runtime-state require --executor graph')
SystemExit('--event-file requires --executor graph')
SystemExit('--dataset-file requires --executor graph')
SystemExit('--poisson-protocol requires --executor graph')
SystemExit('--scale-projection requires --executor graph')
SystemExit('--intervention requires --executor graph')
SystemExit('--inference-timestep-ms requires --executor graph')
SystemExit('graph execution requires exactly one of --input-file, --event-file, --poisson-protocol, or --dataset-file')
SystemExit(str(exc))
ValueError('--scale-projection expects ID=FACTOR')
ValueError(f'--scale-projection repeats projection {projection_id!r}')
ValueError('--intervention expects drop:POPULATION=PROBABILITY or add:POPULATION=RATE_HZ')
ValueError('--dataset-file requires --dataset-encoder')
ValueError('--dataset-file requires --input-dataset-id and --input-split')
ValueError('multi-input dataset graphs require --dataset-input-id')
ValueError('graph training requires --target-file')
ValueError('graph dataset training requires --dataset-target-id')
ValueError('CLI Poisson generation requires exactly one graph input')
ValueError('categorical-rate Poisson requires --input-rates')
Implementation
def main(argv=None):
    """Parse args and run the requested mode. Returns a process exit code."""
    _t0 = _time.monotonic()

    argv = sys.argv[1:] if argv is None else list(argv)
    args = parse_args(argv)
    mode = args.mode

    if args.out_dir is None:
        raise SystemExit(
            f"{mode} writes run artifacts and requires an explicit --out-dir"
        )

    # Every invocation crosses the typed seam. The legacy callback below is
    # deliberately the unchanged handler body; graph requests use the data-only
    # bundle and stable request API directly.
    from snnlab.sim.execution import (
        ExecutionResult,
        execute_request,
        execution_spec_from_args,
    )

    request = execution_spec_from_args(args)

    if request.executor == "legacy" and (
        getattr(args, "load_runtime_state", None)
        or getattr(args, "save_runtime_state", None)
    ):
        raise SystemExit(
            "--load-runtime-state/--save-runtime-state require --executor graph"
        )
    if request.executor == "legacy" and getattr(args, "event_file", None):
        raise SystemExit("--event-file requires --executor graph")
    if request.executor == "legacy" and getattr(args, "dataset_file", None):
        raise SystemExit("--dataset-file requires --executor graph")
    if request.executor == "legacy" and getattr(args, "poisson_protocol", None):
        raise SystemExit("--poisson-protocol requires --executor graph")
    if request.executor == "legacy" and getattr(args, "scale_projection", None):
        raise SystemExit("--scale-projection requires --executor graph")
    if request.executor == "legacy" and getattr(args, "intervention", None):
        raise SystemExit("--intervention requires --executor graph")
    if request.executor == "legacy" and getattr(args, "inference_timestep_ms", None):
        raise SystemExit("--inference-timestep-ms requires --executor graph")

    if request.executor == "graph":
        from dataclasses import replace

        input_file = getattr(args, "input_file", None)
        event_file = getattr(args, "event_file", None)
        poisson_protocol = getattr(args, "poisson_protocol", None)
        dataset_file = getattr(args, "dataset_file", None)
        if (
            sum(
                bool(value)
                for value in (input_file, event_file, poisson_protocol, dataset_file)
            )
            != 1
        ):
            raise SystemExit(
                "graph execution requires exactly one of --input-file, --event-file, --poisson-protocol, or --dataset-file"
            )
        from snnlab.sim.bundle import load_graph_bundle, load_training_recipe
        from snnlab.sim.execution import (
            DatasetEncoder,
            DatasetSnapshotBinding,
            PoissonInputBinding,
            load_dense_array_bindings,
            load_event_stream_bindings,
            load_runtime_state,
            load_target_array_bindings,
            save_runtime_state,
            write_inference_artifacts,
        )

        manifest, graph = load_graph_bundle(args.bundle)
        try:
            inference_overrides = {}
            projection_scales = {}
            for item in getattr(args, "scale_projection", []):
                projection_id, separator, raw_factor = item.partition("=")
                if not separator or not projection_id or not raw_factor:
                    raise ValueError("--scale-projection expects ID=FACTOR")
                if projection_id in projection_scales:
                    raise ValueError(
                        f"--scale-projection repeats projection {projection_id!r}"
                    )
                projection_scales[projection_id] = float(raw_factor)
            if projection_scales:
                inference_overrides["projection_scales"] = projection_scales
            if getattr(args, "inference_timestep_ms", None) is not None:
                inference_overrides["timestep_ms"] = args.inference_timestep_ms
            interventions = []
            for item in getattr(args, "intervention", []):
                target, separator, raw_value = item.partition("=")
                kind, kind_separator, population_id = target.partition(":")
                if (
                    not separator
                    or not kind_separator
                    or not population_id
                    or not raw_value
                    or kind not in {"drop", "add"}
                ):
                    raise ValueError(
                        "--intervention expects drop:POPULATION=PROBABILITY or add:POPULATION=RATE_HZ"
                    )
                value = float(raw_value)
                interventions.append(
                    {
                        "kind": (
                            "drop_spikes" if kind == "drop" else "add_poisson_spikes"
                        ),
                        "population_id": population_id,
                        "probability" if kind == "drop" else "rate_hz": value,
                        "seed": args.seed,
                    }
                )
            if dataset_file:
                if not args.dataset_encoder:
                    raise ValueError("--dataset-file requires --dataset-encoder")
                if not args.input_dataset_id or not args.input_split:
                    raise ValueError(
                        "--dataset-file requires --input-dataset-id and --input-split"
                    )
                input_ids = [row["id"] for row in graph.get("inputs", [])]
                input_id = args.dataset_input_id or (
                    input_ids[0] if len(input_ids) == 1 else None
                )
                if input_id is None:
                    raise ValueError(
                        "multi-input dataset graphs require --dataset-input-id"
                    )
                encoder_kind = args.dataset_encoder.replace("-", "_")
                encoder = DatasetEncoder(
                    encoder_kind,
                    duration_ms=(
                        args.t_ms
                        if encoder_kind in {"rate_poisson", "event_bin"}
                        else None
                    ),
                    max_rate_hz=(
                        args.spike_rate if encoder_kind == "rate_poisson" else None
                    ),
                    seed=request.seed if encoder_kind == "rate_poisson" else 0,
                )
                binding_update = {
                    "dataset_binding": DatasetSnapshotBinding(
                        path=Path(dataset_file),
                        input_id=input_id,
                        target_id=args.dataset_target_id,
                        dataset_id=args.input_dataset_id,
                        split=args.input_split,
                        encoder=encoder,
                        feature_key=args.dataset_feature_key,
                        label_key=args.dataset_label_key,
                        sample_cap=args.max_samples,
                        shuffle=bool(args.input_shuffle),
                        order_seed=request.seed,
                    )
                }
            elif poisson_protocol:
                input_ids = [row["id"] for row in graph.get("inputs", [])]
                if len(input_ids) != 1:
                    raise ValueError(
                        "CLI Poisson generation requires exactly one graph input"
                    )
                dt_ms = float(graph["timebase"]["dt"]["value"])
                steps = duration_steps(args.t_ms, dt_ms)
                rates = (
                    args.input_rates
                    if poisson_protocol == "categorical-rate"
                    else [args.spike_rate]
                )
                if poisson_protocol == "categorical-rate" and not rates:
                    raise ValueError("categorical-rate Poisson requires --input-rates")
                binding_update = {
                    "poisson_bindings": (
                        PoissonInputBinding(
                            input_id=input_ids[0],
                            steps_count=int(steps),
                            batch_size=args.n_batch,
                            rates_hz=rates,
                            seed=args.seed,
                            categorical=poisson_protocol == "categorical-rate",
                        ),
                    )
                }
            else:
                binding_update = (
                    {"event_bindings": load_event_stream_bindings(event_file, graph)}
                    if event_file
                    else {
                        "input_bindings": load_dense_array_bindings(input_file, graph)
                    }
                )
            target_update = {}
            if request.kind == "train":
                target_file = getattr(args, "target_file", None)
                if not target_file and not dataset_file:
                    raise ValueError("graph training requires --target-file")
                if dataset_file and not args.dataset_target_id:
                    raise ValueError(
                        "graph dataset training requires --dataset-target-id"
                    )
                recipe = load_training_recipe(args.bundle, manifest, graph)
                target_update = {
                    "training": recipe,
                    **(
                        {
                            "target_bindings": load_target_array_bindings(
                                target_file, recipe
                            )
                        }
                        if target_file
                        else {}
                    ),
                }
        except ValueError as exc:
            raise SystemExit(str(exc)) from exc
        runtime_state = (
            load_runtime_state(args.load_runtime_state, device=request.device)
            if getattr(args, "load_runtime_state", None)
            else None
        )
        request = replace(
            request,
            **binding_update,
            **target_update,
            protocol={
                "dataset": {
                    key: value
                    for key, value in {
                        "identity": getattr(args, "input_dataset_id", None),
                        "split": getattr(args, "input_split", None),
                        "sample_cap": getattr(args, "max_samples", None),
                        "shuffle": getattr(args, "input_shuffle", None),
                    }.items()
                    if value is not None
                }
            },
            options={
                **request.options,
                "shuffle": bool(getattr(args, "input_shuffle", False)),
                **(
                    {"inference_overrides": inference_overrides}
                    if inference_overrides
                    else {}
                ),
                **({"inference_interventions": interventions} if interventions else {}),
            },
            runtime_state=runtime_state,
        )
        result = execute_request(request)
        out_dir = Path(args.out_dir)
        out_dir.mkdir(parents=True, exist_ok=True)
        write_inference_artifacts(out_dir, result, graph=graph, seed=request.seed)
        if getattr(args, "save_runtime_state", None):
            assert result.runtime_state is not None
            save_runtime_state(args.save_runtime_state, result.runtime_state)
        return 0

    # Build config for non-train modes (build_config syncs the module aliases).
    if mode != "train":
        build_config(args)

    # Apply CLI overrides to models.py globals (all modes, incl. train).
    configure_models(args)

    from snnlab.sim import config as C

    # Determine output directory
    out_dir = Path(args.out_dir)

    # Save run artifacts for all modes
    log = save_run_artifacts(out_dir, args, mode)

    if args._input_auto:
        runlog.phase(log, "input", "auto → dataset (from --dataset/--digit/--sample)")

    def _legacy_request():
        _MODE_HANDLERS[mode](args, C, out_dir, log)
        return ExecutionResult(
            executor="legacy", metrics={"request": request.kind, "routing": "legacy"}
        )

    execute_request(request, legacy=_legacy_request)

    if mode in {"sim", "train"}:
        # Supplied inputs may determine the actual trial length. Keep the
        # original request alongside the number of steps that really ran.
        config_path = out_dir / "config.json"
        completed_config = json.loads(config_path.read_text())
        completed_config.update(duration_metadata(args.t_ms, args.dt))
        completed_config["duration_steps"] = int(M.T_steps)
        completed_config["realized_duration_ms"] = M.T_steps * float(args.dt)
        config_path.write_text(json.dumps(completed_config, indent=2, default=str))

    _elapsed = _time.monotonic() - _t0
    # No device here: it would be a guess. train's summary reports the device it
    # actually used; a sim's device is implicit in its (CPU) result path.
    runlog.done(log, _elapsed)
    runlog.close_events()
    return 0

Constants and type aliases

Initial source expressions are shown, not evaluated runtime values. Legacy configuration may mutate module defaults.

NameAnnotationInitial expressionSource
logunannotatedlogging.getLogger('cli')Source

On this page