File size: 11,812 Bytes
a1dd5ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
"""FDNN-V2: the context head. Same speed class, different training signal.

V1's measured dead end: distilled to an appearance teacher whose own space
separates close/open at 0.60, no student can learn verbs β€” the target lacks
them. V2 keeps V1's skeleton (stem -> FDNN recurrent cell, 0.27 ms/frame)
and warm-starts from its weights; what changes is WHERE the gradient comes
from (docs/FDNNV2_PLAN.md, all research-grounded):

  L1 PREDICTIVE (V-JEPA-style, arXiv 2506.09985): from state h_t, predict
     the embedding of frame t+k at two horizons. The future is free
     supervision on every frame, and dynamics must enter h to predict it.
     This is the biological claim made mechanical: understanding = a scene
     model good enough to predict what happens next.
  L2 ARROW OF TIME (Wei et al., CVPR 2018): classify forward vs reversed
     clips from the final state. Open and close are time-reversals; one
     binary head separates them by construction.
  L3 VERB-FOCUSED CONTRASTIVE (arXiv 2304.06708): sigmoid contrastive
     between clip context embeddings and 7B captions, with hard negatives
     built by verb/direction swaps β€” the joint space queries actually live
     in. A 2-layer text adapter maps SigLIP text embeddings into the
     256-d context space (~0.1 ms per query).

The appearance head keeps its V1 distillation target so every existing
index and the SigLIP text tower continue to work unchanged.

THE GATE IS AN EVENT DETECTOR. The cell's update gate z_t measures how much
of the scene model this frame is allowed to overwrite. Sustained high gate
activity marks an event boundary β€” so the encoder emits event segmentation
for free, and event embeddings are pooled with NOVELTY WEIGHTS (gate
activity), not uniformly: the moment the drawer closes outweighs the seconds
it sat still (the TempMe lesson, arXiv 2409.01156, applied at pooling time).
"""
from __future__ import annotations

import json
import re
from pathlib import Path

import mlx.core as mx
import mlx.nn as nn
import numpy as np

from .fdnnvideo import EMBED_DIM, FDNNVideoEncoder

CTX_DIM = 256

# Verb/direction swaps for L3 hard negatives: shared lexicon (moved to
# lexicon.py so the QUERY path can import it without this module's mlx)
from .lexicon import derived_swaps  # noqa: E402,F401


def swap_verbs(text: str, rng, store=None) -> str | None:
    """One randomly chosen applicable swap -> a hard negative. None if no
    swap applies (caption has no directional content to invert)."""
    # Training negatives obey the no-hardwire rule too: a model taught
    # from hand-written oppositions has the hand-writing baked into its
    # weights, which is worse than a lookup table because it cannot be
    # grepped out afterwards.
    pairs = list(derived_swaps(store)) if store is not None else []
    if not pairs:
        return None
    t = " " + text.lower() + " "
    hits = []
    for a, b in pairs:
        if f" {a} " in t:
            hits.append((a, b))
        if f" {b} " in t:
            hits.append((b, a))
    if not hits:
        # clause order flip is the fallback inversion for compounds
        parts = re.split(r"\band then\b|\bthen\b|,", text)
        if len(parts) >= 2:
            return " then ".join(p.strip() for p in reversed(parts)
                                 if p.strip())
        return None
    a, b = hits[rng.integers(len(hits))]
    return re.sub(rf"\b{re.escape(a)}\b", b, text.lower(), count=1)


class TextAdapter(nn.Module):
    """SigLIP text embedding (1152) -> context space (256). Two layers,
    ~0.7M params, ~0.1 ms β€” the entire query-time cost of verb awareness."""

    def __init__(self, in_dim=EMBED_DIM, dim=CTX_DIM):
        super().__init__()
        self.fc1 = nn.Linear(in_dim, 512)
        self.fc2 = nn.Linear(512, dim)

    def __call__(self, x):
        e = self.fc2(nn.silu(self.fc1(x)))
        return e * mx.rsqrt(mx.sum(e * e, axis=-1, keepdims=True) + 1e-8)


