File size: 5,533 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
"""Shared setup: CPU workarounds + model build for SAM3 multiplex video model."""

import sys
import os

SP = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(SP, "triton_stub"))  # satisfy `import triton`
sys.path.insert(0, os.path.join(SP, "sam3"))

import torch  # noqa: E402

# CPU-run workarounds: sam3 hardcodes .cuda()/.pin_memory() in eval paths
torch.Tensor.cuda = lambda self, *a, **kw: self
torch.Tensor.pin_memory = lambda self, *a, **kw: self

MULTIPLEX_COUNT = 16  # matches released SAM3.1 multiplex checkpoint


# Official facebook/sam3.1 is gated; open mirror hosts the same multiplex weights.
SAM31_HF_REPO = os.environ.get("SAM3_HF_REPO", "AEmotionStudio/sam3.1")
SAM31_CKPT_NAME = "sam3.1_multiplex.pt"


def resolve_checkpoint(checkpoint_path=None):
    """Return local path to sam3.1_multiplex.pt (download via HF if needed)."""
    if checkpoint_path:
        return checkpoint_path
    env = os.environ.get("SAM3_CHECKPOINT")
    if env:
        return env
    from huggingface_hub import hf_hub_download

    return hf_hub_download(repo_id=SAM31_HF_REPO, filename=SAM31_CKPT_NAME)


def hide_triton_stub():
    """Remove the triton stub after all sam3 modules are imported, so torch
    inductor's `import triton` fails cleanly into its `triton = None` path."""
    stub = os.path.join(SP, "triton_stub")
    while stub in sys.path:
        sys.path.remove(stub)
    for k in list(sys.modules):
        if k == "triton" or k.startswith("triton."):
            del sys.modules[k]


def _tracker_state_dict(ckpt):
    """Extract tracker weights from merged SAM3.1 multiplex checkpoint.

    HF / mirror file stores full predictor keys (`tracker.model.*` +
    `detector.*`). Tracker-only module wants the stripped `tracker.model.`
    prefix. Also accepts a bare tracker state_dict or `{"model": ...}`.
    """
    if isinstance(ckpt, dict) and "model" in ckpt and isinstance(ckpt["model"], dict):
        ckpt = ckpt["model"]
    if any(k.startswith("tracker.model.") for k in ckpt):
        return {
            k[len("tracker.model.") :]: v
            for k, v in ckpt.items()
            if k.startswith("tracker.model.")
        }
    return ckpt


def build_model(
    multiplex_count=MULTIPLEX_COUNT,
    checkpoint_path=None,
    load_checkpoint=True,
    seed=1234,
):
    """Build multiplex video tracker. Default: real SAM3.1 weights, mux=16.

    Set load_checkpoint=False for random-weight debug (uses seed).
    """
    from sam3.model_builder import build_sam3_multiplex_video_model

    if not load_checkpoint:
        torch.manual_seed(seed)

    # Always construct without builder-side load: merged HF ckpt needs key strip.
    model = build_sam3_multiplex_video_model(
        checkpoint_path=None,
        load_from_HF=False,
        multiplex_count=multiplex_count,
        use_fa3=False,
        use_rope_real=True,  # avoid complex tensors (CoreML cannot represent them)
        device="cpu",
        strict_state_dict_loading=False,
    )
    # Drop image backbone before load: track_step uses precomputed feats; HF
    # multiplex file also keeps backbone under detector.*, not tracker.model.*.
    model.backbone = None

    if load_checkpoint:
        ckpt_path = resolve_checkpoint(checkpoint_path)
        print(f"loading checkpoint: {ckpt_path}")
        raw = torch.load(ckpt_path, map_location="cpu", weights_only=True)
        sd = _tracker_state_dict(raw)
        missing, unexpected = model.load_state_dict(sd, strict=True)
        assert not missing and not unexpected, (missing, unexpected)
        print(f"loaded tracker weights: {len(sd)} tensors")

    model.eval()
    model.requires_grad_(False)
    return model


def make_multiplex_state(model, num_objects=MULTIPLEX_COUNT):
    return model.multiplex_controller.get_state(
        num_objects, torch.device("cpu"), torch.float32, random=False
    )


def synth_frame_features(model, frame_idx, seed=777):
    """Deterministic synthetic backbone features for one frame (stand-in for the
    stateless image encoder, which converts separately)."""
    g = torch.Generator().manual_seed(seed * 100003 + frame_idx)
    e = model.sam_image_embedding_size  # 72
    sizes = [(4 * e, 4 * e), (2 * e, 2 * e), (e, e)]
    # levels 0/1 are pre-projected by sam_mask_decoder.conv_s0/s1 upstream of
    # track_step (see forward_image), so channels are 32/64, not 256
    chans = [32, 64, 256]
    feats = [torch.randn((h * w, 1, c), generator=g) * 0.5
             for (h, w), c in zip(sizes, chans)]
    return feats, sizes


def pos_embed_72(model):
    """Constant sine positional embedding of the 72x72 grid, (HW, 1, C)."""
    from sam3.model.position_encoding import PositionEmbeddingSine

    pe = PositionEmbeddingSine(
        num_pos_feats=256, normalize=True, scale=None, temperature=10000
    )
    e = model.sam_image_embedding_size
    with torch.no_grad():
        pos = pe(torch.zeros(1, 1, e, e))  # (1,256,e,e)
    return pos.flatten(2).permute(2, 0, 1)  # (HW,1,256)


def init_masks(num_objects=MULTIPLEX_COUNT, size=1008):
    """Non-overlapping square blobs for up to 16 objects: (O,1,size,size)."""
    m = torch.zeros(num_objects, 1, size, size)
    # 4x4 grid of blobs inside 1008
    cols = 4
    cell = size // cols
    margin = cell // 8
    for i in range(num_objects):
        r, c = divmod(i, cols)
        r0, c0 = r * cell + margin, c * cell + margin
        r1, c1 = (r + 1) * cell - margin, (c + 1) * cell - margin
        m[i, 0, r0:r1, c0:c1] = 1.0
    return m