File size: 14,225 Bytes
47d8ad6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
"""DenseTrackStep: fixed-capacity, dict-free rewrite of SAM3 multiplex track_step
(propagation-only mode) around the REAL submodules, for torch.export / CoreML.

Replaces:
  - output_dict["cond_frame_outputs"/"non_cond_frame_outputs"] dict-walk in
    _prepare_memory_conditioned_features (video_tracking_multiplex.py:1289)
    -> fixed rolling tensor banks (register_buffer) + static concat order
  - MultiplexState add/remove dynamics -> frozen state (1 bucket x multiplex_count);
    its mux/demux matrices become trace-time constants
  - Eager's variable-length skip of missing warmup slots -> fixed banks +
    additive -inf key mask from mem_valid / ptr_valid state buffers

Reuses UNCHANGED: transformer.encoder (memory fusion), _forward_sam_heads
(propagation path incl. MultiplexMaskDecoder), _encode_new_memory
(SimpleMaskEncoder), all with real weights and real dims
(hidden_dim=256, num_maskmem=7, 72x72 memory grid, 1008px masks).

Scope / documented approximations:
  - Exactly 1 cond frame (the init frame). kino's use-case prompts once.
  - Warmup handled via validity masks (not per-count model variants).
  - Cond maskmem tpos follows use_maskmem_tpos_v2 from frame_pos (idx 6 only
    once frame_idx >= num_maskmem).
  - max_objects = multiplex_count = 16 (released SAM3.1 ckpt).
"""

import torch
import torch.nn as nn

import common

E = 72          # sam_image_embedding_size
HW = E * E      # 5184
C = 256
NUM_MASKMEM = 7           # 1 cond + 6 rolling non-cond
N_NONCOND = NUM_MASKMEM - 1
MAX_PTRS = 16             # 1 cond + 15 rolling non-cond
N_PTR = MAX_PTRS - 1
OBJ = common.MULTIPLEX_COUNT  # 16
NUM_OBJ_PTR_TOKENS = MAX_PTRS * OBJ  # 256
# Large finite penalty — CoreML/fp16 friendly vs true -inf
# Softmax key-pad penalty. Keep modest: -1e4 overflowed coremltools' numpy
# attention probe during convert and ballooned RAM to ~60GB.
MASK_NEG = -80.0


