Download src/scripts/visualize_architecture.py from Quazim0t0/SpikeWhale-SNN-216M: direct link, hf CLI and curl.
- Browser
- Download file 11.7 kB
-
https://huggingface.co/Quazim0t0/SpikeWhale-SNN-216M/resolve/main/src/scripts/visualize_architecture.py
- Command line
-
hf download hf://Quazim0t0/SpikeWhale-SNN-216M/src/scripts/visualize_architecture.py
-
curl -L -o visualize_architecture.py https://huggingface.co/Quazim0t0/SpikeWhale-SNN-216M/resolve/main/src/scripts/visualize_architecture.py
11.7 kB
| """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() | |