class FDNNv2(nn.Module):
    """V1 skeleton + context head + predictors + arrow-of-time head."""

    def __init__(self, base: FDNNVideoEncoder, ctx_dim=CTX_DIM,
                 horizons=(2, 10), motion_dim=64):
        super().__init__()
        self.base = base
        C = base.channels
        G = base.cfg["glimpse"]
        self.horizons = list(horizons)
        self.motion_dim = motion_dim
        # MOTION PATHWAY β€” the reversal-symmetry breaker. The gated cell is a
        # leaky integrator; for smooth inputs an EMA is nearly order-invariant,
        # and stage-A measured exactly that: AoT stuck at chance (0.47) while
        # dominating the loss. Frame-difference features flip SIGN under time
        # reversal, so direction becomes linearly readable. Velocity is not a
        # nicety here; it is the only anti-symmetric signal in the model.
        self.motion = nn.Linear(G, motion_dim)
        # context head reads (glimpse, state, motion)
        self.head_ctx = nn.Linear(G + C + motion_dim, ctx_dim)
        # Predictors target the model's OWN future glimpse-change (stop-grad).
        # A2 measured why not student-embedding changes: adjacent-frame true
        # change is ~0.2 in norm while the student's own error is ~0.45 β€” the
        # difference of noisy embeddings is noise, and the predictors sat at
        # cosine 0 for 8 epochs. Own-latent prediction is JEPA's actual
        # recipe; the appearance anchor stops the stem collapsing to constants.
        self.pred = [nn.Linear(C + motion_dim, G) for _ in horizons]
        # AoT reads the velocity SEQUENCE through a temporal conv, not the
        # mean: this corpus is reciprocal motion (arm out, arm back), so mean
        # signed velocity cancels β€” measured at chance twice. Order within
        # the window is the signal; a conv kernel can be asymmetric in time.
        self.aot_conv = nn.Conv1d(motion_dim, 32, 5, padding=2)
        self.head_aot = nn.Linear(32, 1)
        self.cfg = {"ctx_dim": ctx_dim, "horizons": self.horizons,
                    "motion_dim": motion_dim, "base": base.cfg}

    # ---- forward over a sequence, exposing everything the losses and the
    # event segmenter need ---------------------------------------------------
    def run(self, seq, h0=None):
        """(B,T,H,W,3) -> dict of app (B,T,1152), ctx (B,T,ctx), preds,
        gates (B,T), h_last."""
        b = self.base
        B, T = seq.shape[0], seq.shape[1]
        g = b.stem(seq.reshape(B * T, *seq.shape[2:])).reshape(B, T, -1)
        h = b.init_state(B) if h0 is None else h0
        hs, gates = [], []
        cell = b.cell
        for t in range(T):
            z = mx.sigmoid(g[:, t] @ cell.Wz + h @ cell.Uz + cell.bz)
            h = cell(g[:, t], h)
            hs.append(h)
            gates.append(mx.mean(z, axis=-1))
        hs = mx.stack(hs, axis=1)                    # (B,T,C)
        gates = mx.stack(gates, axis=1)              # (B,T)
        app = b.head_g(g) + b.head_h(hs)
        app = app * mx.rsqrt(mx.sum(app * app, axis=-1, keepdims=True) + 1e-8)
        # signed velocity of the glimpse; first frame gets zero motion.
        # NORMALISED: measured, ||dg|| is real motion (corr 0.63 with pixel
        # motion, 5x on moving frames) but only ~2% of the feature norm β€” fed
        # raw, it drowned next to signals 50x larger and every motion head
        # starved. Direction is unit-normalised; magnitude re-enters as a
        # bounded gain, so both the WHAT and the HOW-MUCH of motion survive.
        dg = mx.concatenate([mx.zeros_like(g[:, :1]),
                             g[:, 1:] - g[:, :-1]], axis=1)
        mag = mx.sqrt(mx.sum(dg * dg, axis=-1, keepdims=True) + 1e-8)
        mfeat = nn.silu(self.motion(dg / mag)) * mx.tanh(mag)
        cat = mx.concatenate([g, hs, mfeat], axis=-1)
        ctx = self.head_ctx(cat)
        ctx = ctx * mx.rsqrt(mx.sum(ctx * ctx, axis=-1, keepdims=True) + 1e-8)
        hm = mx.concatenate([hs, mfeat], axis=-1)
        preds = [p(hm) for p in self.pred]           # each (B,T,G)
        a = nn.silu(self.aot_conv(mfeat))            # (B,T,32) time-conv
        aot = self.head_aot(mx.mean(a, axis=1))[:, 0]
        return {"app": app, "ctx": ctx, "preds": preds, "gates": gates,
                "aot": aot, "h": h, "g": g}

    # ---- streaming embed (the write path): app + ctx + gate per frame -----
    def embed_stream_np(self, frames_u8, batch=128):
        h = self.base.init_state(1)
        apps, ctxs, gates = [], [], []
        for i in range(0, len(frames_u8), batch):
            x = mx.array(frames_u8[i:i + batch].astype(np.float32)
                         / 127.5 - 1.0)[None]
            out = self.run(x, h0=h)
            h = out["h"]
            apps.append(np.array(out["app"][0], np.float32))
            ctxs.append(np.array(out["ctx"][0], np.float32))
            gates.append(np.array(out["gates"][0], np.float32))
        return (np.concatenate(apps), np.concatenate(ctxs),
                np.concatenate(gates))


