snnlab
API referencesnnlab.sim

snnlab.sim.train

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

Back to sim reference

Training driver for the CLI.

Holds seed_everything and the main train() loop for the PING (COBANet) model on MNIST with a configurable readout, gradient stabilizer, and optimizer.

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
seed_everythingfunction
ShdBinnedDatasetclass
trainfunction

seed_everything

View source

def seed_everything(seed)

Source docstring:

Seed Python, NumPy, and torch RNGs for reproducible runs.

Seeds cover: Python `random`, NumPy global, torch CPU, torch CUDA
(all devices), and torch MPS (via torch.manual_seed, which fans out
to the active backend). Call before dataset load and model init.
ParameterAnnotationDefaultMeaning
seedunannotatedrequiredSeed controlling this operation’s random stream.

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

None
Implementation
def seed_everything(seed):
    """Seed Python, NumPy, and torch RNGs for reproducible runs.

    Seeds cover: Python `random`, NumPy global, torch CPU, torch CUDA
    (all devices), and torch MPS (via torch.manual_seed, which fans out
    to the active backend). Call before dataset load and model init.
    """
    if seed is None:
        return
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)
    log.info(f"  seed={seed} (Python, NumPy, torch)")

ShdBinnedDataset

View source

Source docstring:

Bin SHD event lists to dense (T_steps, n_in) spike tensors on the fly.

SHD ships as per-utterance events (spike time in seconds + unit index in
[0, 700)). Densifying the whole set at the model's native dt would be tens of
GB, so each sample is binned lazily here using the CURRENT module globals
M.T_steps / M.dt — both pinned by set_sim_dt() before the loader runs, so a
fork'd DataLoader worker sees the run's values. The returned (T_steps, n_in)
tensor stacks to (B, T_steps, n_in), which encode_batch() passes straight
through (transposed to (T_steps, B, n_in)) as already-spiked input.

Bases: torch.utils.data.Dataset[tuple[torch.Tensor, int]]. Inherited third-party framework APIs follow their owning library.

Constructor:

ShdBinnedDataset(self, events, labels, n_in)
ParameterAnnotationDefaultMeaning
eventsunannotatedrequiredDefined by the constructor implementation below.
labelsunannotatedrequiredDefined by the constructor implementation below.
n_inunannotatedrequiredDefined by the constructor implementation below.

Instance members assigned by the constructor (expressions are evaluated when constructed):

MemberAnnotationInitial expression
eventsunannotatedevents
labelsunannotatedlabels
n_inunannotatedn_in

ShdBinnedDataset.len

View source

def ShdBinnedDataset.__len__(self)

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

len(self.labels)
Implementation
def __len__(self):
        return len(self.labels)

ShdBinnedDataset.getitem

View source

def ShdBinnedDataset.__getitem__(self, index) -> tuple[torch.Tensor, int]
ParameterAnnotationDefaultMeaning
indexunannotatedrequiredDefined by the source contract and implementation below.

Return annotation: tuple[torch.Tensor, int].

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

(x, int(self.labels[index]))
Implementation
def __getitem__(self, index) -> tuple[torch.Tensor, int]:
        units, times = self.events[index]
        T = M.T_steps
        dt_s = M.dt / 1000.0  # bin width in seconds (dt is ms)
        bins = np.floor(times / dt_s).astype(np.int64)
        # Drop events past the sim window; clamp channels defensively.
        keep = (bins >= 0) & (bins < T) & (units >= 0) & (units < self.n_in)
        x = torch.zeros(T, self.n_in, dtype=torch.float32)
        x[
            torch.from_numpy(bins[keep]),
            torch.from_numpy(units[keep].astype(np.int64)),
        ] = 1.0
        return x, int(self.labels[index])
Complete class implementation
class ShdBinnedDataset(torch.utils.data.Dataset[tuple[torch.Tensor, int]]):
    """Bin SHD event lists to dense (T_steps, n_in) spike tensors on the fly.

    SHD ships as per-utterance events (spike time in seconds + unit index in
    [0, 700)). Densifying the whole set at the model's native dt would be tens of
    GB, so each sample is binned lazily here using the CURRENT module globals
    M.T_steps / M.dt — both pinned by set_sim_dt() before the loader runs, so a
    fork'd DataLoader worker sees the run's values. The returned (T_steps, n_in)
    tensor stacks to (B, T_steps, n_in), which encode_batch() passes straight
    through (transposed to (T_steps, B, n_in)) as already-spiked input.
    """

    def __init__(self, events, labels, n_in):
        self.events = events  # object array; events[i] = (units_i16, times_f32)
        self.labels = labels
        self.n_in = n_in

    def __len__(self):
        return len(self.labels)

    def __getitem__(self, index) -> tuple[torch.Tensor, int]:
        units, times = self.events[index]
        T = M.T_steps
        dt_s = M.dt / 1000.0  # bin width in seconds (dt is ms)
        bins = np.floor(times / dt_s).astype(np.int64)
        # Drop events past the sim window; clamp channels defensively.
        keep = (bins >= 0) & (bins < T) & (units >= 0) & (units < self.n_in)
        x = torch.zeros(T, self.n_in, dtype=torch.float32)
        x[
            torch.from_numpy(bins[keep]),
            torch.from_numpy(units[keep].astype(np.int64)),
        ] = 1.0
        return x, int(self.labels[index])

train

View source

def train(model_name='ping', lr=0.01, weight_decay=0.0, epochs=100, dt=0.1, out_dir=None, device_name=None, w_in=None, w_ee=None, w_ei=None, w_ie=None, w_ii=None, ei_strength=None, ei_ratio=2.0, w_in_initial_zero_fraction=0.0, recurrent_initial_zero_fraction=0.0, dataset='mnist', snapshot_init=True, snapshot_end=True, t_ms=200.0, hidden_sizes=None, max_samples=None, v_grad_dampen=80.0, dales_law=True, batch_size=None, seed=None, readout_w_out_scale=1.0, readout_w_init_mean=None, readout_w_init_std=None, readout_mode='rate', signed_readout=False, readout_bias=False, input_rates=None, tau_gaba=None, fr_reg_upper_target_hz=0.0, fr_reg_upper_strength=0.0, trainable_w_ee=False, trainable_w_ei=False, trainable_w_ie=False, trainable_w_ii=False, state_clamp=False, train_leak=False, tau_m_e_bounds_ms=None, tau_m_i_bounds_ms=None, adaptive_threshold=False, adapt_tau_bounds_ms=None, adapt_strength_init_mv=1.0, adapt_strength_max_mv=None, refractory_e_ms=None, refractory_i_ms=None, refractory_policy='nearest')

Source docstring:

