"""Static architecture portrait of the spiking LM -- straight from a checkpoint, no training information. This is the sibling of `visualize_checkpoint.py` (the live "training scope"), with every training-history element removed: no metrics CSV, no learning curve, no influence-over-training curves. It draws ONLY what is intrinsic to the saved model -- the layers, the neurons, their gene-program identity, and their learned wiring -- and writes a single PNG suitable for posting to the model repo. What it shows ------------- * one stacked plane per LIF layer (embed at the bottom, readout on top); every hidden neuron sits at a fixed (x,y) so it is a vertical column through the depth * neuron COLOUR = its dominant cell-program identity from modulation.P * neuron SIZE = its recurrent wiring strength at that layer (||W_rec row||) * a sample of the strongest inter-layer weights as faint edges * side panels (all derived from the final weights, NOT from training history): - per-program modulation influence = ||(thr,beta,in)_gain|| x std(P_column) - neuron census per gene program - the learned threshold/beta/gain strength for each program Only torch + numpy + matplotlib are needed (no `ml` package, no GPU, no data). Usage ----- PYTHONPATH=src python src/scripts/visualize_architecture.py \ --checkpoint snn_stream_program.pth --out spikewhale_architecture.png """ from __future__ import annotations import argparse import math import numpy as np import torch import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import matplotlib.colors as mcol from matplotlib.lines import Line2D from mpl_toolkits.mplot3d import Axes3D # noqa: F401 (registers 3d projection) # stable colour per program NAME (shared with the other visualisers) PALETTE = { "cell_division": "#E24B4A", "neuron_identity": "#639922", "pluripotency": "#378ADD", "endocrine": "#D4537E", "metabolism": "#BA7517", "signaling_identity": "#7F77DD", } _FALLBACK = ["#E24B4A", "#639922", "#378ADD", "#D4537E", "#BA7517", "#7F77DD", "#888780", "#1D9E75"] _DEFAULT_PROGRAMS = ["cell_division", "neuron_identity", "pluripotency", "endocrine", "metabolism", "signaling_identity"] FG = "#e6edf3"; MUTED = "#8b949e"; BG = "#0d1117"; PANEL = "#161b22"; GRID = "#30363d" def _base(name): return str(name).split("#", 1)[0] def base_programs(names): seen = [] for nm in names or []: b = _base(nm) if b not in seen: seen.append(b) return seen or list(_DEFAULT_PROGRAMS) def prog_color(base, i): return PALETTE.get(base, _FALLBACK[i % len(_FALLBACK)]) def grid_positions(n): cols = int(math.ceil(math.sqrt(n))) rows = int(math.ceil(n / cols)) idx = np.arange(n) x = (idx % cols) / max(cols - 1, 1) y = (idx // cols) / max(rows - 1, 1) return x, y def program_scores(P, program_names): """[H,K] identity matrix + names -> ([H,n_base] per-program scores, base list).""" bases = base_programs(program_names) H, n = P.shape[0], len(bases) bi = {b: i for i, b in enumerate(bases)} scores = np.zeros((H, n), np.float32) counts = np.zeros(n, np.float32) for j, nm in enumerate(program_names or []): b = _base(nm) if b in bi: scores[:, bi[b]] += P[:, j]; counts[bi[b]] += 1 counts[counts == 0] = 1 scores /= counts return scores, bases def extract(ckpt, stride, max_edges): sd = ckpt["model_state_dict"] cfg = ckpt["config"] L = cfg["num_layers"]; H = cfg["hidden"] mod = ckpt.get("modulation", {}) P = mod.get("P") P = P.numpy() if torch.is_tensor(P) else np.zeros((H, len(_DEFAULT_PROGRAMS)), np.float32) names = mod.get("program_names") or list(_DEFAULT_PROGRAMS) scores, bases = program_scores(P, names) dom = scores.argmax(1) rgb = np.array([mcol.to_rgb(prog_color(bases[d], d)) for d in dom], np.float32) enabled = bool(mod.get("enabled", False)) if not enabled or np.allclose(P, 0): rgb[:] = 0.6 # modulation off -> neutral grey counts = [int((dom == i).sum()) for i in range(len(bases))] sel = np.arange(0, H, max(stride, 1)) gx, gy = grid_positions(H) gx, gy = gx[sel], gy[sel] rgb_sel = rgb[sel] layers = [] for l in range(L): w = sd[f"W_rec.{l}.weight"].float() rownorm = w.norm(dim=1).numpy()[sel] rownorm = rownorm / (rownorm.mean() + 1e-6) layers.append({"z": l + 1, "size": 8.0 * rownorm + 2.0}) edges = [] for l in range(1, L): w = sd[f"W_in.{l}.weight"].float().abs() k = min(max_edges, w.numel()) flat = torch.topk(w.flatten(), k).indices.numpy() out_n = flat // w.shape[1]; in_n = flat % w.shape[1] for o, i in zip(out_n, in_n): oi, ii = o // max(stride, 1), i // max(stride, 1) if oi < len(gx) and ii < len(gx): edges.append(((gx[ii], gy[ii], l), (gx[oi], gy[oi], l + 1))) # per-program influence + learned gains (from the FINAL weights only) tg = sd.get("modulation.thr_gain"); bg = sd.get("modulation.beta_gain") ig = sd.get("modulation.in_gain") infl = {b: 0.0 for b in bases} gains = {b: {"thr": 0.0, "beta": 0.0, "in": 0.0, "_n": 0} for b in bases} if tg is not None and P.shape[1] == len(tg): tg, bg, ig = tg.numpy(), bg.numpy(), ig.numpy() for j, nm in enumerate(names): b = _base(nm) if b in infl: infl[b] += float(math.sqrt(tg[j]**2 + bg[j]**2 + ig[j]**2) * P[:, j].std()) gains[b]["thr"] += float(tg[j]); gains[b]["beta"] += float(bg[j]) gains[b]["in"] += float(ig[j]); gains[b]["_n"] += 1 for b in gains: n = max(gains[b].pop("_n"), 1) for k in gains[b]: gains[b][k] /= n tens = [v for v in sd.values() if torch.is_tensor(v)] nparam = sum(v.numel() for v in tens) rms = float(torch.sqrt(sum((v.float() ** 2).sum() for v in tens) / nparam)) fam = cfg.get("family", {}) trait_map = [("use_engram", "Engram"), ("use_hrm", "HRM"), ("use_moe", "MoE"), ("use_mtp", "MTP"), ("use_spikeattn", "SpikeAttn"), ("use_progsem", "ProgSem"), ("use_kuramoto", "Kuramoto"), ("use_jepa", "JEPA")] traits = [lbl for key, lbl in trait_map if fam.get(key)] n_seeds = P.shape[1] // max(len(bases), 1) return { "L": L, "H": H, "gx": gx, "gy": gy, "rgb": rgb_sel, "layers": layers, "edges": edges, "step": int(ckpt.get("step", 0)), "nparam": nparam, "rms": rms, "traits": traits or ["none"], "P_cols": P.shape[1], "n_seeds": n_seeds, "enabled": enabled, "infl": infl, "bases": bases, "counts": counts, "gains": gains, } def draw_3d(ax, d): ax.set_facecolor(BG) for z, lbl in [(0, "embed"), (d["L"] + 1, "readout")]: ax.scatter(d["gx"], d["gy"], np.full_like(d["gx"], z), s=3, c=[(0.5, 0.5, 0.5)], alpha=0.15) ax.text(0.0, 1.05, z, lbl, color=MUTED, fontsize=8) for (a, b) in d["edges"]: ax.plot([a[0], b[0]], [a[1], b[1]], [a[2], b[2]], color="0.6", linewidth=0.3, alpha=0.22) for lay in d["layers"]: z = np.full_like(d["gx"], lay["z"]) ax.scatter(d["gx"], d["gy"], z, s=lay["size"], c=d["rgb"], depthshade=True, edgecolors="none") ax.text(0.0, 1.05, lay["z"], f"L{lay['z']}", color=MUTED, fontsize=8) ax.set_zlim(-0.5, d["L"] + 1.5) ax.set_xticks([]); ax.set_yticks([]) ax.set_zlabel("layer (depth)", color=FG) ax.tick_params(colors=MUTED) ax.view_init(elev=18, azim=-60) for pane in (ax.xaxis, ax.yaxis, ax.zaxis): pane.set_pane_color((0.05, 0.06, 0.09, 1.0)) def _style(ax): ax.set_facecolor(PANEL) for s in ax.spines.values(): s.set_color(GRID) ax.tick_params(colors=MUTED) ax.title.set_color(FG) ax.xaxis.label.set_color(FG); ax.yaxis.label.set_color(FG) def draw_influence(ax, d): _style(ax) items = sorted(d["infl"].items(), key=lambda kv: kv[1]) bases_i = {b: i for i, b in enumerate(d["bases"])} ys = np.arange(len(items)) ax.barh(ys, [v for _, v in items], color=[prog_color(b, bases_i[b]) for b, _ in items], edgecolor="none") ax.set_yticks(ys); ax.set_yticklabels([b.replace("_", " ") for b, _ in items], fontsize=9) ax.set_xlabel("modulation influence ||gain|| x std(P)") ax.set_title("how strongly each gene program shapes the neurons", fontsize=9) ax.grid(axis="x", alpha=0.15) def draw_census(ax, d): _style(ax) bases = d["bases"]; counts = d["counts"] xs = np.arange(len(bases)) ax.bar(xs, counts, color=[prog_color(b, i) for i, b in enumerate(bases)], edgecolor="none") for x, c in zip(xs, counts): ax.text(x, c, f"{c}", ha="center", va="bottom", color=FG, fontsize=8) ax.set_xticks(xs) ax.set_xticklabels([b.replace("_", "\n") for b in bases], fontsize=8) ax.set_ylabel("neurons (dominant program)") ax.set_title("neuron census by gene program", fontsize=9) ax.grid(axis="y", alpha=0.15) def main(): p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) p.add_argument("--checkpoint", default="snn_models/snn_stream_program.pth") p.add_argument("--out", default="spikewhale_architecture.png") p.add_argument("--max-edges", type=int, default=40, help="sampled edges per layer gap") p.add_argument("--stride", type=int, default=1, help="plot every Nth neuron (speed)") p.add_argument("--dpi", type=int, default=150) p.add_argument("--title", default="SpikeWhale-SNN-216M") args = p.parse_args() ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False) d = extract(ckpt, args.stride, args.max_edges) fig = plt.figure(figsize=(16, 9), dpi=args.dpi) fig.patch.set_facecolor(BG) gs = fig.add_gridspec(2, 2, width_ratios=[1.5, 1.0], hspace=0.32, wspace=0.22) ax3d = fig.add_subplot(gs[:, 0], projection="3d") ax_inf = fig.add_subplot(gs[0, 1]) ax_cen = fig.add_subplot(gs[1, 1]) draw_3d(ax3d, d) draw_influence(ax_inf, d) draw_census(ax_cen, d) legend = [Line2D([0], [0], marker="o", ls="", ms=9, mec="none", mfc=prog_color(b, i), label=b.replace("_", " ")) for i, b in enumerate(d["bases"])] leg = fig.legend(handles=legend, loc="lower center", ncol=len(d["bases"]), frameon=True, facecolor=PANEL, edgecolor=GRID, labelcolor=FG, fontsize=9, bbox_to_anchor=(0.5, 0.005), title="neuron colour = dominant gene program") leg.get_title().set_color(FG) fig.suptitle( f"{args.title} ยท a from-scratch spiking neural network, neurons modulated by a virtual-cell gene-program model", color=FG, fontsize=14, y=0.975) fig.text(0.03, 0.90, f"{d['H']:,} neurons x {d['L']} LIF layers | {d['nparam']/1e6:.1f}M params | weight RMS {d['rms']:.3f}\n" f"traits: {', '.join(d['traits'])}\n" f"modulation {'ON' if d['enabled'] else 'off'} " f"({len(d['bases'])} gene programs x {d['n_seeds']} seeds = {d['P_cols']} identity cols)", color=MUTED, fontsize=10, va="top", ha="left", linespacing=1.5) fig.subplots_adjust(top=0.90, bottom=0.10, left=0.02, right=0.97) fig.savefig(args.out, facecolor=fig.get_facecolor()) print("wrote", args.out) print("programs:", d["bases"]) print("neurons/program:", dict(zip(d["bases"], d["counts"]))) if __name__ == "__main__": main()