# ===========================================================================
# Event segmentation from the gate signal + novelty-weighted pooling
# ===========================================================================
def segment_events(gates, ts, min_len=6, smooth=5, thresh_pct=75.0):
    """Gate-activity peaks -> event boundaries.

    A boundary is where smoothed gate activity crosses above its own
    percentile threshold after having been below β€” the scene model is being
    rewritten. Percentile (not absolute) because gate scale is a trained
    quantity; per-stream calibration is free.
    Returns list of (start_idx, end_idx) covering the stream.
    """
    g = np.convolve(gates, np.ones(smooth) / smooth, mode="same")
    thr = np.percentile(g, thresh_pct)
    above = g > thr
    bounds = [0]
    for i in range(1, len(g)):
        if above[i] and not above[i - 1] and i - bounds[-1] >= min_len:
            bounds.append(i)
    bounds.append(len(g))
    return [(a, b) for a, b in zip(bounds[:-1], bounds[1:]) if b - a >= 2]


def pool_event(vecs, gates, lo, hi):
    """Novelty-weighted pool: frames weighted by gate activity, so change
    dominates stillness. Uniform mean is the verb-eraser; this is not."""
    w = gates[lo:hi] + 1e-3
    w = w / w.sum()
    v = (vecs[lo:hi] * w[:, None]).sum(0)
    return v / (np.linalg.norm(v) + 1e-8)


# ===========================================================================
# persistence
# ===========================================================================
def save_v2(model, adapter, meta, path):
    from mlx.utils import tree_flatten
    path = Path(path)
    path.mkdir(parents=True, exist_ok=True)
    np.savez(path / "v2.npz",
             **{k: np.array(v) for k, v in tree_flatten(model.parameters())})
    np.savez(path / "adapter.npz",
             **{k: np.array(v)
                for k, v in tree_flatten(adapter.parameters())})
    (path / "v2.json").write_text(json.dumps({**meta, "cfg": model.cfg},
                                             indent=2))


def load_v2(path):
    from mlx.utils import tree_unflatten
    path = Path(path)
    meta = json.loads((path / "v2.json").read_text())
    base = FDNNVideoEncoder(**meta["cfg"]["base"])
    model = FDNNv2(base, ctx_dim=meta["cfg"]["ctx_dim"],
                   horizons=tuple(meta["cfg"]["horizons"]),
                   motion_dim=meta["cfg"].get("motion_dim", 64))
    z = np.load(path / "v2.npz")
    model.update(tree_unflatten([(k, mx.array(z[k])) for k in z.files]))
    model.base.cell.freeze(keys=["basis_types", "mask"], recurse=False)
    adapter = TextAdapter()
    z = np.load(path / "adapter.npz")
    adapter.update(tree_unflatten([(k, mx.array(z[k])) for k in z.files]))
    mx.eval(model.parameters(), adapter.parameters())
    return model, adapter, meta