Train a model on a supported dataset.
ParameterAnnotationDefaultMeaning
model_nameunannotated'ping'Defined by the source contract and implementation below.
lrunannotated0.01Learning rate; frozen groups require zero and trainable groups require a positive value.
weight_decayunannotated0.0Defined by the source contract and implementation below.
epochsunannotated100Number of training passes over the selected presentations.
dtunannotated0.1Timestep; authoring uses a Quantity and legacy simulation uses milliseconds.
out_dirunannotatedNoneDefined by the source contract and implementation below.
device_nameunannotatedNoneDefined by the source contract and implementation below.
w_inunannotatedNoneDefined by the source contract and implementation below.
w_eeunannotatedNoneDefined by the source contract and implementation below.
w_eiunannotatedNoneDefined by the source contract and implementation below.
w_ieunannotatedNoneDefined by the source contract and implementation below.
w_iiunannotatedNoneDefined by the source contract and implementation below.
ei_strengthunannotatedNoneDefined by the source contract and implementation below.
ei_ratiounannotated2.0Defined by the source contract and implementation below.
w_in_initial_zero_fractionunannotated0.0Defined by the source contract and implementation below.
recurrent_initial_zero_fractionunannotated0.0Defined by the source contract and implementation below.
datasetunannotated'mnist'Defined by the source contract and implementation below.
snapshot_initunannotatedTrueDefined by the source contract and implementation below.
snapshot_endunannotatedTrueDefined by the source contract and implementation below.
t_msunannotated200.0Physical duration in milliseconds.
hidden_sizesunannotatedNoneDefined by the source contract and implementation below.
max_samplesunannotatedNoneDefined by the source contract and implementation below.
v_grad_dampenunannotated80.0Defined by the source contract and implementation below.
dales_lawunannotatedTrueDefined by the source contract and implementation below.
batch_sizeunannotatedNoneNumber of presentations in a binding or mini-batch.
seedunannotatedNoneSeed controlling this operation’s random stream.
readout_w_out_scaleunannotated1.0Defined by the source contract and implementation below.
readout_w_init_meanunannotatedNoneDefined by the source contract and implementation below.
readout_w_init_stdunannotatedNoneDefined by the source contract and implementation below.
readout_modeunannotated'rate'Defined by the source contract and implementation below.
signed_readoutunannotatedFalseDefined by the source contract and implementation below.
readout_biasunannotatedFalseDefined by the source contract and implementation below.
input_ratesunannotatedNoneDefined by the source contract and implementation below.
tau_gabaunannotatedNoneDefined by the source contract and implementation below.
fr_reg_upper_target_hzunannotated0.0Defined by the source contract and implementation below.
fr_reg_upper_strengthunannotated0.0Defined by the source contract and implementation below.
trainable_w_eeunannotatedFalseDefined by the source contract and implementation below.
trainable_w_eiunannotatedFalseDefined by the source contract and implementation below.
trainable_w_ieunannotatedFalseDefined by the source contract and implementation below.
trainable_w_iiunannotatedFalseDefined by the source contract and implementation below.
state_clampunannotatedFalseDefined by the source contract and implementation below.
train_leakunannotatedFalseDefined by the source contract and implementation below.
tau_m_e_bounds_msunannotatedNoneDefined by the source contract and implementation below.
tau_m_i_bounds_msunannotatedNoneDefined by the source contract and implementation below.
adaptive_thresholdunannotatedFalseDefined by the source contract and implementation below.
adapt_tau_bounds_msunannotatedNoneDefined by the source contract and implementation below.
adapt_strength_init_mvunannotated1.0Defined by the source contract and implementation below.
adapt_strength_max_mvunannotatedNoneDefined by the source contract and implementation below.
refractory_e_msunannotatedNoneDefined by the source contract and implementation below.
refractory_i_msunannotatedNoneDefined by the source contract and implementation below.
refractory_policyunannotated'nearest'Defined by the source contract and implementation below.

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

0.0
best_acc

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