class DenseTrackStep(nn.Module):
    def __init__(self, model, ms, pos72):
        super().__init__()
        self.model = model
        self.ms = ms  # frozen MultiplexState: matrices are trace-time constants
        self.multimask = model._use_multimask(False, None)
        assert self.multimask, "propagation multimask expected True"

        # ---- rolling state banks (newest last) ----
        self.register_buffer("mem_bank", torch.zeros(N_NONCOND, 1, C, E, E))
        self.register_buffer("img_bank", torch.zeros(N_NONCOND, HW, 1, C))
        self.register_buffer("ptr_bank", torch.zeros(N_PTR, 1, OBJ, C))
        # 1 = slot filled (attend); 0 = warmup pad (masked)
        self.register_buffer("mem_valid", torch.zeros(N_NONCOND))
        self.register_buffer("ptr_valid", torch.zeros(N_PTR))

        # ---- constants ----
        with torch.no_grad():
            mm_pos = model.maskmem_backbone.position_encoding(
                torch.zeros(1, C, E, E)
            ).to(torch.float32)                      # (1,C,E,E)
        mm_pos = mm_pos.flatten(2).permute(2, 0, 1)  # (HW,1,C)
        tpos = model.maskmem_tpos_enc.detach()       # (7,1,1,C)
        # Non-cond bank: slot s -> t_pos=s+1 -> tpos idx 5-s (fixed).
        # Cond tpos is dynamic from frame_pos (see _cond_tpos).
        noncond_mem_pos = []
        noncond_img_pos = []
        for s in range(N_NONCOND):
            idx = N_NONCOND - 1 - s  # 5..0
            noncond_mem_pos.append(mm_pos + tpos[idx])
            noncond_img_pos.append(pos72 + tpos[idx])
        self.register_buffer("mm_pos", mm_pos.contiguous())
        self.register_buffer("tpos_enc", tpos.contiguous())
        self.register_buffer(
            "noncond_mem_pos", torch.cat(noncond_mem_pos, dim=0).contiguous()
        )
        self.register_buffer(
            "noncond_img_pos", torch.cat(noncond_img_pos, dim=0).contiguous()
        )
        self.register_buffer("pos72", pos72.clone().contiguous())
        self.register_buffer("ptr_tdiffs", torch.arange(1, MAX_PTRS).float())
        self.register_buffer("_one", torch.ones(1))
        self.register_buffer(
            "_arange7", torch.arange(NUM_MASKMEM, dtype=torch.float32)
        )
        self.register_buffer("_oor_idx", torch.tensor(float(NUM_MASKMEM - 1)))

    def _obj_pos(self, frame_pos):
        """Replicates _get_tpos_enc([frame_idx, 1..15], max_abs_pos=16)."""
        from sam3.model.sam3_tracker_utils import get_1d_sine_pe

        pos = torch.cat([frame_pos.reshape(1), self.ptr_tdiffs]) / (MAX_PTRS - 1)
        pe = get_1d_sine_pe(pos, dim=C)                    # (16,C)
        pe = self.model.obj_ptr_tpos_proj(pe)              # (16,C)
        return pe.unsqueeze(1).repeat_interleave(OBJ, dim=0)  # (256,1,C)

    def _cond_tpos(self, frame_pos):
        """use_maskmem_tpos_v2 embedding for cond frame at t=0 (t_pos=frame_idx).

        Soft one-hot over 7 tpos rows — avoids data-dependent integer index
        that torch.export rejects.
        """
        fp = frame_pos.reshape(()).to(self.tpos_enc.dtype)
        in_range = (NUM_MASKMEM - 1) - torch.clamp(fp, min=0, max=float(N_NONCOND))
        sel = torch.where(fp >= float(NUM_MASKMEM), self._oor_idx, in_range)
        # Exact 1 at integer sel, 0 elsewhere (frame_pos is integral in practice)
        w = (1.0 - (self._arange7 - sel).abs()).clamp(min=0.0)
        return (w.view(NUM_MASKMEM, 1, 1, 1) * self.tpos_enc).sum(dim=0)  # (1,1,C)

    def _mem_img_pos_cat(self, frame_pos):
        """Cond spatial+tpos (dynamic) then fixed non-cond bank positions."""
        cond_t = self._cond_tpos(frame_pos)  # (1,1,C)
        cond_mem_pos = self.mm_pos + cond_t
        cond_img_pos = self.pos72 + cond_t
        return (
            torch.cat([cond_mem_pos, self.noncond_mem_pos], dim=0),
            torch.cat([cond_img_pos, self.noncond_img_pos], dim=0),
        )

    def _memory_key_attn_mask(self):
        """Additive SDPA mask (1,1,1,K): 0 attend, MASK_NEG block.

        Key layout matches prompt cat: cond_mem(HW) + mem_bank(N_NONCOND*HW)
        + cond_ptr(OBJ) + flipped ptr_bank(N_PTR*OBJ).
        """
        # mem: cond always valid, then per-slot validity expanded over HW
        mem_slot = torch.cat([self._one, self.mem_valid], dim=0)  # (7,)
        mem_tok = mem_slot.repeat_interleave(HW)  # (7*HW,)
        # ptr: cond always valid; bank is newest-last, prompt uses newest-first
        ptr_slot = torch.cat([self._one, self.ptr_valid.flip(0)], dim=0)  # (16,)
        ptr_tok = ptr_slot.repeat_interleave(OBJ)  # (16*OBJ,)
        valid = torch.cat([mem_tok, ptr_tok], dim=0)  # (K,)
        # (1,1,1,K) broadcasts over batch/heads/queries
        return (MASK_NEG * (1.0 - valid)).view(1, 1, 1, -1)

    def forward(self, vis72, hires0, hires1, cond_mem, cond_img, cond_ptr,
                frame_pos):
        m = self.model

        # ---- dense _prepare_memory_conditioned_features ----
        mem_noncond = (
            self.mem_bank.squeeze(1).flatten(2).permute(0, 2, 1)
            .reshape(N_NONCOND * HW, 1, C)
        )
        mem_tokens = torch.cat(
            [cond_mem.flatten(2).permute(2, 0, 1), mem_noncond], dim=0
        )                                                     # (7*HW,1,C)
        img_tokens = torch.cat(
            [cond_img, self.img_bank.reshape(N_NONCOND * HW, 1, C)], dim=0
        )                                                     # (7*HW,1,C)
        # ptr order: cond, then newest..oldest (t_diff 1..15)
        ptr_noncond = self.ptr_bank.flip(0).squeeze(1).reshape(N_PTR * OBJ, C)
        obj_ptrs = torch.cat(
            [cond_ptr.squeeze(0), ptr_noncond], dim=0
        ).unsqueeze(1)                                        # (256,1,C)

        prompt = torch.cat([mem_tokens, obj_ptrs], dim=0)
        mem_pos_cat, img_pos_cat = self._mem_img_pos_cat(frame_pos)
        prompt_pos = torch.cat([mem_pos_cat, self._obj_pos(frame_pos)], 0)
        key_mask = self._memory_key_attn_mask()

        enc_out = m.transformer.encoder(
            image=vis72,
            src=vis72,
            memory_image=img_tokens,
            memory=prompt,
            image_pos=self.pos72,
            src_pos=self.pos72,
            memory_image_pos=img_pos_cat,
            memory_pos=prompt_pos,
            num_obj_ptr_tokens=NUM_OBJ_PTR_TOKENS,
            memory_key_attn_mask=key_mask,
        )
        pix_feat_with_mem = enc_out["memory"].permute(1, 2, 0).view(1, C, E, E)

        # ---- REAL sam heads (propagation path) ----
        sam_out = m._forward_sam_heads(
            backbone_features=pix_feat_with_mem,
            propagation_high_res_features=[hires0, hires1],
            multimask_output=self.multimask,
            objects_to_interact=list(range(OBJ)),
            multiplex_state=self.ms,
        )
        low_res_masks = sam_out["low_res_masks"]
        high_res_masks = sam_out["high_res_masks"]
        object_score_logits = sam_out["object_score_logits"]
        ious = sam_out["ious"]
        obj_ptr = sam_out["obj_ptr"]

        # ---- REAL memory encoder ----
        maskmem_features, _ = m._encode_new_memory(
            image=None,
            current_vision_feats=[vis72],
            feat_sizes=[(E, E)],
            pred_masks_high_res=high_res_masks,
            object_score_logits=object_score_logits,
            is_mask_from_pts=False,
            conditioning_objects=set(),
            multiplex_state=self.ms,
        )

        # ---- state update: shift + append (static shapes) ----
        self.mem_bank.copy_(
            torch.cat([self.mem_bank[1:], maskmem_features.unsqueeze(0)], 0)
        )
        self.img_bank.copy_(
            torch.cat([self.img_bank[1:], vis72.unsqueeze(0)], 0)
        )
        mux_ptr = self.ms.mux(obj_ptr)  # (1,OBJ,C)
        self.ptr_bank.copy_(
            torch.cat([self.ptr_bank[1:], mux_ptr.unsqueeze(0)], 0)
        )
        self.mem_valid.copy_(torch.cat([self.mem_valid[1:], self._one], 0))
        self.ptr_valid.copy_(torch.cat([self.ptr_valid[1:], self._one], 0))

        return low_res_masks, high_res_masks, object_score_logits, ious


