ProCreations's picture
download
raw
61.8 kB
#!/usr/bin/env python3
"""llmviz - data-driven 3D visualisation of MiniCPM5-2B generating text.
Every surface is a real tensor of the model (texel = weight, mip-mapped), every
sheet is a real activation captured while the model generated text, every arc
is a real attention weight. World scale: 1 unit = 64 weights along an axis.
"""
import os, sys, json, math, time, subprocess, argparse
import numpy as np
import torch
import torch.nn.functional as F
# ------------------------------------------------------------------ constants
UPW = 1.0 / 64.0
NL, D, I, KVD, V, NH, NKV, HD = 42, 2048, 6144, 256, 130560, 16, 2, 128
HID_U, INT_U, KV_U = D * UPW, I * UPW, KVD * UPW # 32, 96, 4
RING_R = V * UPW / (2 * math.pi) # 325.1
RING_SEGS = 96
LAYER_DY = 40.0
TOK_DX = 0.5
Y_TOP = NL * LAYER_DY # 672
RING_LO_Y = -64.0
RING_HI_Y = Y_TOP + 32.0
SHEET_X0 = -40.0
INTER_Y = 0.3
SLAB_H = 0.8
FPS = 60
DUR = 120.0
DEPTH_SOFT = 22.0 # translucency: how fast surfaces behind the nearest one fade (world units)
# palettes: sign tints (negative, positive) + magnitude colormap stops at m = 0, 0.6, 1.0, 1.6, 2.6, 4.0
PAL_SIGN = torch.tensor([
[[0.25, 0.35, 1.00], [1.00, 0.55, 0.20]], # 0 attention weights
[[0.20, 0.70, 1.00], [1.00, 0.40, 0.70]], # 1 activations
[[0.25, 0.35, 1.00], [1.00, 0.55, 0.20]], # 2 norm weights
[[0.16, 0.22, 0.42], [0.16, 0.22, 0.42]], # 3 flat sides
[[0.35, 0.25, 1.00], [1.00, 0.45, 0.30]], # 4 MLP weights
[[0.20, 0.50, 1.00], [1.00, 0.60, 0.25]], # 5 embedding / lm_head
])
CMAP_M = torch.tensor([0.0, 0.6, 1.0, 1.6, 2.6, 4.0])
CMAP = torch.tensor([
[[0.01, 0.03, 0.12], [0.04, 0.10, 0.40], [0.15, 0.40, 1.00], [0.45, 0.85, 1.00], [1.00, 0.95, 0.75], [1.00, 1.00, 1.00]], # weights
[[0.01, 0.06, 0.07], [0.04, 0.26, 0.30], [0.15, 0.90, 0.80], [0.65, 1.00, 0.92], [1.00, 1.00, 0.92], [1.00, 1.00, 1.00]], # activations
[[0.05, 0.04, 0.02], [0.30, 0.22, 0.10], [0.85, 0.65, 0.30], [1.00, 0.85, 0.55], [1.00, 0.95, 0.80], [1.00, 1.00, 1.00]], # norm weights
[[0.10, 0.14, 0.30], [0.10, 0.14, 0.30], [0.10, 0.14, 0.30], [0.10, 0.14, 0.30], [0.10, 0.14, 0.30], [0.10, 0.14, 0.30]], # sides
[[0.02, 0.01, 0.10], [0.08, 0.04, 0.32], [0.30, 0.14, 0.85], [0.70, 0.50, 1.00], [1.00, 0.92, 0.98], [1.00, 1.00, 1.00]], # MLP weights
[[0.01, 0.04, 0.07], [0.03, 0.16, 0.26], [0.10, 0.60, 0.85], [0.50, 0.92, 1.00], [0.95, 1.00, 1.00], [1.00, 1.00, 1.00]], # embeddings
])
ALPHA_OF = dict(slab=0.82, sheet=0.72, inter=0.72, strip=0.9, ring=0.9, side=0.85, normw=0.9)
def tok_x(i):
return SHEET_X0 + (i + 0.5) * TOK_DX
def head_z(h):
return -HID_U / 2 + (h + 0.5) * (HID_U / NH)
def vocab_angle(v):
return 2 * math.pi * (v + 0.5) / V
def hsv(h, s, v):
h = h % 1.0
i = int(h * 6); f = h * 6 - i
p, q, t = v * (1 - s), v * (1 - s * f), v * (1 - s * (1 - f))
return [(v, t, p), (q, v, p), (p, v, t), (p, q, v), (t, p, v), (v, p, q)][i % 6]
HEAD_COL = [hsv(0.0 + h * 0.022, 0.85, 1.0) if h < 8 else hsv(0.47 + (h - 8) * 0.045, 0.8, 1.0) for h in range(NH)]
def sync(dev):
if dev.type == "cuda":
torch.cuda.synchronize()
elif dev.type == "mps":
torch.mps.synchronize()
# ------------------------------------------------------------------ texture atlas
def level_dims(H, W):
out = [(H, W)]
while H > 1 or W > 1:
H = (H + 1) // 2 if H > 1 else 1
W = (W + 1) // 2 if W > 1 else 1
out.append((H, W))
return out
def pool2(x, rms=False):
H, W = x.shape
kh, kw = (2 if H > 1 else 1), (2 if W > 1 else 1)
y = x[None, None]
if rms:
y = y * y
y = F.avg_pool2d(y, (kh, kw), stride=(kh, kw), ceil_mode=True)
if rms:
y = y.sqrt()
return y[0, 0]
class Atlas:
"""All textures (weights + activations) in one flat fp16 buffer with mip pyramids."""
def __init__(self, dev, min_level=0):
self.dev = dev
self.min_level = min_level
self.specs, self.scales, self.sources = [], [], []
def add(self, H, W, source, scale=None):
self.specs.append((H, W)); self.sources.append(source); self.scales.append(scale)
return len(self.specs) - 1
def _levels(self, H, W):
dims = level_dims(H, W)
lo = min(self.min_level, len(dims) - 1)
return lo, dims[lo:]
def build(self):
total = 0
tables = []
for (H, W) in self.specs:
lo, dims = self._levels(H, W)
offs = []
for (h, w) in dims:
offs.append((total, h, w)); total += h * w
tables.append((lo, offs))
print(f"[atlas] {len(self.specs)} textures, {total/1e9:.2f} G texels -> {total*4/1e9:.1f} GB fp16 (val+mag)", flush=True)
self.val = torch.empty(total, dtype=torch.float16, device=self.dev)
self.mag = torch.empty(total, dtype=torch.float16, device=self.dev)
n = len(tables)
LMAX = max(len(t[1]) for t in tables)
lev_off = torch.zeros(n, LMAX, dtype=torch.int64); lev_H = torch.ones(n, LMAX, dtype=torch.int64); lev_W = torch.ones(n, LMAX, dtype=torch.int64)
n_lev = torch.zeros(n, dtype=torch.int64); scale = torch.zeros(n); base_H = torch.zeros(n); base_W = torch.zeros(n)
t0 = time.time()
for ti, ((H, W), src, sc) in enumerate(zip(self.specs, self.sources, self.scales)):
lo, offs = tables[ti]
x = src().to(self.dev, torch.float32)
if x.dim() == 1:
x = x[None]
if sc is None:
sc = float(x.pow(2).mean().sqrt().item()) + 1e-12
val, mag = x, x.abs()
level, li = 0, 0
while True:
if level >= lo and li < len(offs):
off, h, w = offs[li]
assert val.shape == (h, w), (val.shape, h, w, ti)
self.val[off:off + h * w] = val.reshape(-1).half()
self.mag[off:off + h * w] = mag.reshape(-1).half()
lev_off[ti, li] = off; lev_H[ti, li] = h; lev_W[ti, li] = w
li += 1
if li >= len(offs):
break
val = pool2(val); mag = pool2(mag, rms=True); level += 1
for j in range(li, LMAX):
lev_off[ti, j] = lev_off[ti, li - 1]; lev_H[ti, j] = lev_H[ti, li - 1]; lev_W[ti, j] = lev_W[ti, li - 1]
n_lev[ti] = li; scale[ti] = sc; base_H[ti], base_W[ti] = offs[0][1], offs[0][2]
if ti % 100 == 0:
print(f"[atlas] {ti}/{n} {time.time()-t0:.0f}s", flush=True)
self.lev_off, self.lev_H, self.lev_W, self.n_lev = lev_off.to(self.dev), lev_H.to(self.dev), lev_W.to(self.dev), n_lev.to(self.dev)
self.scale, self.base_H, self.base_W = scale.to(self.dev), base_H.to(self.dev), base_W.to(self.dev)
self.sources = None
print(f"[atlas] built in {time.time()-t0:.0f}s", flush=True)
def sample(self, tid, u, v, lod):
"""tid [M]; u,v,lod [M] (u along rows/H, v along cols/W). Returns (val, mag) normalised by texture scale."""
nl = self.n_lev[tid]
lod = torch.minimum(lod.clamp(min=0), (nl - 1).float())
l0 = lod.floor().long()
f = lod - l0.float()
l1 = torch.minimum(l0 + 1, nl - 1)
outs = []
for lv in (l0, l1):
off = self.lev_off[tid, lv]; H = self.lev_H[tid, lv]; W = self.lev_W[tid, lv]
r = u * H.float() - 0.5; c = v * W.float() - 0.5
r0 = r.floor(); c0 = c.floor()
fr = (r - r0)[:, None]; fc = (c - c0)[:, None]
r0 = r0.long(); c0 = c0.long()
Hm, Wm = H - 1, W - 1
ra = torch.minimum(r0.clamp(min=0), Hm); rb = torch.minimum((r0 + 1).clamp(min=0), Hm)
ca = torch.minimum(c0.clamp(min=0), Wm); cb = torch.minimum((c0 + 1).clamp(min=0), Wm)
base = off + ra * W
base2 = off + rb * W
i00 = base + ca; i01 = base + cb; i10 = base2 + ca; i11 = base2 + cb
wgt = torch.cat([(1 - fr) * (1 - fc), (1 - fr) * fc, fr * (1 - fc), fr * fc], 1)
vv = torch.stack([self.val[i00], self.val[i01], self.val[i10], self.val[i11]], 1).float()
mm = torch.stack([self.mag[i00], self.mag[i01], self.mag[i10], self.mag[i11]], 1).float()
outs.append(((vv * wgt).sum(1), (mm * wgt).sum(1)))
val = outs[0][0] * (1 - f) + outs[1][0] * f
mag = outs[0][1] * (1 - f) + outs[1][1] * f
s = self.scale[tid]
return val / s, mag / s
# ------------------------------------------------------------------ scene
class Scene:
def __init__(self, acts, weights, dev, min_level=0):
self.dev = dev
meta = json.loads(str(acts["meta"]))
self.meta = meta
self.P = meta["prompt_len"]
self.T = len(meta["tokens"])
self.tokens = np.array(meta["tokens"])
self.acts = acts
A = self.atlas = Atlas(dev, min_level=min_level)
self.tex = {}
def wsrc(name):
return lambda: weights(name)
self.tex["embed"] = A.add(V, D, wsrc("model.embed_tokens.weight"))
self.tex["lm_head"] = A.add(V, D, wsrc("lm_head.weight"))
self.tex["norm"] = A.add(1, D, wsrc("model.norm.weight"), scale=1.0)
for l in range(NL):
p = f"model.layers.{l}."
self.tex[f"q{l}"] = A.add(D, D, wsrc(p + "self_attn.q_proj.weight"))
self.tex[f"k{l}"] = A.add(KVD, D, wsrc(p + "self_attn.k_proj.weight"))
self.tex[f"v{l}"] = A.add(KVD, D, wsrc(p + "self_attn.v_proj.weight"))
self.tex[f"o{l}"] = A.add(D, D, wsrc(p + "self_attn.o_proj.weight"))
self.tex[f"g{l}"] = A.add(I, D, wsrc(p + "mlp.gate_proj.weight"))
self.tex[f"u{l}"] = A.add(I, D, wsrc(p + "mlp.up_proj.weight"))
self.tex[f"d{l}"] = A.add(D, I, wsrc(p + "mlp.down_proj.weight"))
self.tex[f"ln1w{l}"] = A.add(1, D, wsrc(p + "input_layernorm.weight"), scale=1.0)
self.tex[f"ln2w{l}"] = A.add(1, D, wsrc(p + "post_attention_layernorm.weight"), scale=1.0)
T = self.T
def asrc(key, l=None):
def f():
a = acts[key] if l is None else acts[key][l]
return torch.from_numpy(np.ascontiguousarray(a.astype(np.float32)))
return f
def rms(a):
return float(np.sqrt((a.astype(np.float32) ** 2).mean())) + 1e-9
self.tex["sheet0"] = A.add(T, D, asrc("emb"), scale=rms(acts["emb"]))
for l in range(NL):
self.tex[f"sheet{l+1}"] = A.add(T, D, asrc("resid", l), scale=rms(acts["resid"][l]))
self.tex[f"inter{l}"] = A.add(T, I, asrc("inter", l), scale=rms(acts["inter"][l]))
for key in ("ln1", "q", "k", "v", "attn_out", "ln2", "mlp_out"):
self.tex[f"{key}{l}"] = A.add(T, acts[key].shape[-1], asrc(key, l), scale=rms(acts[key][l]))
self.tex["sheetF"] = A.add(T, D, asrc("final_norm"), scale=rms(acts["final_norm"]))
self.tex["flat"] = A.add(1, 1, lambda: torch.ones(1, 1), scale=1.0)
A.build()
self.build_quads()
def build_quads(self):
Q = []
T = self.T
def add(p0, e1, e2, tex, pal, base, strength, kind, layer=-1, uv=(0, 0, 1, 1), sub=0, alpha=None):
Q.append(dict(p0=p0, e1=e1, e2=e2, tex=tex, pal=pal, base=base, strength=strength, kind=kind, layer=layer, uv=uv, sub=sub, alpha=alpha))
def slab(p0, e1, e2, tex, kind, l):
add(p0, e1, e2, tex, 4 if kind in ("wg", "wu", "wd") else 0, 0.05, 0.95, kind, l, alpha=ALPHA_OF["slab"])
# thickness: four side faces (flat material), hanging below the top face
a = np.array(p0, float); e1 = np.array(e1, float); e2 = np.array(e2, float)
nrm = np.cross(e1, e2); nrm = nrm / np.linalg.norm(nrm)
dn = -nrm * SLAB_H
for (q0, ee) in ((a, e1), (a + e1, e2), (a + e1 + e2, -e1), (a + e2, -e2)):
add(tuple(q0 + dn), tuple(ee), tuple(-dn), self.tex["flat"], 3, 0.6, 0.0, "side", l, alpha=ALPHA_OF["side"])
for l in range(NL):
y = l * LAYER_DY
# attention weights: standing panels in the XY plane (rows = output dim along x, cols = input dim vertical)
slab((-128, y, 0), (32, 0, 0), (0, 32, 0), self.tex[f"q{l}"], "wq", l)
slab((-94, y, 0), (4, 0, 0), (0, 32, 0), self.tex[f"k{l}"], "wk", l)
slab((-88, y, 0), (4, 0, 0), (0, 32, 0), self.tex[f"v{l}"], "wv", l)
slab((-80, y, 0), (32, 0, 0), (0, 32, 0), self.tex[f"o{l}"], "wo", l)
# MLP: gate / down / up standing panels (96 wide x 32 tall), intermediate sheet lying under them
slab((36, y, -20), (96, 0, 0), (0, 32, 0), self.tex[f"g{l}"], "wg", l)
slab((36, y, 20), (96, 0, 0), (0, 32, 0), self.tex[f"u{l}"], "wu", l)
slab((36, y, 0), (0, 32, 0), (96, 0, 0), self.tex[f"d{l}"], "wd", l)
add((-44, y + 0.3, -16), (0.7, 0, 0), (0, 0, 32), self.tex[f"ln1w{l}"], 2, 0.10, 0.8, "ln1w", l, alpha=ALPHA_OF["normw"])
add((28.5, y + 0.3, -16), (0.7, 0, 0), (0, 0, 32), self.tex[f"ln2w{l}"], 2, 0.10, 0.8, "ln2w", l, alpha=ALPHA_OF["normw"])
add((SHEET_X0, y + 0.3, -16), (T * TOK_DX, 0, 0), (0, 0, 32), self.tex[f"sheet{l}"], 1, 0.02, 1.6, "sheet", l, alpha=ALPHA_OF["sheet"])
add((36, y + INTER_Y, -T * TOK_DX / 2), (0, 0, T * TOK_DX), (96, 0, 0), self.tex[f"inter{l}"], 1, 0.02, 1.6, "inter", l, alpha=ALPHA_OF["inter"])
add((-46.2, y + 0.3, -16), (0.7, 0, 0), (0, 0, 32), self.tex[f"ln1{l}"], 1, 0.0, 1.2, "s_ln1", l, alpha=ALPHA_OF["strip"])
add((-128, y + 0.3, 1.2), (0, 0, 0.8), (32, 0, 0), self.tex[f"q{l}"], 1, 0.0, 1.2, "s_q", l, alpha=ALPHA_OF["strip"])
add((-94, y + 0.3, 1.2), (0, 0, 0.8), (4, 0, 0), self.tex[f"k{l}"], 1, 0.0, 1.2, "s_k", l, alpha=ALPHA_OF["strip"])
add((-88, y + 0.3, 1.2), (0, 0, 0.8), (4, 0, 0), self.tex[f"v{l}"], 1, 0.0, 1.2, "s_v", l, alpha=ALPHA_OF["strip"])
add((-80, y + 0.3, 1.2), (0, 0, 0.8), (32, 0, 0), self.tex[f"attn_out{l}"], 1, 0.0, 1.2, "s_o", l, alpha=ALPHA_OF["strip"])
add((30.5, y + 0.3, -16), (0.7, 0, 0), (0, 0, 32), self.tex[f"ln2{l}"], 1, 0.0, 1.2, "s_ln2", l, alpha=ALPHA_OF["strip"])
add((134, y + 0.3, -16), (0.8, 0, 0), (0, 0, 32), self.tex[f"mlp_out{l}"], 1, 0.0, 1.2, "s_mlp", l, alpha=ALPHA_OF["strip"])
y = Y_TOP
add((SHEET_X0, y + 0.3, -16), (T * TOK_DX, 0, 0), (0, 0, 32), self.tex["sheetF"], 1, 0.02, 1.6, "sheet", NL, alpha=ALPHA_OF["sheet"])
add((-44, y, -16), (0.7, 0, 0), (0, 0, 32), self.tex["norm"], 2, 0.10, 0.8, "normw", NL, alpha=ALPHA_OF["normw"])
for ring, (ty, tex) in enumerate(((RING_LO_Y, self.tex["embed"]), (RING_HI_Y, self.tex["lm_head"]))):
for s in range(RING_SEGS):
a0, a1 = 2 * math.pi * s / RING_SEGS, 2 * math.pi * (s + 1) / RING_SEGS
p0 = (RING_R * math.cos(a0), ty, RING_R * math.sin(a0))
p1 = (RING_R * math.cos(a1), ty, RING_R * math.sin(a1))
add(p0, (p1[0] - p0[0], 0, p1[2] - p0[2]), (0, HID_U, 0), tex, 5, 0.05, 0.9,
"ring_lo" if ring == 0 else "ring_hi", -1, uv=(s / RING_SEGS, 0, (s + 1) / RING_SEGS, 1), sub=s, alpha=ALPHA_OF["ring"])
self.Q = Q
dev = self.dev
f32 = lambda key: torch.tensor([q[key] for q in Q], dtype=torch.float32, device=dev)
self.P0, self.E1, self.E2 = f32("p0"), f32("e1"), f32("e2")
self.TEX = torch.tensor([q["tex"] for q in Q], dtype=torch.int64, device=dev)
self.PALI = torch.tensor([q["pal"] for q in Q], dtype=torch.int64, device=dev)
self.BASE, self.STR = f32("base"), f32("strength")
self.ALPHA = f32("alpha")
self.UV0 = f32("uv")
self.kind = [q["kind"] for q in Q]
self.idx = {}
for i, q in enumerate(Q):
self.idx.setdefault(q["kind"], {})[(q["layer"], q["sub"])] = i
self.side_ids = {}
for i, q in enumerate(Q):
if q["kind"] == "side":
self.side_ids.setdefault(q["layer"], []).append(i)
print(f"[scene] {len(Q)} quads", flush=True)
# ------------------------------------------------------------------ timeline
class Timeline:
SUB = dict(ln1=(0.0, 0.10), qkv=(0.10, 0.28), attn=(0.28, 0.56), o=(0.56, 0.66), ln2=(0.66, 0.72),
gu=(0.72, 0.84), inter=(0.84, 0.92), down=(0.92, 1.0))
def __init__(self, P, T):
self.P = P
self.pf_spark0, self.pf_asc0, self.pf_asc1 = 4.0, 7.5, 16.0
steps = [16, 10, 7, 5, 3.5, 2.5, 2.0, 1.6, 1.3, 1.1]
t = self.pf_asc1
tok = []
k = 0
while True:
if k < len(steps):
s = steps[k]
else:
frac = (t - 65.0) / (112.0 - 65.0)
s = 1.0 - 0.4 * max(0.0, min(1.0, frac))
if t + s > 112.5 or P + k >= T:
break
pred, drop = 0.20 * s, 0.16 * s
tok.append(dict(pred0=t, drop0=t + pred, asc0=t + pred + drop, asc1=t + s))
t += s
k += 1
self.tok = tok
self.K = len(tok)
self.end_gen = t
T_used = P + self.K
self.T_used = T_used
ARR = np.full((NL + 1, T_used), np.inf)
ARR_INT = np.full((NL, T_used), np.inf)
ARR_SUB = {k: np.full((NL, T_used), np.inf) for k in self.SUB}
TAU = np.ones(T_used)
for i in range(P):
ARR[0, i] = self.pf_spark0 + 2.0 + 0.04 * i
slot = (self.pf_asc1 - self.pf_asc0) / NL
for l in range(NL):
ARR[l + 1, :P] = self.pf_asc0 + (l + 1) * slot
ARR_INT[l, :P] = self.pf_asc0 + (l + self.SUB["inter"][0]) * slot
for kk, (a, b) in self.SUB.items():
ARR_SUB[kk][l, :P] = self.pf_asc0 + (l + a) * slot
TAU[:P] = 1.2
for k, ev in enumerate(tok):
i = P + k
ARR[0, i] = ev["asc0"]
slot = (ev["asc1"] - ev["asc0"]) / NL
for l in range(NL):
ARR[l + 1, i] = ev["asc0"] + (l + 1) * slot
ARR_INT[l, i] = ev["asc0"] + (l + self.SUB["inter"][0]) * slot
for kk, (a, b) in self.SUB.items():
ARR_SUB[kk][l, i] = ev["asc0"] + (l + a) * slot
TAU[i] = max(0.35, min(2.0, slot * 3.0))
self.ARR, self.ARR_INT, self.ARR_SUB, self.TAU = ARR, ARR_INT, ARR_SUB, TAU
self.slot_of = np.array([(self.pf_asc1 - self.pf_asc0) / NL] * P + [(ev["asc1"] - ev["asc0"]) / NL for ev in tok])
print(f"[timeline] {self.K} generated tokens shown (of {T-P} captured), generation ends at {t:.1f}s", flush=True)
def pulse(self, t):
out = []
if self.pf_asc0 <= t < self.pf_asc1:
out.append((-1, (t - self.pf_asc0) / (self.pf_asc1 - self.pf_asc0) * NL))
for k, ev in enumerate(self.tok):
if ev["asc0"] <= t < ev["asc1"]:
out.append((self.P + k, (t - ev["asc0"]) / (ev["asc1"] - ev["asc0"]) * NL))
return out
def phase(self, t):
for k, ev in enumerate(self.tok):
if ev["pred0"] <= t < ev["drop0"]:
return "pred", k, (t - ev["pred0"]) / (ev["drop0"] - ev["pred0"])
if ev["drop0"] <= t < ev["asc0"]:
return "drop", k, (t - ev["drop0"]) / (ev["asc0"] - ev["drop0"])
return None, -1, 0.0
# ------------------------------------------------------------------ camera
class Camera:
def __init__(self, keys):
self.keys = sorted(keys, key=lambda k: k[0])
self.ts = np.array([k[0] for k in self.keys])
self.vals = np.array([list(k[1]) + list(k[2]) + [k[3]] for k in self.keys], dtype=np.float64)
def at(self, t):
ts, vals = self.ts, self.vals
n = len(ts)
if t <= ts[0]:
v = vals[0]
elif t >= ts[-1]:
v = vals[-1]
else:
i = int(np.searchsorted(ts, t, side="right") - 1)
i0, i1, i2, i3 = max(i - 1, 0), i, i + 1, min(i + 2, n - 1)
h = ts[i2] - ts[i1]
s = (t - ts[i1]) / h
def tangent(a, b, c, ta, tb, tc):
if tb == ta:
return (c - b) / (tc - tb)
if tc == tb:
return (b - a) / (tb - ta)
return 0.5 * ((b - a) / (tb - ta) + (c - b) / (tc - tb))
m1 = tangent(vals[i0], vals[i1], vals[i2], ts[i0], ts[i1], ts[i2]) * h
m2 = tangent(vals[i1], vals[i2], vals[i3], ts[i1], ts[i2], ts[i3]) * h
s2, s3 = s * s, s * s * s
v = (2 * s3 - 3 * s2 + 1) * vals[i1] + (s3 - 2 * s2 + s) * m1 + (-2 * s3 + 3 * s2) * vals[i2] + (s3 - s2) * m2
return v[0:3], v[3:6], float(v[6])
def orbit(ang, r, h):
return (r * math.cos(ang), h, r * math.sin(ang))
def default_camera(tl):
K = tl.tok
keys = []
def k(t, eye, tgt, fov=46):
keys.append((t, eye, tgt, fov))
F = math.pi / 2 # facade direction
k(0.0, orbit(F - 0.75, 2500, 950), (0, 820, 0), 40)
k(4.0, orbit(F - 0.45, 2300, 850), (0, 820, 0), 40)
k(7.5, orbit(F - 0.2, 700, 130), (0, 70, 0), 46)
k(10.0, orbit(F - 0.05, 520, 230), (0, 160, 0), 48)
k(13.0, orbit(F + 0.2, 520, 920), (0, 860, 0), 48)
k(16.0, orbit(F + 0.4, 520, 1730), (0, 1670, 0), 48)
e0 = K[0]
k(e0["pred0"] + 1.2, orbit(F + 0.6, 720, 2120), (0, 1720, 0), 50)
k(e0["drop0"] + 0.2, orbit(F + 0.8, 950, 1950), (0, 1450, 0), 50)
k(e0["asc0"], orbit(F + 1.0, 1100, 260), (0, 150, 0), 48)
a0, a1 = e0["asc0"], e0["asc1"]
for f in np.linspace(0, 1, 7):
ang = F + 0.6 - f * 1.2
yy = f * Y_TOP
e = orbit(ang, 215, yy + 26)
k(a0 + f * (a1 - a0), (e[0] - 12, e[1], e[2]), (-14, yy + 9, 0), 50)
e1 = K[1]
k(e1["pred0"] + 0.8, orbit(F - 0.9, 820, 2080), (0, 1710, 0), 48)
k(e1["asc0"], orbit(F - 0.7, 700, 320), (0, 210, 0), 48)
k(e1["asc1"], orbit(F - 0.3, 600, 1120), (0, 1010, 0), 48)
e3 = K[3]
k(e3["asc0"], orbit(F + 0.1, 650, 520), (0, 430, 0), 48)
e6 = K[6]
k(e6["asc0"], orbit(F + 0.5, 820, 920), (0, 810, 0), 46)
t = e6["asc0"]; t_end = 112.0
steps = 7
for j in range(1, steps + 1):
f = j / steps
tt = t + f * (t_end - t)
k(tt, orbit(F + 0.5 - f * 1.5, 900 + f * 1300, 820 + f * 200), (0, 830 - 40 * f, 0), 46 - 4 * f)
k(120.0, orbit(F - 1.3, 2400, 1000), (0, 820, 0), 42)
return Camera(keys)
# ------------------------------------------------------------------ renderer
class Renderer:
def __init__(self, scene, tl, cam, W, H, ss, dev):
self.sc, self.tl, self.cam = scene, tl, cam
self.W, self.H, self.ss = W, H, ss
self.Ws, self.Hs = W * ss, H * ss
self.dev = dev
A = scene.acts
self.attn, self.inter, self.probs = A["attn"], A["inter"], A["probs"]
self.tokens = scene.tokens
self.pal_sign = PAL_SIGN.to(dev); self.cmap = CMAP.to(dev); self.cmap_m = CMAP_M.to(dev)
self.headcol_np = np.array(HEAD_COL, np.float32)
self.light = torch.tensor([0.35, 1.0, 0.25], device=dev); self.light = self.light / self.light.norm()
self.prof = bool(os.environ.get("LLMVIZ_PROF"))
self._prof = {}
self.bits = 8
# ---- camera ---------------------------------------------------------
def setup_camera(self, eye, tgt, fov):
dev = self.dev
eye = torch.tensor(eye, dtype=torch.float32, device=dev)
tgt = torch.tensor(tgt, dtype=torch.float32, device=dev)
fwd = tgt - eye; fwd = fwd / fwd.norm()
up0 = torch.tensor([0.0, 1.0, 0.0], device=dev)
right = torch.cross(fwd, up0, dim=0); right = right / right.norm()
up = torch.cross(right, fwd, dim=0)
tanf = math.tan(math.radians(fov) / 2)
aspect = self.Ws / self.Hs
ys = (0.5 - (torch.arange(self.Hs, device=dev, dtype=torch.float32) + 0.5) / self.Hs) * 2 * tanf
xs = ((torch.arange(self.Ws, device=dev, dtype=torch.float32) + 0.5) / self.Ws - 0.5) * 2 * tanf * aspect
dirs = fwd[None, None] + xs[None, :, None] * right[None, None] + ys[:, None, None] * up[None, None]
self.dirs = dirs / dirs.norm(dim=-1, keepdim=True)
self.eye, self.fwd, self.right, self.up, self.tanf, self.aspect = eye, fwd, right, up, tanf, aspect
self.px_per_unit = self.Hs / (2 * tanf) # at depth 1
def project(self, p):
d = p - self.eye
zc = d @ self.fwd; xc = d @ self.right; yc = d @ self.up
ok = zc > 0.5
zs = zc.clamp(min=0.5)
px = (xc / zs / (self.tanf * self.aspect) * 0.5 + 0.5) * self.Ws
py = (0.5 - yc / zs / self.tanf * 0.5) * self.Hs
return px, py, zc, ok
# ---- rasterise --------------------------------------------------------
def raster(self, st):
sc, dev = self.sc, self.dev
Hs, Ws = self.Hs, self.Ws
flash = st["flash"]
npx = Hs * Ws
zbuf = torch.full((npx,), 1e30, device=dev)
acc_rgb = torch.zeros(npx, 3, device=dev)
acc_w = torch.zeros(npx, device=dev)
acc_log = torch.zeros(npx, device=dev)
P0, E1, E2 = sc.P0, sc.E1, sc.E2
N = P0.shape[0]
corners = torch.stack([P0, P0 + E1, P0 + E2, P0 + E1 + E2], 1)
px, py, dep, ok = self.project(corners.reshape(-1, 3))
px, py, ok = px.view(N, 4), py.view(N, 4), ok.view(N, 4)
anyok = ok.any(1)
big = anyok & ~ok.all(1)
x0 = torch.where(big, torch.zeros_like(px[:, 0]), px.min(1).values.floor()).clamp(0, Ws)
x1 = torch.where(big, torch.full_like(px[:, 0], Ws), px.max(1).values.ceil() + 1).clamp(0, Ws)
y0 = torch.where(big, torch.zeros_like(py[:, 0]), py.min(1).values.floor()).clamp(0, Hs)
y1 = torch.where(big, torch.full_like(py[:, 0], Hs), py.max(1).values.ceil() + 1).clamp(0, Hs)
vis = anyok & (x1 > x0) & (y1 > y0) & (flash > 0)
x0, x1, y0, y1 = x0.long(), x1.long(), y0.long(), y1.long()
ids = torch.nonzero(vis).squeeze(1)
if ids.numel() == 0:
return acc_rgb.view(Hs, Ws, 3), zbuf.view(Hs, Ws)
bw_c, bh_c = (x1 - x0)[ids].cpu().numpy(), (y1 - y0)[ids].cpu().numpy()
ids_c = ids.cpu().numpy()
# buckets by size class (sqrt2 steps of the larger side, pow2 of the smaller)
buckets = {}
for i, w_, h_ in zip(ids_c, bw_c, bh_c):
key = (int(math.ceil(math.log2(max(h_, 1)))), int(math.ceil(math.log2(max(w_, 1)))))
buckets.setdefault(key, []).append((int(i), int(h_), int(w_)))
MAXPIX = 40_000_000 if dev.type == "cuda" else 6_000_000
groups = []
for key, items in buckets.items():
th = max(it[1] for it in items); tw = max(it[2] for it in items)
per = max(1, MAXPIX // max(th * tw, 1))
for s in range(0, len(items), per):
chunk = items[s:s + per]
groups.append((torch.tensor([c[0] for c in chunk], device=dev), th, tw))
geo = dict(x0=x0, y0=y0, x1=x1, y1=y1)
# pass 1: nearest depth
for (qi, th, tw) in groups:
g = self._geom(qi, th, tw, geo, st)
if g is None:
continue
idx, inside, t = g["idx"], g["inside"], g["t"]
zbuf.scatter_reduce_(0, idx[inside], t[inside], "amin")
self.mark("pass1")
# pass 2: weighted blended translucent shading
npix = 0
for (qi, th, tw) in groups:
g = self._geom(qi, th, tw, geo, st)
if g is None:
continue
idx, inside, t = g["idx"], g["inside"], g["t"]
w = torch.exp(-(t - zbuf[idx]).clamp(min=0) / DEPTH_SOFT)
sel = inside & (w > 0.04)
if not sel.any():
continue
rgb, alpha = self._shade(qi, g, sel, st)
ww = w[sel] * alpha
ii = idx[sel]
acc_rgb.index_put_((ii,), rgb * ww[:, None], accumulate=True)
acc_w.index_put_((ii,), ww, accumulate=True)
acc_log.index_put_((ii,), torch.log1p(-alpha), accumulate=True)
npix += int(sel.sum()) if self.prof else 0
self.mark("pass2")
reveal = torch.exp(acc_log)
color = acc_rgb / acc_w.clamp(min=1e-6)[:, None] * (1 - reveal)[:, None]
if self.prof:
print(f"[raster] {ids.numel()} quads {len(groups)} groups {npix/1e6:.1f} Mpix shaded", flush=True)
return color.view(Hs, Ws, 3), zbuf.view(Hs, Ws)
def mark(self, name):
if self.prof:
sync(self.dev)
now = time.time()
self._prof[name] = self._prof.get(name, 0.0) + now - self._t0
self._t0 = now
def _geom(self, qi, th, tw, geo, st):
sc, dev = self.sc, self.dev
n = qi.numel()
X0, Y0, X1, Y1 = geo["x0"], geo["y0"], geo["x1"], geo["y1"]
ys = Y0[qi][:, None] + torch.arange(th, device=dev)[None]
xs = X0[qi][:, None] + torch.arange(tw, device=dev)[None]
vy = ys < Y1[qi][:, None]; vx = xs < X1[qi][:, None]
ysc = ys.clamp(max=self.Hs - 1); xsc = xs.clamp(max=self.Ws - 1)
valid = vy[:, :, None] & vx[:, None, :]
dirs = self.dirs[ysc[:, :, None], xsc[:, None, :]]
p0, e1, e2 = sc.P0[qi], sc.E1[qi], sc.E2[qi]
nrm = torch.cross(e1, e2, dim=1); nrm = nrm / nrm.norm(dim=1, keepdim=True)
denom = (dirs * nrm[:, None, None]).sum(-1)
num = ((p0 - self.eye) * nrm).sum(-1)
safe = denom.abs() > 1e-7
t = num[:, None, None] / torch.where(safe, denom, torch.ones_like(denom))
hit = safe & (t > 0.3)
Pw = self.eye + t[..., None] * dirs
dd = Pw - p0[:, None, None]
u = (dd * e1[:, None, None]).sum(-1) / (e1 * e1).sum(1)[:, None, None]
v = (dd * e2[:, None, None]).sum(-1) / (e2 * e2).sum(1)[:, None, None]
inside = hit & valid & (u >= 0) & (u < 1) & (v >= 0) & (v < 1)
rm_off, rm_len, rm_buf = st["rm_off"], st["rm_len"], st["rm_buf"]
ro, rl = rm_off[qi], rm_len[qi]
ridx = (ro[:, None, None] + (u.clamp(0, 0.999999) * rl[:, None, None].float()).long()).clamp(0, rm_buf.shape[0] - 1)
rmv = rm_buf[ridx] # [n,th,tw,2]
inside = inside & (rmv[..., 1] > 0)
if not inside.any():
return None
idx = (ysc[:, :, None] * self.Ws + xsc[:, None, :]).expand(n, th, tw)
# depth is measured along the camera forward axis (consistent across quads)
zc = (t * (dirs @ self.fwd))
return dict(idx=idx, inside=inside, t=zc, u=u, v=v, dirs=dirs, nrm=nrm, rmb=rmv[..., 0])
def _shade(self, qi, g, sel, st):
sc, A, dev = self.sc, self.sc.atlas, self.dev
n = qi.numel()
u, v = g["u"], g["v"]
uv = st["uv"][qi]
Hb = A.base_H[sc.TEX[qi]] * (uv[:, 2] - uv[:, 0]); Wb = A.base_W[sc.TEX[qi]] * (uv[:, 3] - uv[:, 1])
def grad(a):
gx = torch.zeros_like(a); gy = torch.zeros_like(a)
if a.shape[2] > 1:
gx[:, :, 1:] = a[:, :, 1:] - a[:, :, :-1]; gx[:, :, 0] = gx[:, :, 1]
if a.shape[1] > 1:
gy[:, 1:, :] = a[:, 1:, :] - a[:, :-1, :]; gy[:, 0, :] = gy[:, 1, :]
return gx, gy
ux, uy = grad(u); vx_, vy_ = grad(v)
Hb3, Wb3 = Hb[:, None, None], Wb[:, None, None]
rho = torch.maximum((ux * Hb3) ** 2 + (vx_ * Wb3) ** 2, (uy * Hb3) ** 2 + (vy_ * Wb3) ** 2).clamp(min=1e-12)
lod = 0.5 * torch.log2(rho) + 0.3
qsel = torch.arange(n, device=dev)[:, None, None].expand(n, u.shape[1], u.shape[2])[sel]
q_ids = qi[qsel]
uu, vv, ll = u[sel], v[sel], lod[sel]
ut = uv[qsel, 0] + uu * (uv[qsel, 2] - uv[qsel, 0])
vt = uv[qsel, 1] + vv * (uv[qsel, 3] - uv[qsel, 1])
val, mag = A.sample(sc.TEX[q_ids], ut, vt, ll)
pali = sc.PALI[q_ids]
m = mag.clamp(min=0)
isnw = pali == 2
if isnw.any():
m = torch.where(isnw, 1.0 + (val - 1.0).abs() * 1.5, m)
# colormap by magnitude (piecewise linear over CMAP_M stops)
mm = m.clamp(0, float(self.cmap_m[-1]))
seg = torch.bucketize(mm, self.cmap_m[1:-1]) # 0..4
m0 = self.cmap_m[seg]; m1 = self.cmap_m[seg + 1]
f = ((mm - m0) / (m1 - m0)).clamp(0, 1)[:, None]
c0 = self.cmap[pali, seg]; c1 = self.cmap[pali, seg + 1]
col = c0 * (1 - f) + c1 * f
lum = (col * torch.tensor([0.30, 0.59, 0.11], device=self.dev)).sum(1, keepdim=True)
# sign tint where the mean is resolved (fine mip levels)
sgn = (val / (mag + 1e-6)).clamp(-1, 1)
w = (sgn + 1) * 0.5
sign_col = self.pal_sign[pali, 0] * (1 - w[:, None]) + self.pal_sign[pali, 1] * w[:, None]
sat = (sgn.abs().clamp(0, 1) ** 0.8)[:, None] * 0.75
col = col * (1 - sat) + sign_col * lum * 1.6 * sat
bright = torch.ones_like(m)
rgb = col * (sc.BASE[q_ids] + sc.STR[q_ids] * bright)[:, None]
rgb = rgb * (st["flash"][q_ids] * g["rmb"][sel])[:, None]
# outline
e1n = sc.E1[q_ids].norm(dim=1); e2n = sc.E2[q_ids].norm(dim=1)
edge = torch.minimum(torch.minimum(uu, 1 - uu) * e1n, torch.minimum(vv, 1 - vv) * e2n)
edge_f = (1 - (edge / 0.14).clamp(0, 1)) ** 2
rgb = rgb + edge_f[:, None] * 0.55 * self.cmap[pali, 2] * st["flash"][q_ids][:, None]
# lighting: key light + grazing darkening + distance fog
nrm = g["nrm"][qsel]
ndl = (nrm @ self.light).abs()
cosv = (g["dirs"][sel] * nrm).sum(-1).abs()
shade = (0.65 + 0.35 * ndl) * (0.6 + 0.4 * cosv)
fog = torch.exp(-g["t"][sel] * 0.00025)
rgb = rgb * (shade * fog)[:, None]
alpha = sc.ALPHA[q_ids]
return rgb, alpha
# ---- glow -------------------------------------------------------------
def glow_begin(self):
self._glow = {k: ([], []) for k in ("s", "m", "l", "xl")}
def add_glow(self, pts, cols, cls="m", line=False):
"""pts [M,3]; cols [M,3]. line=True: cols = per-pixel target brightness x world length of the sample
(scaled at splat time by pixels-per-unit so lines look the same at any distance); else absolute energy."""
if isinstance(pts, np.ndarray):
pts = torch.from_numpy(pts.astype(np.float32)).to(self.dev)
if isinstance(cols, np.ndarray):
cols = torch.from_numpy(cols.astype(np.float32)).to(self.dev)
if pts.shape[0] == 0:
return
flag = torch.full((pts.shape[0], 1), 1.0 if line else 0.0, device=self.dev)
self._glow[cls][0].append(pts); self._glow[cls][1].append(torch.cat([cols, flag], 1))
def glow_end(self, zbuf):
Hs, Ws, dev = self.Hs, self.Ws, self.dev
out = torch.zeros(Hs * Ws, 3, device=dev)
sig = dict(s=1.4, m=2.4, l=5.0, xl=12.0)
for cls, (pl, cl) in self._glow.items():
if not pl:
continue
pts = torch.cat(pl); cols = torch.cat(cl)
px, py, zc, ok = self.project(pts)
ok = ok & (px > -4) & (px < Ws + 4) & (py > -4) & (py < Hs + 4)
if not ok.any():
continue
px, py, zc, cols = px[ok], py[ok], zc[ok], cols[ok]
s_px = sig[cls] * self.ss
ppu = self.px_per_unit / zc.clamp(min=1.0) # pixels per world unit at that depth
line_scale = ppu * (2.5 * s_px)
dist_att = (ppu / (2.0 * self.ss)).clamp(0.08, 1.0)
cols = cols[:, :3] * torch.where(cols[:, 3:4] > 0.5, line_scale[:, None], torch.ones_like(line_scale)[:, None]) * dist_att[:, None]
xi = px.long().clamp(0, Ws - 1); yi = py.long().clamp(0, Hs - 1)
zs = zbuf.view(-1)[yi * Ws + xi]
occl = ((zs - zc) / (0.015 * zc + 0.4)).clamp(0, 1) # 1 = in front of surfaces
cols = cols * (0.10 + 0.90 * occl)[:, None] * torch.exp(-zc * 0.00025)[:, None]
buf = torch.zeros(Hs * Ws, 3, device=dev)
fx = px.floor(); fy = py.floor(); ax = px - fx; ay = py - fy
fx = fx.long(); fy = fy.long()
for (dx, dy, wgt) in ((0, 0, (1 - ax) * (1 - ay)), (1, 0, ax * (1 - ay)), (0, 1, (1 - ax) * ay), (1, 1, ax * ay)):
xx = fx + dx; yy = fy + dy
m = (xx >= 0) & (xx < Ws) & (yy >= 0) & (yy < Hs)
buf.index_put_(((yy * Ws + xx)[m],), cols[m] * wgt[m][:, None], accumulate=True)
img = buf.view(Hs, Ws, 3).permute(2, 0, 1)[None]
img = gauss_blur(img, s_px)
out += img[0].permute(1, 2, 0).reshape(-1, 3)
return out.view(Hs, Ws, 3)
def _dens(self, a, b, n_per_unit):
"""samples per world unit so that consecutive samples are <= ~2 px apart on screen"""
mid = torch.from_numpy(((a + b) * 0.5).astype(np.float32)).to(self.dev)
zc = self.project(mid)[2].clamp(min=1.0).cpu().numpy()
ppu = self.px_per_unit / zc
return np.maximum(n_per_unit, ppu / 2.0)
def poly_pts(self, poly, bright, n_per_unit=2.0):
"""dense samples along a polyline [K,3]; bright [K,3] per-vertex per-pixel brightness (line mode)."""
a, b = poly[:-1], poly[1:]
L = np.linalg.norm(b - a, axis=1)
dens = self._dens(a, b, n_per_unit)
out_p, out_c = [], []
for i in range(len(a)):
n = max(2, int(L[i] * dens[i]))
t = np.linspace(0, 1, n, endpoint=False)[:, None]
out_p.append(a[i] + (b[i] - a[i]) * t)
out_c.append((bright[i] * (1 - t) + bright[i + 1] * t) * (L[i] / n))
return np.concatenate(out_p), np.concatenate(out_c)
def line_pts(self, a, b, brightness_per_unit, n_per_unit=2.0):
"""dense points along segments a->b [M,3]; brightness per world unit of length."""
L = np.linalg.norm(b - a, axis=1)
dens = self._dens(a, b, n_per_unit)
out_p, out_c = [], []
for i in range(len(a)):
n = max(2, min(int(L[i] * dens[i]), 4000))
s = np.linspace(0, 1, n)[:, None]
out_p.append(a[i] + (b[i] - a[i]) * s)
out_c.append(np.repeat(brightness_per_unit[i][None], n, 0) * (L[i] / n))
return np.concatenate(out_p), np.concatenate(out_c)
def glow_prims(self, t, st):
sc, tl, dev = self.sc, self.tl, self.dev
P, T = tl.P, tl.T_used
latest = st["latest"]
ph, k, frac = st["phase"]
ppu = self.px_per_unit
# ---- residual beams: faint vertical lines per token up to the highest sheet reached
reached = (t - tl.ARR >= 0)
top = reached.sum(0) - 1
toks = np.where(top >= 0)[0]
if len(toks):
a = np.stack([tok_x(toks), np.zeros(len(toks)), np.zeros(len(toks))], 1)
b = a.copy(); b[:, 1] = np.minimum(top[toks], NL) * LAYER_DY
pts, cols = self.line_pts(a, b, np.tile(np.array([[0.25, 0.45, 0.9]]) * 0.006, (len(toks), 1)), 1.0)
self.add_glow(pts, cols, "s", line=True)
# ---- attention arcs
idxs, ages = latest["attn"]
arcs = []
for l in range(NL):
if idxs[l] < 0:
continue
i = idxs[l]
tau = tl.TAU[i] * 1.2
if ages[l] > tau * 2.2:
continue
fade = math.exp(-ages[l] / tau)
slot = tl.slot_of[i]
a0, a1 = tl.SUB["attn"]
grow = min(1.0, max(0.0, (ages[l] + (a1 - a0) * slot) / max((a1 - a0) * slot, 1e-6)))
qi_list = list(range(P)) if i < P else [i]
for q in qi_list:
w = self.attn[l, :, q, :q + 1].astype(np.float32)
thr = 0.06 if i < P else 0.03
# keep at most 6 strongest keys per head
keep = np.zeros_like(w, dtype=bool)
kk = min(6, w.shape[1])
top_j = np.argpartition(-w, kk - 1, axis=1)[:, :kk]
keep[np.arange(NH)[:, None], top_j] = True
hs, js = np.nonzero((w > thr) & keep)
if len(hs) == 0:
continue
ww = w[hs, js]
y = l * LAYER_DY + 0.6
xa = np.full(len(hs), tok_x(q)); xb = tok_x(js)
zz = head_z(hs)
a = np.stack([xa, np.full(len(hs), y), zz], 1); b = np.stack([xb, np.full(len(hs), y), zz], 1)
b = a + (b - a) * grow
hgt = np.minimum((2.5 + 0.8 * np.abs(xb - xa) * grow) * (0.7 + 0.6 * hs / (NH - 1)), 44)
c = self.headcol_np[hs] * (ww ** 0.6)[:, None] * fade * (0.35 if i >= P else 0.12)
arcs.append((a, b, hgt, c))
if arcs:
a = np.concatenate([x[0] for x in arcs]); b = np.concatenate([x[1] for x in arcs])
h = np.concatenate([x[2] for x in arcs]); c = np.concatenate([x[3] for x in arcs])
self._draw_arcs(a, b, h, c)
# ---- MLP sparks
idxs, ages = latest["inter"]
for l in range(NL):
i = idxs[l]
if i < 0:
continue
tau = tl.TAU[i] * 1.2
if ages[l] > tau * 3:
continue
fade = math.exp(-ages[l] / tau)
vec = self.inter[l, i].astype(np.float32)
scale = np.sqrt((vec ** 2).mean()) + 1e-6
top = np.argpartition(-np.abs(vec), 48)[:48]
mag = np.abs(vec[top]) / scale
rise = min(1.0, ages[l] / max(tau * 0.4, 1e-3))
xx = 36 + (top + 0.5) / I * 96
zz = -T * TOK_DX / 2 + (i + 0.5) * TOK_DX
yy = l * LAYER_DY + 1.0 + rise * (33.0 + 4.0 * np.minimum(mag, 6))
pts = np.stack([xx, yy, np.full(48, zz)], 1)
col = np.where((vec[top] > 0)[:, None], np.array([[1.0, 0.45, 0.7]]), np.array([[0.2, 0.85, 1.0]]))
cols = col * (fade * 18.0 * np.minimum(mag, 6) / 3)[:, None]
self.add_glow(pts, cols, "m")
# ---- pulses
for (i, lf) in tl.pulse(t):
toks_ = list(range(P)) if i < 0 else [i]
l = int(lf); f = lf - l
y = lf * LAYER_DY
n = len(toks_)
xs_ = tok_x(np.array(toks_))
a = np.stack([xs_, np.full(n, max(y - 20, 0)), np.zeros(n)], 1)
b = np.stack([xs_, np.full(n, y), np.zeros(n)], 1)
base = np.tile(np.array([[0.75, 0.88, 1.0]]) * (0.35 if i >= 0 else 0.12), (n, 1))
pts, cols = self.line_pts(a, b, base * 1.2, 3.0)
frac_ = ((pts[:, 1] - a[0, 1]) / max(y - a[0, 1], 1e-6)) ** 2
self.add_glow(pts, cols * frac_[:, None], "m", line=True)
head = np.stack([xs_, np.full(n, y), np.zeros(n)], 1)
self.add_glow(head, np.tile(np.array([[30.0, 34.0, 40.0]]) * (1.0 if i >= 0 else 0.12), (n, 1)), "l")
if i >= 0:
self.add_glow(head, np.tile(np.array([[40.0, 46.0, 60.0]]), (n, 1)), "xl")
slot = (tl.pf_asc1 - tl.pf_asc0) / NL if i < 0 else tl.slot_of[i]
if slot > 0.06 and i >= 0:
self._tour_spark(i, l, f)
# ---- prediction: top candidate tokens glow on the lm_head ring
if getattr(self, "_pred", None) is not None:
top, pw, env = self._pred
ang = vocab_angle(top.astype(np.float64))
base = np.stack([RING_R * np.cos(ang), np.full(len(top), RING_HI_Y + HID_U), RING_R * np.sin(ang)], 1)
self.add_glow(base, np.array([[1.0, 0.95, 0.8]]) * (env * 18.0 * np.sqrt(pw))[:, None], "l")
# pillars of light: height proportional to probability
tip = base.copy(); tip[:, 1] += 24 + 130 * pw * env
pts, cols = self.line_pts(base, tip, np.tile(np.array([[0.9, 0.85, 0.7]]) * 1.2 * env, (len(top), 1)) * np.sqrt(pw)[:, None], 1.5)
self.add_glow(pts, cols, "m", line=True)
# ---- prediction / drop
if ph in ("pred", "drop"):
tokid = int(self.tokens[P + k])
ang = vocab_angle(tokid)
xr, zr = RING_R * math.cos(ang), RING_R * math.sin(ang)
topP = np.array([xr, RING_HI_Y + HID_U / 2, zr])
if ph == "pred":
s = min(1.0, max(0.0, (frac - 0.45) / 0.55))
self.add_glow(topP[None], np.array([[26.0, 26.0, 28.0]]) * s, "l")
a = np.array([[xr, RING_HI_Y, zr]]); b = np.array([[xr, RING_HI_Y + HID_U, zr]])
pts, cols = self.line_pts(a, b, np.array([[0.9, 0.9, 1.0]]) * 3 * s, 2.0)
self.add_glow(pts, cols, "m", line=True)
else:
botP = np.array([xr, RING_LO_Y + HID_U / 2, zr])
dst = np.array([tok_x(P + k), 0.0, 0.0])
mid = (botP + dst) * 0.5 + np.array([0, 90.0, 0])
def pos(fj):
g1 = min(1.0, fj / 0.45); g2 = max(0.0, (fj - 0.45) / 0.55)
if g2 <= 0:
return topP + (botP - topP) * g1
return (1 - g2) ** 2 * botP + 2 * (1 - g2) * g2 * mid + g2 ** 2 * dst
trail = [pos(frac - j * 0.003) for j in range(40) if frac - j * 0.003 >= 0]
trail = np.array(trail)
if len(trail) > 1:
wts = np.linspace(1, 0, len(trail)) ** 2
pts, cols = self.poly_pts(trail, np.array([[0.9, 0.9, 1.0]]) * 2.5 * wts[:, None], 1.5)
self.add_glow(pts, cols, "m", line=True)
self.add_glow(trail[:1], np.array([[40.0, 40.0, 46.0]]), "l")
self.add_glow(trail[:1], np.array([[50.0, 50.0, 60.0]]), "xl")
f2 = max(0.0, (frac - 0.45) / 0.55)
if f2 > 0:
a = np.array([[xr, RING_LO_Y, zr]]); b = np.array([[xr, RING_LO_Y + HID_U, zr]])
pts, cols = self.line_pts(a, b, np.array([[0.8, 0.85, 1.0]]) * 3 * (1 - f2), 2.0)
self.add_glow(pts, cols, "m", line=True)
# ---- prefill sparks
if tl.pf_spark0 <= t < tl.pf_asc0 + 0.5:
heads, hcols, trails, tcols = [], [], [], []
for i in range(P):
t0 = tl.pf_spark0 + 0.04 * i
fr = (t - t0) / 2.0
if fr < 0 or fr > 1.15:
continue
tokid = int(self.tokens[i]); ang = vocab_angle(tokid)
src = np.array([RING_R * math.cos(ang), RING_LO_Y + HID_U / 2, RING_R * math.sin(ang)])
dst = np.array([tok_x(i), 0.0, 0.0])
mid = (src + dst) * 0.5 + np.array([0, 110.0, 0])
gs = np.array([min(1.0, fr) - j * 0.01 for j in range(40)])
gs = gs[gs >= 0]
if len(gs) == 0:
continue
p = (1 - gs)[:, None] ** 2 * src + 2 * ((1 - gs) * gs)[:, None] * mid + (gs[:, None] ** 2) * dst
heads.append(p[0]); hcols.append([12.0, 12.0, 14.0])
if len(p) > 1:
wts = np.linspace(1, 0, len(p)) ** 2
pp, cc = self.poly_pts(p, np.array([[0.8, 0.85, 1.0]]) * 1.5 * wts[:, None], 1.5)
trails.append(pp); tcols.append(cc)
a = np.array([[src[0], RING_LO_Y, src[2]]]); b = np.array([[src[0], RING_LO_Y + HID_U, src[2]]])
pts, cols = self.line_pts(a, b, np.array([[0.6, 0.7, 0.9]]) * 3 * (1 - min(fr, 1)), 2.0)
trails.append(pts); tcols.append(cols)
if heads:
self.add_glow(np.array(heads), np.array(hcols) * 2.5, "l")
self.add_glow(np.concatenate(trails), np.concatenate(tcols), "m", line=True)
def _draw_arcs(self, a, b, h, c):
M = len(a)
# sample count from projected size
pa = self.project(torch.from_numpy(a.astype(np.float32)).to(self.dev)); pb = self.project(torch.from_numpy(b.astype(np.float32)).to(self.dev))
L = (torch.sqrt((pa[0] - pb[0]) ** 2 + (pa[1] - pb[1]) ** 2) + torch.from_numpy(h.astype(np.float32)).to(self.dev) * 2 * self.px_per_unit / pa[2].clamp(min=1)).cpu().numpy()
n_all = np.clip(L / (2.0 * self.ss), 6, 300).astype(int)
for lo, hi, n in ((0, 12, 12), (12, 40, 40), (40, 120, 120), (120, 10 ** 9, 300)):
m = (n_all >= lo) & (n_all < hi)
if not m.any():
continue
aa, bb, hh, cc = a[m], b[m], h[m], c[m]
s = np.linspace(0, 1, n)[None, :, None]
mid = (aa + bb) * 0.5 + np.array([0, 1, 0]) * hh[:, None]
p = (1 - s) ** 2 * aa[:, None] + 2 * (1 - s) * s * mid[:, None] + s ** 2 * bb[:, None]
# arc length per sample (world units) so brightness is per unit length
seg = np.linalg.norm(np.diff(p, axis=1), axis=2)
seg = np.concatenate([seg, seg[:, -1:]], 1)
cols = cc[:, None, :] * seg[:, :, None] * 1.3
self.add_glow(p.reshape(-1, 3), cols.reshape(-1, 3), "m", line=True)
def _tour_spark(self, i, l, f):
y = l * LAYER_DY
x = tok_x(i)
zi = -self.tl.T_used * TOK_DX / 2 + (i + 0.5) * TOK_DX
path = [((x, y + 0.6, 0), 0.0), ((-45.5, y + 0.6, 0), 0.10), ((-112, y + 16, 1.5), 0.19), ((-92, y + 16, 1.5), 0.24),
((-86, y + 16, 1.5), 0.28), ((x, y + 5, 0), 0.42), ((x, y + 0.6, 0), 0.56), ((-64, y + 16, 1.5), 0.62),
((x, y + 0.6, 0), 0.66), ((29, y + 0.6, 0), 0.72), ((84, y + 35, -20), 0.78), ((84, y + 35, 20), 0.84),
((84, y + 1.0, zi), 0.90), ((84, y + 35, 0), 0.95), ((x, y + 0.6, 0), 0.98), ((x, y + LAYER_DY, 0), 1.0)]
def pos(fj):
for s in range(len(path) - 1):
(p0, t0), (p1, t1) = path[s], path[s + 1]
if t0 <= fj <= t1:
g = (fj - t0) / max(t1 - t0, 1e-6)
return np.array(p0) * (1 - g) + np.array(p1) * g
return np.array(path[-1][0])
trail = np.array([pos(f - j * 0.003) for j in range(40) if f - j * 0.003 >= 0])
if len(trail) < 2:
return
wts = np.linspace(1, 0, len(trail)) ** 2
pts, cols = self.poly_pts(trail, np.array([[1.0, 0.85, 0.55]]) * 2.5 * wts[:, None], 2.0)
self.add_glow(pts, cols, "m", line=True)
self.add_glow(trail[:1], np.array([[24.0, 20.0, 14.0]]), "l")
# ---- frame state ------------------------------------------------------
def frame_state(self, t):
sc, tl, dev = self.sc, self.tl, self.dev
N = sc.P0.shape[0]
P, T = tl.P, tl.T_used
flash = np.ones(N, np.float32)
if not hasattr(self, "_uv0"):
self._uv0 = sc.UV0.cpu().numpy()
uv = self._uv0.copy()
rm_off = np.zeros(N, np.int64); rm_len = np.ones(N, np.int64)
rm = [np.array([[1.0, 1.0]], np.float32)]
cursor = [1]
def push(arr):
rm.append(arr.astype(np.float32)); o = cursor[0]; cursor[0] += arr.shape[0]; return o
age = t - tl.ARR
vis = age >= 0
tau = tl.TAU[None, :]
br = 1.0 + 1.2 * np.exp(-np.maximum(age, 0) / tau) * vis
Tfull = sc.T
for l in range(NL + 1):
i = sc.idx["sheet"][(l, 0)]
arr = np.zeros((Tfull, 2), np.float32)
arr[:T, 0] = br[l]; arr[:T, 1] = vis[l]
rm_off[i] = push(arr); rm_len[i] = Tfull
age_i = t - tl.ARR_INT
vis_i = age_i >= 0
br_i = 1.0 + 1.5 * np.exp(-np.maximum(age_i, 0) / tau) * vis_i
for l in range(NL):
i = sc.idx["inter"][(l, 0)]
arr = np.zeros((Tfull, 2), np.float32)
arr[:T, 0] = br_i[l]; arr[:T, 1] = vis_i[l]
rm_off[i] = push(arr); rm_len[i] = Tfull
strip_map = dict(s_ln1="ln1", s_q="qkv", s_k="qkv", s_v="qkv", s_o="o", s_ln2="ln2", s_mlp="down")
slab_map = dict(wq="qkv", wk="qkv", wv="qkv", wo="o", wg="gu", wu="gu", wd="down", ln1w="ln1", ln2w="ln2")
latest = {}
for key in tl.SUB:
a = t - tl.ARR_SUB[key]
m = a >= 0
idxs = np.where(m.any(1), T - 1 - np.argmax(m[:, ::-1], axis=1), -1)
ages = np.where(idxs >= 0, a[np.arange(NL), np.clip(idxs, 0, T - 1)], np.inf)
latest[key] = (idxs, ages)
for kind, key in strip_map.items():
idxs, ages = latest[key]
for l in range(NL):
i = sc.idx[kind][(l, 0)]
if idxs[l] < 0:
flash[i] = 0.0
else:
uv[i, 0] = idxs[l] / Tfull; uv[i, 2] = (idxs[l] + 1) / Tfull
flash[i] = 1.0 + 1.2 * math.exp(-ages[l] / tl.TAU[idxs[l]])
for kind, key in slab_map.items():
idxs, ages = latest[key]
for l in range(NL):
i = sc.idx[kind][(l, 0)]
if idxs[l] >= 0:
flash[i] = 1.0 + 0.8 * math.exp(-ages[l] / (0.6 * tl.TAU[idxs[l]]))
ph, k, frac = tl.phase(t)
ring_hi = np.ones((V, 2), np.float32)
ring_lo = np.ones((V, 2), np.float32)
show_k, env = None, 0.0
if ph == "pred":
show_k, env = k, min(1.0, frac * 3.0)
elif ph == "drop":
show_k, env = k, max(0.0, 1.0 - frac * 1.2)
self._pred = None
if show_k is not None:
pr = self.probs[show_k].astype(np.float32)
pm = max(float(pr.max()), 1e-6)
# spread each token's probability over ~3 world units of ring so it is visible
ker = np.exp(-0.5 * (np.arange(-240, 241) / 80.0) ** 2); ker /= ker.sum()
spread = np.convolve(pr / pm, ker, mode="same")
spread = spread / max(float(spread.max()), 1e-9)
ring_hi[:, 0] = 1.0 + env * 7.0 * np.sqrt(spread)
top = np.argpartition(-pr, 32)[:32]
self._pred = (top, pr[top] / pm, env)
if ph == "pred":
i = sc.idx["normw"][(NL, 0)]; flash[i] = 1.0 + 2.5 * math.sin(math.pi * min(frac, 1))
o_hi = push(ring_hi); o_lo = push(ring_lo)
seg = V // RING_SEGS
for s in range(RING_SEGS):
ii = sc.idx["ring_hi"][(-1, s)]; rm_off[ii] = o_hi + s * seg; rm_len[ii] = seg
ii = sc.idx["ring_lo"][(-1, s)]; rm_off[ii] = o_lo + s * seg; rm_len[ii] = seg
rm_buf = np.concatenate(rm, 0)
return dict(flash=torch.from_numpy(flash).to(dev), uv=torch.from_numpy(uv).to(dev),
rm_off=torch.from_numpy(rm_off).to(dev), rm_len=torch.from_numpy(rm_len).to(dev),
rm_buf=torch.from_numpy(rm_buf).to(dev), latest=latest, phase=(ph, k, frac))
# ---- post -------------------------------------------------------------
def post(self, color, glow, fade):
ss = self.ss
glow = glow / (1 + glow / 8.0)
img = (color + glow).permute(2, 0, 1)[None]
if ss > 1:
img = F.avg_pool2d(img, ss)
bright = (img - 1.0).clamp(min=0)
acc = torch.zeros_like(img)
cur = bright
wts = [0.5, 0.45, 0.4, 0.35, 0.3, 0.25]
for i in range(6):
cur = F.avg_pool2d(cur, 2, ceil_mode=True)
blur = gauss_blur(cur, 1.5)
acc = acc + F.interpolate(blur, size=img.shape[-2:], mode="bilinear", align_corners=False) * wts[i]
img = (img + acc * 0.3) * fade
x = img.clamp(min=0)
x = (x * (2.51 * x + 0.03)) / (x * (2.43 * x + 0.59) + 0.14)
x = x.clamp(0, 1) ** (1 / 2.2)
H, W = x.shape[-2:]
yy = torch.linspace(-1, 1, H, device=x.device)[:, None]; xx = torch.linspace(-1, 1, W, device=x.device)[None]
x = x * (1 - 0.25 * (xx ** 2 + yy ** 2 * 0.8) ** 1.2)
if self.bits == 16:
x = x + (torch.rand_like(x) - 0.5) / 65535.0
return (x.clamp(0, 1) * 65535 + 0.5).to(torch.int32)[0].permute(1, 2, 0).contiguous()
x = x + (torch.rand_like(x) - 0.5) / 255.0
return (x.clamp(0, 1) * 255 + 0.5).to(torch.uint8)[0].permute(1, 2, 0).contiguous()
def render(self, frame):
t = frame / FPS
self._t0 = time.time(); self._prof = {}
eye, tgt, fov = self.cam.at(t)
if os.environ.get("LLMVIZ_CAM"):
c = [float(x) for x in os.environ["LLMVIZ_CAM"].split(",")]
eye, tgt, fov = c[0:3], c[3:6], c[6]
self.setup_camera(eye, tgt, fov)
st = self.frame_state(t); self.mark("state")
color, zbuf = self.raster(st)
self.glow_begin()
self.glow_prims(t, st); self.mark("glow_build")
glow = self.glow_end(zbuf); self.mark("glow_splat")
fade = min(1.0, t / 2.5) * min(1.0, max(0.0, (DUR - t) / 2.5))
out = self.post(color, glow, fade); self.mark("post")
if self.prof:
print(f"[prof] t={t:.2f} " + " ".join(f"{k}={v:.2f}" for k, v in self._prof.items()), flush=True)
return out
def gauss_blur(x, sigma):
r = max(1, int(math.ceil(sigma * 3)))
k = torch.exp(-torch.arange(-r, r + 1, device=x.device, dtype=torch.float32) ** 2 / (2 * sigma * sigma))
k = k / k.sum()
C = x.shape[1]
x = F.conv2d(F.pad(x, (r, r, 0, 0), mode="replicate"), k.view(1, 1, 1, -1).expand(C, 1, 1, -1), groups=C)
x = F.conv2d(F.pad(x, (0, 0, r, r), mode="replicate"), k.view(1, 1, -1, 1).expand(C, 1, -1, 1), groups=C)
return x
# ------------------------------------------------------------------ loading / main
def load_weights(model_dir):
from safetensors import safe_open
files = [os.path.join(model_dir, f) for f in os.listdir(model_dir) if f.endswith(".safetensors")]
handles = [safe_open(f, framework="pt") for f in files]
table = {}
for h in handles:
for k in h.keys():
table[k] = h
return lambda name: table[name].get_tensor(name)
def build(args):
if args.device:
dev = torch.device(args.device)
elif torch.cuda.is_available():
dev = torch.device("cuda")
elif torch.backends.mps.is_available():
dev = torch.device("mps")
else:
dev = torch.device("cpu")
print("[main] device", dev, flush=True)
acts = np.load(args.acts)
acts = {k: acts[k] for k in acts.files}
weights = load_weights(args.model)
sc = Scene(acts, weights, dev, min_level=args.min_level)
tl = Timeline(sc.P, sc.T)
cam = default_camera(tl)
return Renderer(sc, tl, cam, args.width, args.height, args.ss, dev)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--acts", default="acts.npz")
ap.add_argument("--model", default="model")
ap.add_argument("--width", type=int, default=3840)
ap.add_argument("--height", type=int, default=2160)
ap.add_argument("--ss", type=int, default=2)
ap.add_argument("--min-level", type=int, default=0)
ap.add_argument("--stills", type=str, default=None, help="comma list of times (s) -> PNGs")
ap.add_argument("--prefix", default="still")
ap.add_argument("--start", type=int, default=0)
ap.add_argument("--end", type=int, default=int(DUR * FPS))
ap.add_argument("--out", default="out.mp4")
ap.add_argument("--ffmpeg", default="ffmpeg")
ap.add_argument("--crf", type=int, default=15)
ap.add_argument("--preset", default="slow")
ap.add_argument("--threads", type=int, default=0)
ap.add_argument("--codec", default="hevc", choices=["hevc", "h264"])
ap.add_argument("--bits", type=int, default=16)
ap.add_argument("--device", default=None)
args = ap.parse_args()
R = build(args)
if args.stills:
from PIL import Image
R.bits = 8
for tt in [float(x) for x in args.stills.split(",")]:
t0 = time.time()
img = R.render(int(round(tt * FPS))).cpu().numpy()
fn = f"{args.prefix}_{tt:06.2f}.png"
Image.fromarray(img).save(fn)
print(f"[still] t={tt} -> {fn} ({time.time()-t0:.2f}s)", flush=True)
return
R.bits = args.bits
pix_in = "rgb48le" if args.bits == 16 else "rgb24"
cmd = [args.ffmpeg, "-y", "-hide_banner", "-loglevel", "error", "-f", "rawvideo", "-pix_fmt", pix_in,
"-s", f"{args.width}x{args.height}", "-r", str(FPS), "-i", "-"]
if args.codec == "hevc":
xp = "keyint=120:min-keyint=60:log-level=error" + (f":pools={args.threads}" if args.threads else "")
cmd += ["-c:v", "libx265", "-preset", args.preset, "-crf", str(args.crf), "-pix_fmt", "yuv420p10le", "-tag:v", "hvc1",
"-x265-params", xp]
else:
cmd += ["-c:v", "libx264", "-preset", args.preset, "-crf", str(args.crf), "-pix_fmt", "yuv420p",
"-profile:v", "high", "-level", "5.2", "-x264-params", "keyint=120:min-keyint=60"]
cmd += ["-color_primaries", "bt709", "-color_trc", "bt709", "-colorspace", "bt709", "-movflags", "+faststart"]
if args.threads and args.codec != "hevc":
cmd += ["-threads", str(args.threads)]
cmd += [args.out]
print("[main] ffmpeg:", " ".join(cmd), flush=True)
proc = subprocess.Popen(cmd, stdin=subprocess.PIPE)
t0 = time.time()
for fr in range(args.start, args.end):
img = R.render(fr)
arr = img.cpu().numpy()
proc.stdin.write((arr.astype("<u2") if args.bits == 16 else arr).tobytes())
if (fr - args.start) % 30 == 0:
el = time.time() - t0; done = fr - args.start + 1
print(f"[render] frame {fr} ({done}/{args.end-args.start}) {el:.0f}s {el/done:.2f}s/frame", flush=True)
proc.stdin.close(); proc.wait()
print(f"[main] done {args.out} in {time.time()-t0:.0f}s rc={proc.returncode}", flush=True)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
61.8 kB
·
Xet hash:
114c3a12131cdd5b9659403a7e39880dab1b5d9aca3d36a9b4f273ffa3e08d37

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.