Explicit exception expression
ValueError('fr_reg_upper_target_hz must be non-negative and finite')
ValueError('fr_reg_upper_strength must be non-negative and finite')
ValueError('readout_w_init_mean and readout_w_init_std must be specified together')
ValueError('readout_w_init_std must be non-negative')
ValueError('direct readout initialization cannot be combined with readout_w_out_scale')
ValueError('variable input rates are supported only for image datasets')
ValueError('input_rates must contain non-negative rates')
Implementation
def train(
    model_name="ping",
    lr=0.01,
    weight_decay=0.0,
    epochs=100,
    dt=0.1,
    out_dir=None,
    device_name=None,
    w_in=None,
    w_ee=None,
    w_ei=None,
    w_ie=None,
    w_ii=None,
    ei_strength=None,
    ei_ratio=2.0,
    w_in_initial_zero_fraction=0.0,
    recurrent_initial_zero_fraction=0.0,
    dataset="mnist",
    snapshot_init=True,
    snapshot_end=True,
    t_ms=200.0,
    hidden_sizes=None,
    max_samples=None,
    v_grad_dampen=80.0,
    dales_law=True,
    batch_size=None,
    seed=None,
    readout_w_out_scale=1.0,
    readout_w_init_mean=None,
    readout_w_init_std=None,
    readout_mode="rate",
    signed_readout=False,
    readout_bias=False,
    input_rates=None,
    tau_gaba=None,
    fr_reg_upper_target_hz=0.0,
    fr_reg_upper_strength=0.0,
    trainable_w_ee=False,
    trainable_w_ei=False,
    trainable_w_ie=False,
    trainable_w_ii=False,
    state_clamp=False,
    train_leak=False,
    tau_m_e_bounds_ms=None,
    tau_m_i_bounds_ms=None,
    adaptive_threshold=False,
    adapt_tau_bounds_ms=None,
    adapt_strength_init_mv=1.0,
    adapt_strength_max_mv=None,
    refractory_e_ms=None,
    refractory_i_ms=None,
    refractory_policy="nearest",
):
    """Train a model on a supported dataset."""
    from torch.utils.data import DataLoader, TensorDataset

    seed_everything(seed)

    if fr_reg_upper_target_hz < 0 or not math.isfinite(fr_reg_upper_target_hz):
        raise ValueError("fr_reg_upper_target_hz must be non-negative and finite")
    if fr_reg_upper_strength < 0 or not math.isfinite(fr_reg_upper_strength):
        raise ValueError("fr_reg_upper_strength must be non-negative and finite")

    # Setup dt and all derived constants
    # ─────────────────────────────────────────────────────────────────────────
    # M MODULE GLOBALS INITIALIZATION
    # ─────────────────────────────────────────────────────────────────────────
    # The models.py module (M) uses global state that must be set BEFORE building
    # the network or loading weights. This is a deliberate torch.compile choice:
    # dt-dependent constants are computed in forward() (not pre-computed globals),
    # but network topology constants (N_HID, N_INH, HIDDEN_SIZES) must be fixed.
    # ─────────────────────────────────────────────────────────────────────────

    # Pin dt / T_ms / T_steps to THIS run's values before any forward() call.
    # Without this, forward() falls back to the models.py defaults (dt = 0.25,
    # T_steps from the 1000 ms default) and the network trains at the wrong
    # timestep regardless of --dt — the 8befe44 regression this restores.
    set_sim_dt(dt, t_ms)

    # Optional τ_GABA override: forward() recomputes decay_gaba from the module
    # globals M.tau_gaba and M.dt each call, so setting M.tau_gaba here is what
    # actually takes effect. nb041 sweeps tau_gaba to vary the gamma (E-I
    # oscillation) frequency across models.
    if tau_gaba is not None:
        M.tau_gaba = float(tau_gaba)

    # Determine network hidden layer sizes: use CLI arg or smart default per dataset
    # Smart defaults: mnist=1024
    if hidden_sizes is None:
        default = DATASET_N_HIDDEN_DEFAULTS.get(dataset, 256)
        hidden_sizes = [default]
        log.info(f"  n_hidden auto → {hidden_sizes} (smart default for {dataset})")

    # Initialize M module globals: M.N_HID, M.N_INH, M.HIDDEN_SIZES
    # These are read by build_net and model.forward() to determine network topology.
    # Must be set before build_net() below, or network initialization will be wrong.
    setup_model_globals(hidden_sizes)
    M.V_GRAD_DAMPEN = v_grad_dampen
    if batch_size is not None:
        M.BATCH_SIZE = batch_size

    device = torch.device(device_name) if device_name else _auto_device()
    if input_rates is not None:
        if dataset not in IMAGE_DATASETS:
            raise ValueError(
                "variable input rates are supported only for image datasets"
            )
        input_rates = tuple(float(rate) for rate in input_rates)
        if not input_rates or any(rate < 0 for rate in input_rates):
            raise ValueError("input_rates must contain non-negative rates")
    rate_values = (
        torch.tensor(input_rates, dtype=torch.float32) if input_rates else None
    )
    validation_draw_count = VALIDATION_DRAW_COUNT if dataset in IMAGE_DATASETS else 1
    validation_encoder_seeds = VALIDATION_ENCODER_SEEDS[:validation_draw_count]
    validation_rate_seeds = VALIDATION_RATE_SEEDS[:validation_draw_count]
    train_rate_gen = torch.Generator().manual_seed((seed or 0) + 82_001)
    prediction_rate_gen = torch.Generator().manual_seed((seed or 0) + 82_003)

    def sample_input_rates(batch_size, generator):
        if rate_values is None:
            return None
        chosen = torch.randint(len(rate_values), (batch_size,), generator=generator)
        return rate_values[chosen].to(device)

    if out_dir is None:
        out_dir = Path.cwd() / "artifacts" / "training" / model_name
    else:
        out_dir = Path(out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)

    # Data — single canonical loader.
    runlog.phase(log, "loading", dataset)
    X_tr, X_te, y_tr, y_te = load_dataset(dataset, max_samples=max_samples, split=True)
    # Output classes are a property of the labels, not a hardcode. Set the module
    # global so build_net sizes W_ff[-1] correctly (mnist=10, shd=20).
    M.N_OUT = int(max(int(y_tr.max()), int(y_te.max()))) + 1
    bs = batch_size if batch_size is not None else BATCH_SIZE
    # Event datasets (SHD) arrive as an object array of per-sample event lists →
    # bin to (T_steps, n_in) spikes lazily. Image datasets (MNIST) arrive as
    # dense (N, n_in) pixel rows and Poisson-encode per batch. N_IN comes from the
    # data in both cases (784 for MNIST pixels, 700 for SHD channels).
    if getattr(X_tr, "dtype", None) is not None and X_tr.dtype.kind == "O":
        M.N_IN = SHD_N_IN
        train_ds = ShdBinnedDataset(X_tr, y_tr, n_in=SHD_N_IN)
        test_ds = ShdBinnedDataset(X_te, y_te, n_in=SHD_N_IN)
    else:
        M.N_IN = X_tr.shape[1]
        train_ds = TensorDataset(torch.from_numpy(X_tr), torch.from_numpy(y_tr))
        test_ds = TensorDataset(torch.from_numpy(X_te), torch.from_numpy(y_te))
    train_loader = DataLoader(train_ds, batch_size=bs, shuffle=True)
    test_loader = DataLoader(test_ds, batch_size=bs)

    # Model — symmetry-break standard-snn (dense W_in needs heterogeneous init phases).
    # Skip randomize_init when Kaiming init is used: Kaiming already gives
    # heterogeneous per-neuron tuning, so scattering mem phases is redundant.
    # Uniform randomize_init across all models — symmetry breaking matters
    # for all architectures. Only skip when kaiming (already heterogeneous).
    if (readout_w_init_mean is None) != (readout_w_init_std is None):
        raise ValueError(
            "readout_w_init_mean and readout_w_init_std must be specified together"
        )
    if readout_w_init_std is not None and readout_w_init_std < 0:
        raise ValueError("readout_w_init_std must be non-negative")
    if readout_w_init_mean is not None and readout_w_out_scale != 1.0:
        raise ValueError(
            "direct readout initialization cannot be combined with readout_w_out_scale"
        )
    if readout_w_init_mean is not None:
        assert readout_w_init_std is not None
    readout_w_init = (
        (readout_w_init_mean, readout_w_init_std)
        if readout_w_init_mean is not None
        else None
    )
    net = build_net(
        model_name,
        refractory_e_ms=refractory_e_ms,
        refractory_i_ms=refractory_i_ms,
        refractory_policy=refractory_policy,
        w_in=w_in,
        w_in_initial_zero_fraction=w_in_initial_zero_fraction,
        w_ee=w_ee,
        w_ei=w_ei,
        w_ie=w_ie,
        w_ii=w_ii,
        ei_strength=ei_strength,
        ei_ratio=ei_ratio,
        recurrent_initial_zero_fraction=recurrent_initial_zero_fraction,
        device=device,
        randomize_init=True,
        dales_law=dales_law,
        hidden_sizes=hidden_sizes,
        readout_mode=readout_mode,
        signed_readout=signed_readout,
        readout_bias=readout_bias,
        readout_w_init=readout_w_init,
        trainable_w_ee=trainable_w_ee,
        trainable_w_ei=trainable_w_ei,
        trainable_w_ie=trainable_w_ie,
        trainable_w_ii=trainable_w_ii,
        state_clamp=state_clamp,
        train_leak=train_leak,
        tau_m_e_bounds_ms=tau_m_e_bounds_ms,
        tau_m_i_bounds_ms=tau_m_i_bounds_ms,
        adaptive_threshold=adaptive_threshold,
        adapt_tau_bounds_ms=adapt_tau_bounds_ms,
        adapt_strength_init_mv=adapt_strength_init_mv,
        adapt_strength_max_mv=adapt_strength_max_mv,
    )
    timing_config = {
        **duration_metadata(t_ms, dt),
        **refractory_metadata(
            net.refractory_e_ms, net.refractory_i_ms, dt,
            policy=net.refractory_policy,
        ),
    }
    if readout_mode != "rate":
        log.info(f"  readout_mode={readout_mode}")
    if readout_w_out_scale != 1.0:
        with torch.no_grad():
            net.W_ff[-1].mul_(readout_w_out_scale)
            if hasattr(net, "b_ff") and len(net.b_ff) > 0:
                net.b_ff[-1].mul_(readout_w_out_scale)
        log.info(
            f"  readout_w_out_scale={readout_w_out_scale:g} "
            f"(W_ff[-1] and b_ff[-1] scaled at init)"
        )
        # Scaling is part of initialization, so provenance must describe the
        # tensor actually handed to the optimizer rather than the pre-scale draw.
        current = net.weight_final_statistics()["W_out"]
        recorded = net.weight_initialization["W_out"]
        recorded["scaling_convention"] = "fan_in_normalized_then_scaled"
        recorded["post_initialization_scale"] = float(readout_w_out_scale)
        recorded["statistics"]["initialization_zero_count"] = current["zero_count"]
        recorded["statistics"]["initialization_zero_fraction"] = current[
            "zero_fraction"
        ]
        recorded["statistics"]["all_entries"] = current["all_entries"]
        recorded["statistics"]["realized_column_sum"] = current["realized_column_sum"]
    readout_init = net.W_ff[-1].detach()
    readout_init_stats = {
        "mean": float(readout_init.mean()),
        "std": float(readout_init.std(unbiased=False)),
        "min": float(readout_init.min()),
        "max": float(readout_init.max()),
        "zero_fraction": float((readout_init == 0).float().mean()),
    }
    log.info(
        "  readout_w_init_realized="
        f"mean={readout_init_stats['mean']:.6g}, "
        f"std={readout_init_stats['std']:.6g}, "
        f"min={readout_init_stats['min']:.6g}, "
        f"max={readout_init_stats['max']:.6g}, "
        f"zero_fraction={readout_init_stats['zero_fraction']:.6g}"
    )
    if not dales_law:
        log.info("  dales_law=False (signed weights, no clamp)")
    n_params = sum(p.numel() for p in net.parameters())
    n_trainable = sum(p.numel() for p in net.parameters() if p.requires_grad)

    runlog.phase(
        log,
        "building",
        f"{n_trainable:,} params · {len(X_tr)} train · {len(X_te)} test",
    )

    # Save config for reproducibility. `mode` leads so every config.json (train
    # and sim/infer) carries it — a reader can dispatch on config["mode"] without
    # KeyError. Ignored by --load-config (not in the config→arg map).
    config = {
        "mode": "train",
        "model": model_name,
        "lr": lr,
        "weight_decay": weight_decay,
        "epochs": epochs,
        "dt": dt,
        "t_ms": t_ms,
        **timing_config,
        "dataset": dataset,
        "n_hidden": M.N_HID,
        "n_inh": M.N_INH,
        "n_in": M.N_IN,
        "n_out": M.N_OUT,
        "tau_ampa_ms": float(M.tau_ampa),
        "w_in": list(w_in) if w_in else None,
        "w_ee": list(w_ee) if w_ee else None,
        "ei_strength": ei_strength,
        "ei_ratio": ei_ratio,
        "w_in_initial_zero_fraction": w_in_initial_zero_fraction,
        "recurrent_initial_zero_fraction": recurrent_initial_zero_fraction,
        "input_rate": M.max_rate_hz,
        "input_rates": list(input_rates) if input_rates else None,
        "input_rate_sampling": "uniform_categorical_per_presentation"
        if input_rates
        else "fixed",
        "v_grad_dampen": v_grad_dampen,
        # The CLI applies this process-wide model constant before dispatching
        # into train().  Record the resolved value alongside the other
        # scientific parameters so a campaign can validate and replay it.
        "surrogate_slope": M.SURROGATE_SLOPE,
        "batch_size": bs,
        "grad_clip": GRAD_CLIP,
        "max_samples": max_samples,
        "dataset_split": (
            {
                "source_train_partition": "official_mnist_train",
                "source_test_partition": "official_mnist_test",
                "validation_fraction": 0.1,
                "split_seed": 42,
                "optimizer_train_samples": len(y_tr),
                "validation_samples": len(y_te),
                "official_test_samples": 10000,
                "checkpoint_selection_partition": "validation",
                "official_test_used_during_training": False,
            }
            if dataset == "mnist"
            else {"contract": "official_shd_train_test"}
        ),
        "validation_encoder_draws": {
            "count": validation_draw_count,
            "encoder_seeds": list(validation_encoder_seeds),
            "input_rate_seeds": list(validation_rate_seeds),
            "aggregation_unit": "validation_sample_then_encoder_draw",
            "checkpoint_selection": (
                "minimum_mean_cross_entropy; tie maximum_mean_accuracy; "
                "tie earliest_epoch"
            ),
        },
        "n_params": n_params,
        "n_trainable": n_trainable,
        "dales_law": dales_law,
        "readout_mode": readout_mode,
        "signed_readout": signed_readout,
        "readout_bias": readout_bias,
        "readout_w_init": {
            "distribution": "lower_clamped_normal",
            "units": "stored_weight",
            "mean": readout_w_init_mean,
            "std": readout_w_init_std,
        }
        if readout_w_init is not None
        else None,
        "readout_w_init_mean": readout_w_init_mean,
        "readout_w_init_std": readout_w_init_std,
        "readout_w_init_realized": readout_init_stats,
        "weight_initialization": net.weight_initialization,
        "readout_reduction": (
            M.CUMULATIVE_READOUT_REDUCTION
            if readout_mode == "cumulative-potential"
            else "output_spike_count"
            if readout_mode == "spike-count"
            else "output_spike_count_divided_by_duration_seconds"
            if readout_mode == "spike-rate"
            else None
        ),
        "readout_units": (
            "spikes"
            if readout_mode == "spike-count"
            else "Hz"
            if readout_mode == "spike-rate"
            else None
        ),
        "readout_tau_bounds_ms": (
            list(M.CUMULATIVE_READOUT_TAU_BOUNDS_MS)
            if readout_mode == "cumulative-potential"
            else None
        ),
        "readout_reference": (
            M.CUMULATIVE_READOUT_REFERENCE
            if readout_mode == "cumulative-potential"
            else None
        ),
        "hidden_sizes": hidden_sizes,
        "trainable_w_ee": trainable_w_ee,
        "trainable_w_ei": trainable_w_ei,
        "trainable_w_ie": trainable_w_ie,
        "trainable_w_ii": trainable_w_ii,
        "state_clamp": state_clamp,
        "train_leak": train_leak,
        "tau_m_e_bounds_ms": list(tau_m_e_bounds_ms)
        if tau_m_e_bounds_ms is not None
        else list(M.TRAINABLE_TAU_M_E_BOUNDS_MS),
        "tau_m_i_bounds_ms": list(tau_m_i_bounds_ms)
        if tau_m_i_bounds_ms is not None
        else list(M.TRAINABLE_TAU_M_I_BOUNDS_MS),
        "adaptive_threshold": adaptive_threshold,
        "adapt_tau_bounds_ms": list(adapt_tau_bounds_ms)
        if adapt_tau_bounds_ms is not None
        else list(M.ADAPT_TAU_BOUNDS_MS),
        "adapt_strength_init_mv": adapt_strength_init_mv,
        "adapt_strength_max_mv": adapt_strength_max_mv
        if adapt_strength_max_mv is not None
        else M.ADAPT_STRENGTH_MAX_MV,
        "seed": seed,
        "tau_gaba_ms": float(M.tau_gaba),
        # Swept / recipe-varying knobs — must be structured fields, not just
        # buried in run.sh: the firing-rate target is the independent variable
        # of the activity frontier.
        "fr_reg_upper_target_hz": fr_reg_upper_target_hz,
        "fr_reg_upper_strength": fr_reg_upper_strength,
        "fr_reg_contract": {
            "aggregation": "sample-wise-population-mean",
            "population": "hidden-excitatory",
            "units": "hz",
            "penalty": "one-sided-quadratic",
            "normalization": ["batch", "population-width", "duration", "hidden-layers"],
        },
        "readout_w_out_scale": (
            readout_w_out_scale if readout_w_init is None else None
        ),
        # Provenance (git SHA, run_id, started_at, device, torch version,
        # python env hash) — keeps train-mode config.json at parity with
        # sim/image modes.
        **runlog.provenance(),
    }
    with open(out_dir / "config.json", "w") as f:
        json.dump(config, f, indent=2)
    with open(out_dir / "run.sh", "w") as f:
        f.write("#!/bin/bash\n")
        f.write(" ".join(sys.argv) + "\n")

    # Train uses the module-level log (already set up by save_run_artifacts)

    # Pre-generate a fixed reference spike train so the init/end snapshots use the
    # same sample regardless of train/test shuffling. Image data: Poisson-encode a
    # fixed digit-0 sample. Event data (SHD): bin the first test utterance — there
    # is no pixel image to encode.
    if dataset in IMAGE_DATASETS:
        ref_pixel_vec, _ref_image = _load_dataset_image(
            dataset, digit_class=0, sample_idx=0
        )
        ref_input = torch.from_numpy(ref_pixel_vec).float().unsqueeze(0).to(device)
        torch.manual_seed(0)
        pixels = ref_input.clamp(0, 1)
        p = M.max_rate_hz * dt / 1000.0
        ref_spikes = (
            (torch.rand(M.T_steps, 1, M.N_IN, device=device) < pixels * p)
            .float()
            .squeeze(1)
        )
    else:
        # (T_steps, N_IN) already-binned spikes for the first test sample.
        ref_spikes = test_ds[0][0].to(device)

    if snapshot_init:
        C.cfg.n_e = M.N_HID
        C.cfg.n_i = M.N_INH
        net.recording = True
        with torch.no_grad():
            net(input_spikes=ref_spikes)
        net.recording = False
        rec = net.spike_record

        spk_e = _to_np(rec[primary_hid_key(rec)])
        _ik = primary_inh_key(rec)
        spk_i = _to_np(rec[_ik]) if _ik else None

        snapshot_init_state = compute_metrics(
            spk_e, spk_i, dt, model_name, n_e=M.N_HID, n_i=M.N_INH
        )
        runlog.metrics_line(log, snapshot_init_state, label="init")

    else:
        snapshot_init_state = None

    if epochs == 0:
        log.info("  epochs=0, use --epochs N to train")
        # Even in probe mode, write a minimal metrics.json so callers can
        # inspect init state without parsing logs.
        from datetime import datetime, timezone

        metrics_blob = {
            "mode": "train",
            "model": model_name,
            "run_finished_at": datetime.now(timezone.utc).isoformat(timespec="seconds"),
            "config": {
                "dt": dt,
                "t_ms": t_ms,
                **timing_config,
                "epochs": 0,
                "lr": lr,
                "input_rate": M.max_rate_hz,
                "w_in": list(w_in) if w_in else None,
                "w_in_initial_zero_fraction": w_in_initial_zero_fraction,
                "recurrent_initial_zero_fraction": recurrent_initial_zero_fraction,
                "ei_strength": ei_strength,
                "ei_ratio": ei_ratio,
                "n_hidden": M.N_HID,
                "n_inh": M.N_INH,
                "n_in": M.N_IN,
                "max_samples": max_samples,
                "dataset": dataset,
                "v_grad_dampen": v_grad_dampen,
                "n_params": n_params,
                "n_trainable": n_trainable,
            },
            "init": snapshot_init_state,
            "epochs": [],
            "end": None,
            "best_acc": 0.0,
            "best_epoch": 0,
            "total_elapsed_s": 0.0,
        }
        with open(out_dir / "metrics.json", "w") as f:
            json.dump(metrics_blob, f, indent=2, default=float)
        log.info(f"  done (probe only) \u2192 {out_dir}")
        return 0.0

    # AdamW (decoupled weight decay); identical to Adam when weight_decay=0, so
    # the default leaves every existing run unchanged. Weight decay bounds the
    # free signed recurrent weights (W_ee especially) that otherwise grow without
    # limit under --no-dales-law and tip the forward pass into NaN divergence —
    # the constraint Dale's law provides implicitly via its non-negativity clamp.
    opt = torch.optim.AdamW(net.parameters(), lr=lr, weight_decay=weight_decay)
    # Dale's law via projected gradient: clamp the constrained weights back onto
    # the non-negative orthant after every step. Registered once as a step-post
    # hook so it can't be forgotten (no-op when --no-dales-law allows signed
    # weights). Runs eager, outside the compiled graph.
    if not net.signed_weights:
        opt.register_step_post_hook(lambda *_: net.project_dales())
    _ce = torch.nn.CrossEntropyLoss()

    def loss_fn(logits, y):
        return _ce(logits, y)

    # Training loop

    best_acc = 0.0
    best_validation_loss = float("inf")
    best_state = None
    no_improve = 0
    t_start = _time.perf_counter()
    prev_lr = lr
    epoch_records: list = []  # accumulated for metrics.json
    best_epoch = 0
    wtracker = runlog.WarningTracker()
    jsonl = runlog.MetricsJsonl(out_dir / "metrics.jsonl")
    # The first training batch pays the one-time torch.compile cost (models.py
    # compiles the per-timestep body lazily on first call). Warn so the opening
    # silence reads as "compiling", not "hung". Skip when compile is disabled.
    if os.environ.get("PINGLAB_NO_COMPILE") != "1":
        runlog.phase(log, "compiling", "first batch · one-time · ~10–60s")
    runlog.epoch_header(log)
    # MPS has no max_memory_allocated — sample per epoch and keep the peak.
    # CUDA tracks peak natively (queried at end of run).
    peak_mem_mps = 0
    if device.type == "cuda":
        torch.cuda.reset_peak_memory_stats(device)

    # Running mean of epoch wall-times for a stable ETA. The first epoch is an
    # outlier (compile + warm-up), so once we have ≥2 epochs we average only the
    # post-compile ones; ETA stops lurching after epoch 2.
    epoch_times: list[float] = []

    for epoch in range(epochs):
        t_epoch = _time.perf_counter()

        # Train
        net.train()
        total_loss = 0.0
        n_batches = 0
        grad_sum = 0.0
        grad_max = 0.0            # peak pre-clip global grad norm this epoch
        n_grad = 0
        n_skipped_steps = 0
        n_nan_forward = 0         # batches whose forward pass returned NaN logits
        layer_ratio_sum = {}
        grad_norm_sum = {}
        n_samples_train = 0
        t_train_compute = _time.perf_counter()
        n_train_batches = len(train_loader)
        hb = runlog.Heartbeat()
        for batch_idx, (X_b, y_b) in enumerate(train_loader):
            X_b, y_b = X_b.to(device), y_b.to(device)
            spk = encode_batch(
                X_b,
                dt,
                max_rate_hz=sample_input_rates(len(X_b), train_rate_gen),
            )
            logits = net(input_spikes=spk)
            if torch.isnan(logits).any():
                opt.zero_grad()
                n_nan_forward += 1
                continue
            loss = loss_fn(logits, y_b)
            spike_counts = getattr(net, "last_spike_counts", None)
            if spike_counts is not None:
                # Quadratic firing-rate regularizer: penalise overshoot above θ_u.
                if fr_reg_upper_strength > 0:
                    loss = loss + _firing_rate_penalty(
                        spike_counts,
                        fr_reg_upper_target_hz,
                        fr_reg_upper_strength,
                        (M.T_steps * dt) / 1000.0,
                    )
            opt.zero_grad()
            loss.backward()
            for pname, p in net.named_parameters():
                if p.grad is None or pname.startswith("b_") or pname.endswith(".bias"):
                    continue
                gnorm = p.grad.norm().item()
                grad_norm_sum[pname] = grad_norm_sum.get(pname, 0.0) + gnorm
                wn = p.norm().item()
                if wn > 0:
                    layer_ratio_sum[pname] = (
                        layer_ratio_sum.get(pname, 0.0) + gnorm / wn
                    )
            gn = torch.nn.utils.clip_grad_norm_(net.parameters(), GRAD_CLIP)
            gn_f = float(gn)
            if not math.isfinite(gn_f):
                opt.zero_grad(set_to_none=True)
                n_skipped_steps += 1
            else:
                opt.step()  # project_dales runs via the step-post hook
                total_loss += loss.item()
                n_batches += 1
                grad_sum += gn_f
                grad_max = max(grad_max, gn_f)
                n_grad += 1
            n_samples_train += y_b.size(0)
            # Within-epoch heartbeat: a live partial row (running loss under the
            # loss column, batch counter in the dynamics gap, elapsed under dt).
            running_loss = total_loss / max(n_batches, 1)
            hb.beat(
                log,
                runlog.epoch_progress(
                    epoch + 1,
                    epochs,
                    f"batch {batch_idx + 1}/{n_train_batches}",
                    _time.perf_counter() - t_train_compute,
                    loss=running_loss,
                ),
            )
        train_compute_s = _time.perf_counter() - t_train_compute
        if device.type == "mps":
            peak_mem_mps = max(peak_mem_mps, torch.mps.current_allocated_memory())
        avg_grad = grad_sum / max(n_grad, 1)
        grad_ratios = {n: s / max(n_grad, 1) for n, s in layer_ratio_sum.items()}
        grad_norms = {n: s / max(n_grad, 1) for n, s in grad_norm_sum.items()}

        # Validation uses a fixed panel of independent Poisson draws. Reusing
        # the same panel each epoch keeps checkpoint comparisons paired and
        # reproducible while estimating performance over encoder noise rather
        # than over one privileged realization.
        t_eval = _time.perf_counter()
        net.eval()
        correct = total = 0
        test_loss_sum = 0.0
        test_batches = 0
        # Accumulate per-cell rate (Hz) over the validation set so the
        # per-epoch metrics record reflects validation means rather than
        # only the single-trial observation rate.
        test_rate_e_sum = 0.0
        test_rate_i_sum = 0.0
        # Logit-discrimination accumulators. Cross-entropy keeps falling after
        # accuracy plateaus because it rewards margin, not just correctness
        # (nb024): track how separated/confident the logits are per epoch so
        # that confidence inflation can be told apart from accuracy gains.
        margin_sum = 0.0       # mean (z_true - z_runner_up), the decision margin
        conf_sum = 0.0         # mean softmax prob of the true class
        logit_scale_sum = 0.0  # mean |logit|, the raw logit magnitude
        output_spike_sum = 0.0
        output_silent_samples = 0
        output_sample_count = 0
        output_class_spikes = torch.zeros(M.N_OUT, dtype=torch.float64)
        output_by_rate: dict[float, dict[str, float]] = {}
        n_test_batches = len(test_loader)
        validation_draws = []
        with torch.no_grad():
            for draw_index, (encoder_seed, rate_seed) in enumerate(
                zip(validation_encoder_seeds, validation_rate_seeds, strict=True)
            ):
                eval_gen = torch.Generator().manual_seed(encoder_seed)
                draw_rate_gen = torch.Generator().manual_seed(rate_seed)
                draw_correct = draw_total = 0
                draw_loss_sum = 0.0
                for X_b, y_b in test_loader:
                    X_b, y_b = X_b.to(device), y_b.to(device)
                    sampled_rates = sample_input_rates(len(X_b), draw_rate_gen)
                    spk = encode_batch(
                        X_b,
                        dt,
                        generator=eval_gen,
                        max_rate_hz=sampled_rates,
                    )
                    logits_t = net(input_spikes=spk)
                    B = y_b.size(0)
                    batch_loss_sum = loss_fn(logits_t, y_b).item() * B
                    draw_loss_sum += batch_loss_sum
                    test_loss_sum += batch_loss_sum
                    test_batches += 1
                    batch_correct = (logits_t.argmax(1) == y_b).sum().item()
                    draw_correct += batch_correct
                    correct += batch_correct
                    draw_total += B
                    total += B
                    # Per-sample margin, confidence and logit scale.
                    z_true = logits_t.gather(1, y_b.unsqueeze(1)).squeeze(1)
                    z_other = logits_t.clone()
                    z_other.scatter_(1, y_b.unsqueeze(1), float("-inf"))
                    z_runner = z_other.max(1).values
                    margin_sum += float((z_true - z_runner).sum().item())
                    conf_sum += float(
                        torch.softmax(logits_t, dim=1)
                        .gather(1, y_b.unsqueeze(1))
                        .sum()
                        .item()
                    )
                    logit_scale_sum += float(logits_t.abs().mean(1).sum().item())
                    if readout_mode in ("spike-count", "spike-rate"):
                        output_counts = net.last_output_spike_counts.detach().cpu()
                        per_sample_counts = output_counts.sum(dim=1)
                        output_spike_sum += float(per_sample_counts.sum().item())
                        output_silent_samples += int((per_sample_counts == 0).sum().item())
                        output_sample_count += B
                        output_class_spikes += output_counts.sum(dim=0).to(torch.float64)
                        rates = (
                            sampled_rates.detach().cpu()
                            if sampled_rates is not None
                            else torch.full((B,), float(M.max_rate_hz))
                        )
                        predictions = logits_t.argmax(1).detach().cpu()
                        labels = y_b.detach().cpu()
                        for rate in torch.unique(rates).tolist():
                            mask = rates == rate
                            bucket = output_by_rate.setdefault(
                                float(rate),
                                {"n_samples": 0.0, "n_correct": 0.0,
                                 "n_silent": 0.0, "n_spikes": 0.0},
                            )
                            bucket["n_samples"] += int(mask.sum().item())
                            bucket["n_correct"] += int(
                                (predictions[mask] == labels[mask]).sum().item()
                            )
                            bucket["n_silent"] += int(
                                (per_sample_counts[mask] == 0).sum().item()
                            )
                            bucket["n_spikes"] += float(
                                per_sample_counts[mask].sum().item()
                            )
                    # net.rates is set by _set_meta after every forward pass;
                    # values are already per-cell Hz averaged over the batch.
                    batch_rates = getattr(net, "rates", None) or {}
                    for k, v in batch_rates.items():
                        if k.startswith("hid"):
                            test_rate_e_sum += float(v) * B
                        elif k.startswith("inh"):
                            test_rate_i_sum += float(v) * B
                    hb.beat(
                        log,
                        runlog.epoch_progress(
                            epoch + 1,
                            epochs,
                            f"eval draw {draw_index + 1}/{validation_draw_count} "
                            f"batch {(test_batches - 1) % n_test_batches + 1}/"
                            f"{n_test_batches}",
                            _time.perf_counter() - t_eval,
                        ),
                    )
                validation_draws.append(
                    {
                    "draw": draw_index + 1,
                    "encoder_seed": encoder_seed,
                    "input_rate_seed": rate_seed,
                    "n_samples": draw_total,
                    "accuracy_pct": 100.0 * draw_correct / draw_total,
                    "cross_entropy": draw_loss_sum / draw_total,
                    }
                )

        eval_s = _time.perf_counter() - t_eval
        hb.clear()  # erase the live progress line; the finished row prints in its place

        acc = 100.0 * correct / total
        avg_train = total_loss / max(n_batches, 1)
        avg_test = test_loss_sum / max(total, 1)
        test_rate_e = test_rate_e_sum / total if total else 0.0
        test_rate_i = test_rate_i_sum / total if total else 0.0
        test_margin = margin_sum / total if total else 0.0
        test_confidence = conf_sum / total if total else 0.0
        test_logit_scale = logit_scale_sum / total if total else 0.0
        test_output_spikes_per_sample = (
            output_spike_sum / output_sample_count if output_sample_count else None
        )
        test_output_silent_fraction = (
            output_silent_samples / output_sample_count if output_sample_count else None
        )
        class_spike_total = float(output_class_spikes.sum().item())
        test_output_class_spike_fraction = (
            [float(value / class_spike_total) for value in output_class_spikes.tolist()]
            if class_spike_total
            else [0.0] * M.N_OUT
            if output_sample_count
            else None
        )
        test_output_by_input_rate = [
            {
                "rate_hz": rate,
                "n_samples": int(bucket["n_samples"]),
                "accuracy_pct": 100.0 * bucket["n_correct"] / bucket["n_samples"],
                "spikes_per_sample": bucket["n_spikes"] / bucket["n_samples"],
                "silent_fraction": bucket["n_silent"] / bucket["n_samples"],
            }
            for rate, bucket in sorted(output_by_rate.items())
        ]

        new_best = avg_test < best_validation_loss or (
            avg_test == best_validation_loss and acc > best_acc
        )
        if new_best:
            best_validation_loss = avg_test
            best_acc = acc
            best_epoch = epoch + 1
            best_state = {k: v.cpu().clone() for k, v in net.state_dict().items()}
            no_improve = 0
        else:
            no_improve += 1

        cur_lr = opt.param_groups[0]["lr"]
        elapsed = _time.perf_counter() - t_epoch

        if cur_lr != prev_lr:
            log.info(f"  ⚡ lr → {cur_lr:.0e}")
            prev_lr = cur_lr

        # Compute firing-rate metrics for the JSON.
        t_observe = _time.perf_counter()
        rng_state = torch.get_rng_state()
        net.recording = True
        with torch.no_grad():
            net(input_spikes=ref_spikes)
        net.recording = False
        torch.set_rng_state(rng_state)

        _rec = net.spike_record
        spk_e = _to_np(_rec[primary_hid_key(_rec)])
        _ik = primary_inh_key(_rec)
        spk_i = _to_np(_rec[_ik]) if _ik else None
        epoch_metrics = compute_metrics(
            spk_e, spk_i, dt, model_name, n_e=M.N_HID, n_i=M.N_INH
        )
        observe_s = _time.perf_counter() - t_observe

        # Trainable-parameter Frobenius norms — surfaced per epoch so
        # convergence audits (nb024) can tell whether the optimiser is
        # still actively moving each named parameter. Mirrors the
        # structure of grad_ratios (one key per named parameter, scalar
        # value). All grads-requiring parameters included.
        weight_norms = {
            name: float(p.detach().norm().item())
            for name, p in net.named_parameters()
            if p.requires_grad
        }

        # Record this epoch into the structured metrics history
        record = {
            "ep": epoch + 1,
            "acc": acc,
            "loss": avg_train,
            "test_loss": avg_test,
            "test_rate_e": test_rate_e,
            "test_rate_i": test_rate_i,
            "test_margin": test_margin,
            "test_confidence": test_confidence,
            "test_logit_scale": test_logit_scale,
            "test_output_spikes_per_sample": test_output_spikes_per_sample,
            "test_output_silent_fraction": test_output_silent_fraction,
            "test_output_class_spike_fraction": test_output_class_spike_fraction,
            "test_output_by_input_rate": test_output_by_input_rate,
            "validation_draws": validation_draws,
            "lr": cur_lr,
            "elapsed_s": elapsed,
            "train_compute_s": train_compute_s,
            "eval_s": eval_s,
            "observe_s": observe_s,
            "samples": n_samples_train,
            "grad_norm": avg_grad,
            "grad_norm_max": grad_max,
            "grad_ratios": grad_ratios,
            "grad_norms": grad_norms,
            "weight_norms": weight_norms,
            "skipped_steps": n_skipped_steps,
            "nan_forward_batches": n_nan_forward,
            "new_best": new_best,
        }
        # Flatten per-parameter gradient norms to scalar fields so they land in
        # the per-epoch jsonl sidecar (e.g. gnorm__W_ei.1, gnorm__W_ie.1).
        for _pname, _gn in grad_norms.items():
            record[f"gnorm__{_pname}"] = _gn
        if epoch_metrics:
            record.update(epoch_metrics)
        epoch_records.append(record)

        # jsonl sidecar — one line per epoch (skip dict fields, they
        # round-trip via metrics.json).
        jsonl.write(
            **{
                k: v
                for k, v in record.items()
                if k
                not in ("grad_ratios", "grad_norms", "weight_norms", "validation_draws")
            }
            )

        # Structured progress line + warning tracker
        e_rate = epoch_metrics.get("rate_e", 0.0) if epoch_metrics else 0.0
        i_rate = epoch_metrics.get("rate_i") if epoch_metrics else None
        cv = epoch_metrics.get("cv", 0.0) if epoch_metrics else 0.0
        activity = epoch_metrics.get("act", 0.0) * 100 if epoch_metrics else 0.0
        flags = wtracker.tick(epoch + 1, acc, activity, avg_train)
        # ETA from a running mean of epoch wall-times, not just the last epoch.
        # Epoch 1 is a compile-inflated outlier, so drop it from the mean once
        # we have a post-compile epoch to average — this stops the ETA lurching.
        epoch_times.append(elapsed)
        eta_sample = epoch_times[1:] if len(epoch_times) > 1 else epoch_times
        mean_epoch_s = sum(eta_sample) / len(eta_sample)
        eta = (epochs - epoch - 1) * mean_epoch_s
        runlog.print_epoch(
            log,
            epoch + 1,
            epochs,
            acc,
            avg_train,
            e_rate,
            i_rate,
            cv,
            activity,
            elapsed,
            eta,
            new_best=new_best,
            warnings=flags,
        )

    total_time = _time.perf_counter() - t_start

    # Tier 1 perf block. Per-epoch breakdown is in epoch_records
    # (train_compute_s / eval_s / observe_s); these are the whole-run
    # aggregates the dashboards read.
    perf_block: dict = {
        "device": {"type": device.type},
        "torch_version": torch.__version__,
    }
    try:
        if device.type == "cuda":
            perf_block["device"]["name"] = torch.cuda.get_device_name(device)
            perf_block["peak_memory_bytes"] = int(
                torch.cuda.max_memory_allocated(device)
            )
        elif device.type == "mps" and peak_mem_mps > 0:
            # current_allocated_memory sampled per-epoch (max). Reflects active
            # tensor allocation, not the MPS driver's cached pool.
            perf_block["peak_memory_bytes"] = int(peak_mem_mps)
    except Exception:
        pass
    if epoch_records:
        perf_block["epoch1_total_s"] = float(epoch_records[0]["elapsed_s"])
        warm = epoch_records[1:] if len(epoch_records) > 1 else epoch_records
        warm_total_s = sum(r["elapsed_s"] for r in warm)
        warm_compute_s = sum(r.get("train_compute_s", 0.0) for r in warm)
        warm_samples = sum(r.get("samples", 0) for r in warm)
        perf_block["epoch_warm_mean_s"] = float(warm_total_s / max(len(warm), 1))
        if warm_compute_s > 0:
            perf_block["samples_per_sec_warm"] = float(warm_samples / warm_compute_s)

    # Snapshot end state
    end_state = None
    if snapshot_end:
        C.cfg.n_e = M.N_HID
        C.cfg.n_i = M.N_INH
        net.recording = True
        with torch.no_grad():
            net(input_spikes=ref_spikes)
        net.recording = False
        rec = net.spike_record

        spk_e = _to_np(rec[primary_hid_key(rec)])
        _ik = primary_inh_key(rec)
        spk_i = _to_np(rec[_ik]) if _ik else None
        end_state = compute_metrics(
            spk_e, spk_i, dt, model_name, n_e=M.N_HID, n_i=M.N_INH
        )

    # Save weights — best-accuracy state (deployment) AND the final-epoch state.
    # best_epoch can be far earlier than the last epoch: accuracy plateaus while
    # the firing-rate attractor keeps drifting (nb024), so the end-of-training
    # network is a genuinely different, scientifically interesting object. Saving
    # only best_state would make that drifted network unrecoverable. net still
    # holds the final weights here (the end-state snapshot below is no_grad, and
    # best_state is not reloaded until test_predictions further down).
    if best_state is not None:
        torch.save(best_state, out_dir / "weights.pth")
    torch.save(net.state_dict(), out_dir / "weights_final.pth")
    checkpoints = {
        "final_epoch": {
            "filename": "weights_final.pth",
            "epoch": int(epochs),
            "sha256": _sha256_file(out_dir / "weights_final.pth"),
        },
    }
    if best_state is not None:
        checkpoints["best_validation"] = {
            "filename": "weights.pth",
            "epoch": int(best_epoch),
            "sha256": _sha256_file(out_dir / "weights.pth"),
            "selection_metric": "validation_cross_entropy_mean_over_encoder_draws",
            "tie_breaker": "validation_accuracy_pct_mean_over_encoder_draws",
            "validation_draw_count": validation_draw_count,
        }

    # Write structured metrics for tests/analysis (parallels output.log).
    # run_finished_at is the canonical "when did this run actually produce
    # its numbers" timestamp — consumers should prefer it over file mtime,
    # which gets clobbered by git checkout, file copies, etc.
    from datetime import datetime, timezone

    metrics_path = out_dir / "metrics.json"
    metrics_blob = {
        "mode": "train",
        "schema_version": 1,
        "model": model_name,
        "run_finished_at": datetime.now(timezone.utc).isoformat(timespec="seconds"),
        "config": {
            "dt": dt,
            "t_ms": t_ms,
            **timing_config,
            "epochs": epochs,
            "lr": lr,
            "weight_decay": weight_decay,
            "input_rate": M.max_rate_hz,
            "input_rates": list(input_rates) if input_rates else None,
            "input_rate_sampling": "uniform_categorical_per_presentation"
            if input_rates
            else "fixed",
            "w_in": list(w_in) if w_in else None,
            "w_ee": list(w_ee) if w_ee else None,
            "w_in_initial_zero_fraction": w_in_initial_zero_fraction,
            "recurrent_initial_zero_fraction": recurrent_initial_zero_fraction,
            "ei_strength": ei_strength,
            "ei_ratio": ei_ratio,
            # Full training-dynamics config carried here too, so metrics.json is
            # self-sufficient for the frontier/ladder analyses without opening the
            # co-located config.json — every knob that shapes the loss landscape.
            "tau_gaba_ms": float(M.tau_gaba),
            "fr_reg_upper_target_hz": fr_reg_upper_target_hz,
            "fr_reg_upper_strength": fr_reg_upper_strength,
            "fr_reg_contract": config["fr_reg_contract"],
            "readout_w_out_scale": (
                readout_w_out_scale if readout_w_init is None else None
            ),
            "readout_w_init": config["readout_w_init"],
            "readout_w_init_mean": readout_w_init_mean,
            "readout_w_init_std": readout_w_init_std,
            "readout_w_init_realized": readout_init_stats,
            "weight_initialization": net.weight_initialization,
            "readout_mode": readout_mode,
            "signed_readout": signed_readout,
            "readout_bias": readout_bias,
            "readout_reduction": (
                M.CUMULATIVE_READOUT_REDUCTION
                if readout_mode == "cumulative-potential"
                else "output_spike_count"
                if readout_mode == "spike-count"
                else "output_spike_count_divided_by_duration_seconds"
                if readout_mode == "spike-rate"
                else None
            ),
            "readout_units": (
                "spikes"
                if readout_mode == "spike-count"
                else "Hz"
                if readout_mode == "spike-rate"
                else None
            ),
            "readout_tau_bounds_ms": (
                list(M.CUMULATIVE_READOUT_TAU_BOUNDS_MS)
                if readout_mode == "cumulative-potential"
                else None
            ),
            "readout_reference": (
                M.CUMULATIVE_READOUT_REFERENCE
                if readout_mode == "cumulative-potential"
                else None
            ),
            "dales_law": dales_law,
            "trainable_w_ee": trainable_w_ee,
            "trainable_w_ei": trainable_w_ei,
            "trainable_w_ie": trainable_w_ie,
            "trainable_w_ii": trainable_w_ii,
            "state_clamp": state_clamp,
            "seed": seed,
            "n_hidden": M.N_HID,
            "n_inh": M.N_INH,
            "n_in": M.N_IN,
            "n_out": M.N_OUT,
            "tau_ampa_ms": float(M.tau_ampa),
            "max_samples": max_samples,
            "dataset": dataset,
            "dataset_split": config["dataset_split"],
            "validation_encoder_draws": config["validation_encoder_draws"],
            "v_grad_dampen": v_grad_dampen,
            "batch_size": bs,
            "grad_clip": GRAD_CLIP,
            "n_params": n_params,
            "n_trainable": n_trainable,
        },
        "init": snapshot_init_state,
        "epochs": epoch_records,
        "end": end_state,
        "weight_final": net.weight_final_statistics(),
        "best_acc": best_acc,
        "best_validation_loss": best_validation_loss,
        "best_epoch": best_epoch,
        "checkpoints": checkpoints,
        "total_elapsed_s": total_time,
        "perf": perf_block,
    }
    with open(metrics_path, "w") as f:
        json.dump(metrics_blob, f, indent=2, default=float)

    jsonl.close()

    # Historical filename retained for artifact compatibility; these are
    # validation predictions used during checkpoint selection, not test data.
    if best_state is not None:
        net.load_state_dict(best_state, strict=False)
    net.eval()
    preds = []
    idx = 0
    with torch.no_grad():
        for X_b, y_b in test_loader:
            X_b, y_b = X_b.to(device), y_b.to(device)
            spk = encode_batch(
                X_b,
                dt,
                max_rate_hz=sample_input_rates(len(X_b), prediction_rate_gen),
            )
            logits_t = net(input_spikes=spk)
            p = logits_t.argmax(1)
            for i in range(y_b.size(0)):
                preds.append(
                    {
                        "idx": idx,
                        "true": int(y_b[i].item()),
                        "pred": int(p[i].item()),
                        "correct": bool(p[i].item() == y_b[i].item()),
                        "logits": [float(x) for x in logits_t[i].tolist()],
                    }
                )
                idx += 1
    runlog.write_test_predictions(out_dir / "test_predictions.json", preds)

    # Structured summary block
    dyn = None
    if end_state:
        dyn = {
            "E": f"{end_state.get('rate_e', 0):.0f}Hz",
            "I": (
                f"{end_state.get('rate_i', 0):.0f}Hz"
                if end_state.get("rate_i") not in (None, 0.0)
                else "—"
            ),
            "CV": f"{end_state.get('cv', 0):.2f}",
            "act": f"{end_state.get('act', 0) * 100:.0f}%",
        }
    runlog.summary(
        log,
        best_acc=best_acc,
        final_acc=acc,
        best_epoch=best_epoch,
        total_epochs=epochs,
        runtime_s=total_time,
        perf=perf_block,
        device=device.type,
        dynamics=dyn,
        out_dir=out_dir,
        warnings=wtracker.summary_lines(),
    )
    return best_acc

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
BATCH_SIZEunannotated64Source
GRAD_CLIPunannotated1.0Source
VALIDATION_DRAW_COUNTunannotated3Source
VALIDATION_ENCODER_SEEDSunannotatedtuple((EVAL_SEED + draw_index for draw_index in range(VALIDATION_DRAW_COUNT)))Source
VALIDATION_RATE_SEEDSunannotatedtuple((EVAL_SEED + 10000 + draw_index for draw_index in range(VALIDATION_DRAW_COUNT)))Source
IMAGE_DATASETSunannotated{'mnist'}Source

On this page