Buckets:
| #!/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.