""" visualize.py -- live visualizer for Mycel-LM: watch the fungal colony grow in 3D as the model generates, token by token. Captures per generated token (via instrument.Recorder, which MycelBlock writes to): * tip POSITIONS (all pos_dim axes) for one layer's colony -- the 3D scatter * per-tip growth TRAIL over the growth steps -- the hyphae extending this token * per-layer colony density + trait-station sites + attention + token stream Renders one self-contained HTML dashboard (data embedded -> opens via file://). The colony is a Three.js network you can drag-rotate / scroll-zoom: hyphal tips linked to their nearest neighbours as filaments (so it reads as a mycelial web, not loose dots), tips coloured by local density, with trait STATIONS as orange wire-spheres and faint grey growth trails. Three.js loads from a CDN (needs network the first time). Usage: python visualize.py --prompt "the mycelium spreads" --tokens 50 """ import argparse, json, os, sys, webbrowser import torch from model import QuazimotoLM, QuazimotoConfig import instrument PKG_DIR = os.path.dirname(os.path.abspath(__file__)) def find_ckpt(path): if path and os.path.isfile(path): return path folder = (os.path.dirname(path) if path else os.path.join(PKG_DIR, "chkpt")) or "." if not os.path.isdir(folder): return None pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")] return max(pts, key=os.path.getmtime) if pts else None def load_tokenizer(tok_dir): sys.path.insert(0, tok_dir) from spike_tokenizer import SpikeTokenizer return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json")) @torch.no_grad() def run_capture(model, cfg, tok, ids, n_tokens, temperature, top_k, device): rec = instrument.Recorder(phase_layer=cfg.n_layer // 2, attn_layer=cfg.n_layer // 2) instrument.set_rec(rec) idx = torch.tensor([ids], device=device) try: for _ in range(n_tokens): cond = idx[:, -cfg.block_size:] rec.begin() logits, _, _ = model(cond) lg = torch.nan_to_num(logits[:, -1, :].float(), nan=0.0, posinf=1e4, neginf=-1e4) / max(temperature, 1e-6) if top_k: v = torch.topk(lg, min(top_k, lg.size(-1))).values lg = lg.masked_fill(lg < v[:, [-1]], float("-inf")) probs = torch.softmax(lg, dim=-1) nxt = (torch.argmax(lg, -1, keepdim=True) if temperature <= 1e-3 else torch.multinomial(probs, 1)) top = torch.topk(torch.softmax(logits[:, -1].float(), -1), 5) rec.end(token=int(nxt), char=tok.decode([int(nxt)], skip_special_tokens=False), top=[[int(t), round(float(p), 3)] for t, p in zip(top.indices[0], top.values[0])]) idx = torch.cat([idx, nxt], dim=1) finally: instrument.set_rec(None) return rec.frames, idx[0].tolist() def main(): p = argparse.ArgumentParser(description="Mycel-LM 3D colony visualizer") p.add_argument("--ckpt", default="") p.add_argument("--tok_dir", default=PKG_DIR) p.add_argument("--prompt", default="the mycelium spreads") p.add_argument("--tokens", type=int, default=50) p.add_argument("--temperature", type=float, default=0.0) p.add_argument("--top_k", type=int, default=40) p.add_argument("--device", default="cpu") p.add_argument("--out", default=os.path.join(PKG_DIR, "viz.html")) p.add_argument("--no_open", action="store_true") args = p.parse_args() path = find_ckpt(args.ckpt) if path is None: print("No checkpoint found; pass --ckpt or train one."); return ckpt = torch.load(path, map_location=args.device, weights_only=False) cfg = QuazimotoConfig(**ckpt["family_config"]) model = QuazimotoLM(cfg); model.load_state_dict(ckpt["model"], strict=False) model.to(args.device).eval() tok = load_tokenizer(args.tok_dir) print(f"loaded {path} (step {ckpt.get('step')}) | capturing {args.tokens} tokens ...") ids = tok.encode(args.prompt, add_special_tokens=False) frames, _ = run_capture(model, cfg, tok, ids, args.tokens, args.temperature, args.top_k or None, args.device) pl = cfg.n_layer // 2 blk = model.layers[pl].quaz stations = blk.stations.anchors.detach().cpu().tolist() if getattr(blk, "use_stations", False) else [] data = { "meta": {"ckpt": os.path.basename(path), "step": ckpt.get("step"), "n_layer": cfg.n_layer, "n_tips": cfg.n_tips, "pos_dim": cfg.mycel_pos_dim, "bound": cfg.osc_bound, "phase_layer": pl, "prompt": args.prompt, "stations": stations, "tip_rings": bool(getattr(cfg, "use_tip_rings", False)), "tip_ring_size": getattr(cfg, "tip_ring_size", 0)}, "prompt_chars": [tok.decode([t], skip_special_tokens=False) for t in ids], "frames": frames, } html = HTML_TEMPLATE.replace("/*__DATA__*/", json.dumps(data)) with open(args.out, "w", encoding="utf-8") as f: f.write(html) print(f"wrote {args.out} ({os.path.getsize(args.out)//1024} KB)") if not args.no_open: webbrowser.open("file://" + os.path.abspath(args.out)) HTML_TEMPLATE = r""" Mycel-LM live

Mycel-LM · colony 3D

Token stream

Trait activity (this token)

Top predictions

Colony — layer (drag rotate, scroll zoom; wire-spheres = trait stations)

Colony density per layer (bright = dense/clustered)

Attention — layer (last token → context)

Notes

""" if __name__ == "__main__": main()