Chimera-64M / instrument.py
Quazim0t0's picture
Upload instrument.py with huggingface_hub
aa7bd25 verified
Raw History Blame Contribute Delete
3.57 kB
"""
instrument.py -- zero-overhead-when-off capture hooks for the live visualizer.
A single global Recorder is consulted by the model's hot paths. When no recorder
is active (the default), every call is a single `is None` check, so training and
normal generation are unaffected. visualize.py activates a recorder, runs a
forward per generated token, and reads back per-token state.
"""
_REC = None
def get_rec():
return _REC
def set_rec(r):
global _REC
_REC = r
def _r(x, nd=4):
return round(float(x), nd)
class Recorder:
"""Collects one frame per generated token. Modules append in execution order,
so list position == layer index (attention then quaz, per layer)."""
def __init__(self, phase_layer=None, attn_layer=None):
self.frames = []
self.cur = None
self.enabled = False
self.phase_layer = phase_layer # which layer's individual phases to keep
self.attn_layer = attn_layer # which layer's attention map to keep
self._spec_tmp = []
# ---- frame lifecycle (driven by visualize.py) ----
def begin(self):
self.cur = {"rings": [], "phases": None, "attn": None,
"spec": [], "quaz_norm": [], "traits": {}}
self._spec_tmp = []
self.enabled = True
def end(self, **meta):
self.enabled = False
if self.cur is not None:
self.cur.update(meta)
self.frames.append(self.cur)
self.cur = None
# ---- module callbacks (guarded by `enabled` at the call site) ----
def log_ring(self, R, psi, phases):
li = len(self.cur["rings"])
self.cur["rings"].append({"R": [_r(x, 4) for x in R],
"psi": [_r(x, 4) for x in psi]})
if phases is not None and (self.phase_layer is None or li == self.phase_layer):
self.cur["phases"] = {"layer": li, "theta": [_r(x, 3) for x in phases]}
def log_tip_phase(self, psi, r):
"""Attach per-tip ring phase psi and coherence r to THIS layer's colony frame
(only if it is the kept phase_layer). Lets the 3D view colour tips by ring
phase so synchronization waves are visible."""
ph = self.cur.get("phases")
if ph is not None and ph.get("layer") == len(self.cur["rings"]) - 1:
ph["tip_psi"] = [_r(x, 3) for x in psi]
ph["tip_r"] = [_r(x, 3) for x in r]
def log_traj(self, traj):
"""Attach the tip growth trajectory (list over growth steps, each a flattened
[N*pos_dim] snapshot) to THIS layer's colony frame (phase_layer only). Lets the
3D view draw each tip's hypha extending across the growth steps."""
ph = self.cur.get("phases")
if ph is not None and ph.get("layer") == len(self.cur["rings"]) - 1:
ph["traj"] = [[_r(x, 3) for x in step] for step in traj]
def log_attn(self, layer_idx, w):
if self.attn_layer is None or layer_idx == self.attn_layer:
self.cur["attn"] = {"layer": layer_idx, "w": [_r(x, 4) for x in w]}
def log_quaz_norm(self, n):
self.cur["quaz_norm"].append(_r(n, 4))
def push_spec(self, route): # one ring's [n_spec] route weights
self._spec_tmp.append([_r(x, 4) for x in route])
def flush_spec(self): # called once per quaz block (a layer)
if self._spec_tmp:
self.cur["spec"].append(self._spec_tmp)
self._spec_tmp = []
def log_trait(self, name, val):
self.cur["traits"][name] = _r(val, 4)