SpikeWhale-SNN-216M / src /scripts /visualize_architecture.py
Quazim0t0's picture
Add static architecture portrait (PNG + visualize_architecture.py) and embed in card
2c04a66 verified
Raw History Blame Contribute Delete
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()