def build_wrapper():
    model = common.build_model()
    pos72 = common.pos_embed_72(model)
    ms = common.make_multiplex_state(model)
    return DenseTrackStep(model, ms, pos72).eval(), model


def seed_banks(wrapper, cache):
    """Fill banks from eager frames so state matches eager just before TARGET."""
    tgt = cache["target_frame"]
    frames = cache["frames"]
    for s in range(N_NONCOND):  # slot s <- frame tgt-6+s
        f = tgt - N_NONCOND + s
        wrapper.mem_bank[s] = frames[f]["maskmem_features"]
        wrapper.img_bank[s] = frames[f]["image_features"]
        wrapper.mem_valid[s] = 1.0
    for p in range(N_PTR):      # slot p <- frame tgt-15+p
        f = tgt - N_PTR + p
        wrapper.ptr_bank[p] = frames[f]["obj_ptr"]
        wrapper.ptr_valid[p] = 1.0


def seed_banks_for_frame(wrapper, cache, frame_idx):
    """Seed banks to match eager state *entering* frame_idx (warmup-aware)."""
    wrapper.mem_bank.zero_()
    wrapper.img_bank.zero_()
    wrapper.ptr_bank.zero_()
    wrapper.mem_valid.zero_()
    wrapper.ptr_valid.zero_()
    frames = cache["frames"]
    # mem: last up to N_NONCOND non-cond frames before frame_idx
    mem_src = [f for f in range(max(1, frame_idx - N_NONCOND), frame_idx)]
    for i, f in enumerate(mem_src):
        slot = N_NONCOND - len(mem_src) + i
        wrapper.mem_bank[slot] = frames[f]["maskmem_features"]
        wrapper.img_bank[slot] = frames[f]["image_features"]
        wrapper.mem_valid[slot] = 1.0
    ptr_src = [f for f in range(max(1, frame_idx - N_PTR), frame_idx)]
    for i, f in enumerate(ptr_src):
        slot = N_PTR - len(ptr_src) + i
        wrapper.ptr_bank[slot] = frames[f]["obj_ptr"]
        wrapper.ptr_valid[slot] = 1.0


