File size: 3,572 Bytes
aa7bd25
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""
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)