snnlab.sim.train
Complete declared API of the train module, with signatures, data fields, validation and source.
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.
| Symbol | Kind |
|---|---|
| seed_everything | function |
| ShdBinnedDataset | class |
| train | function |
seed_everything
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
seed | unannotated | required | Seed controlling this operation’s random stream. |
Return expressions (branch-dependent; names refer to the linked implementation):
NoneImplementation
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
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)| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
events | unannotated | required | Defined by the constructor implementation below. |
labels | unannotated | required | Defined by the constructor implementation below. |
n_in | unannotated | required | Defined by the constructor implementation below. |
Instance members assigned by the constructor (expressions are evaluated when constructed):
| Member | Annotation | Initial expression |
|---|---|---|
events | unannotated | events |
labels | unannotated | labels |
n_in | unannotated | n_in |
ShdBinnedDataset.len
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
def ShdBinnedDataset.__getitem__(self, index) -> tuple[torch.Tensor, int]| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
index | unannotated | required | Defined 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
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.| Parameter | Annotation | Default | Meaning |
|---|---|---|---|
model_name | unannotated | 'ping' | Defined by the source contract and implementation below. |
lr | unannotated | 0.01 | Learning rate; frozen groups require zero and trainable groups require a positive value. |
weight_decay | unannotated | 0.0 | Defined by the source contract and implementation below. |
epochs | unannotated | 100 | Number of training passes over the selected presentations. |
dt | unannotated | 0.1 | Timestep; authoring uses a Quantity and legacy simulation uses milliseconds. |
out_dir | unannotated | None | Defined by the source contract and implementation below. |
device_name | unannotated | None | Defined by the source contract and implementation below. |
w_in | unannotated | None | Defined by the source contract and implementation below. |
w_ee | unannotated | None | Defined by the source contract and implementation below. |
w_ei | unannotated | None | Defined by the source contract and implementation below. |
w_ie | unannotated | None | Defined by the source contract and implementation below. |
w_ii | unannotated | None | Defined by the source contract and implementation below. |
ei_strength | unannotated | None | Defined by the source contract and implementation below. |
ei_ratio | unannotated | 2.0 | Defined by the source contract and implementation below. |
w_in_initial_zero_fraction | unannotated | 0.0 | Defined by the source contract and implementation below. |
recurrent_initial_zero_fraction | unannotated | 0.0 | Defined by the source contract and implementation below. |
dataset | unannotated | 'mnist' | Defined by the source contract and implementation below. |
snapshot_init | unannotated | True | Defined by the source contract and implementation below. |
snapshot_end | unannotated | True | Defined by the source contract and implementation below. |
t_ms | unannotated | 200.0 | Physical duration in milliseconds. |
hidden_sizes | unannotated | None | Defined by the source contract and implementation below. |
max_samples | unannotated | None | Defined by the source contract and implementation below. |
v_grad_dampen | unannotated | 80.0 | Defined by the source contract and implementation below. |
dales_law | unannotated | True | Defined by the source contract and implementation below. |
batch_size | unannotated | None | Number of presentations in a binding or mini-batch. |
seed | unannotated | None | Seed controlling this operation’s random stream. |
readout_w_out_scale | unannotated | 1.0 | Defined by the source contract and implementation below. |
readout_w_init_mean | unannotated | None | Defined by the source contract and implementation below. |
readout_w_init_std | unannotated | None | Defined by the source contract and implementation below. |
readout_mode | unannotated | 'rate' | Defined by the source contract and implementation below. |
signed_readout | unannotated | False | Defined by the source contract and implementation below. |
readout_bias | unannotated | False | Defined by the source contract and implementation below. |
input_rates | unannotated | None | Defined by the source contract and implementation below. |
tau_gaba | unannotated | None | Defined by the source contract and implementation below. |
fr_reg_upper_target_hz | unannotated | 0.0 | Defined by the source contract and implementation below. |
fr_reg_upper_strength | unannotated | 0.0 | Defined by the source contract and implementation below. |
trainable_w_ee | unannotated | False | Defined by the source contract and implementation below. |
trainable_w_ei | unannotated | False | Defined by the source contract and implementation below. |
trainable_w_ie | unannotated | False | Defined by the source contract and implementation below. |
trainable_w_ii | unannotated | False | Defined by the source contract and implementation below. |
state_clamp | unannotated | False | Defined by the source contract and implementation below. |
train_leak | unannotated | False | Defined by the source contract and implementation below. |
tau_m_e_bounds_ms | unannotated | None | Defined by the source contract and implementation below. |
tau_m_i_bounds_ms | unannotated | None | Defined by the source contract and implementation below. |
adaptive_threshold | unannotated | False | Defined by the source contract and implementation below. |
adapt_tau_bounds_ms | unannotated | None | Defined by the source contract and implementation below. |
adapt_strength_init_mv | unannotated | 1.0 | Defined by the source contract and implementation below. |
adapt_strength_max_mv | unannotated | None | Defined by the source contract and implementation below. |
refractory_e_ms | unannotated | None | Defined by the source contract and implementation below. |
refractory_i_ms | unannotated | None | Defined by the source contract and implementation below. |
refractory_policy | unannotated | 'nearest' | Defined by the source contract and implementation below. |
Return expressions (branch-dependent; names refer to the linked implementation):
0.0best_accExplicit 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_accConstants and type aliases
Initial source expressions are shown, not evaluated runtime values. Legacy configuration may mutate module defaults.
| Name | Annotation | Initial expression | Source |
|---|---|---|---|
log | unannotated | logging.getLogger('cli') | Source |
BATCH_SIZE | unannotated | 64 | Source |
GRAD_CLIP | unannotated | 1.0 | Source |
VALIDATION_DRAW_COUNT | unannotated | 3 | Source |
VALIDATION_ENCODER_SEEDS | unannotated | tuple((EVAL_SEED + draw_index for draw_index in range(VALIDATION_DRAW_COUNT))) | Source |
VALIDATION_RATE_SEEDS | unannotated | tuple((EVAL_SEED + 10000 + draw_index for draw_index in range(VALIDATION_DRAW_COUNT))) | Source |
IMAGE_DATASETS | unannotated | {'mnist'} | Source |