def frame_inputs(model, cache, frame_idx):
    feats, _ = common.synth_frame_features(model, frame_idx)
    vis72 = feats[2].contiguous()
    hires0 = feats[0].permute(1, 2, 0).view(1, 32, 4 * E, 4 * E).contiguous()
    hires1 = feats[1].permute(1, 2, 0).view(1, 64, 2 * E, 2 * E).contiguous()
    cond = cache["cond"]
    return (vis72, hires0, hires1,
            cond["maskmem_features"].contiguous(),
            cond["image_features"].contiguous(),
            cond["obj_ptr"].contiguous(),
            torch.tensor([float(frame_idx)]))


def main():
    cache = torch.load("eager_cache.pt", weights_only=False)
    wrapper, model = build_wrapper()

    def report(name, a, b):
        d = (a - b).abs().max().item()
        r = b.abs().max().item()
        print(f"  {name:24s} max|diff|={d:.3e}  (ref max|val|={r:.3e})")
        return d

    # ---- steady-state parity ----
    seed_banks(wrapper, cache)
    tgt = cache["target_frame"]
    inputs = frame_inputs(model, cache, tgt)
    with torch.no_grad():
        low, high, osl, ious = wrapper(*inputs)
    t = cache["target"]
    print(f"parity vs eager track_step @ frame {tgt} (steady):")
    d1 = report("pred_masks(low)", low, t["pred_masks"])
    d2 = report("pred_masks_high_res", high, t["pred_masks_high_res"])
    d3 = report("object_score_logits", osl, t["object_score_logits"])
    d4 = report("ious", ious, t["ious"])
    d5 = report("new maskmem_features", wrapper.mem_bank[-1],
                t["maskmem_features"].squeeze(0))
    d6 = report("new obj_ptr (muxed)", wrapper.ptr_bank[-1], t["obj_ptr"])
    ok_steady = max(d1, d2, d3, d4, d5, d6) < 1e-3

    # ---- warmup parity (frame 3) ----
    warm = cache.get("warmup_frame", 3)
    if "warmup" in cache:
        seed_banks_for_frame(wrapper, cache, warm)
        with torch.no_grad():
            low, high, osl, ious = wrapper(*frame_inputs(model, cache, warm))
        w = cache["warmup"]
        print(f"parity vs eager track_step @ frame {warm} (warmup):")
        w1 = report("pred_masks(low)", low, w["pred_masks"])
        w2 = report("pred_masks_high_res", high, w["pred_masks_high_res"])
        w3 = report("object_score_logits", osl, w["object_score_logits"])
        w4 = report("ious", ious, w["ious"])
        ok_warm = max(w1, w2, w3, w4) < 1e-3
    else:
        print("warmup cache missing — skip warmup check")
        ok_warm = True

    print("PARITY:", "PASS" if (ok_steady and ok_warm) else "FAIL")


if __name__ == "__main__":
    main()