sam3.1-coreml / tracker /scripts /eager_baseline.py
sdkv2's picture
Merge tracker + backbone, drop kino references
47d8ad6 verified
Raw
History Blame Contribute Delete
4.12 kB
"""Eager SAM3 track_step baseline: frame 0 mask-init + frames 1..TARGET propagation.
Caches per-frame outputs needed by the dense wrapper + TARGET/warmup eager
outputs for parity checking. One-time cost; results go to eager_cache.pt.
"""
import time
import torch
import common
TARGET = 16 # first frame with FULL banks: mem needs f-1..f-6, ptrs need f-1..f-15
WARMUP = 3 # partial banks — validates validity-mask path
def bf_prop(model, frame_idx):
feats, sizes = common.synth_frame_features(model, frame_idx)
pos72 = common.POS72
return {
"vision_feats": feats,
"vision_masks": [None, None, None],
"vision_pos_embeds": [None, None, pos72],
"feat_sizes": sizes,
}
def bf_inter(model, frame_idx):
feats, sizes = common.synth_frame_features(model, frame_idx)
return {"vision_feats": feats, "feat_sizes": sizes}
def main():
model = common.build_model()
common.POS72 = common.pos_embed_72(model)
ms = common.make_multiplex_state(model)
output_dict = {"cond_frame_outputs": {}, "non_cond_frame_outputs": {}}
with torch.no_grad():
# ---- frame 0: init via mask input (mode: mask_as_output) ----
t0 = time.time()
out0 = model.track_step(
frame_idx=0,
is_init_cond_frame=True,
backbone_features_interactive=bf_inter(model, 0),
backbone_features_propagation=bf_prop(model, 0),
image=None,
point_inputs=None,
mask_inputs=common.init_masks(),
gt_masks=None,
frames_to_add_correction_pt=[],
output_dict=output_dict,
num_frames=64,
multiplex_state=ms,
)
output_dict["cond_frame_outputs"][0] = out0
print(f"frame 0 (init): {time.time()-t0:.1f}s")
for k, v in out0.items():
if torch.is_tensor(v):
print(f" {k}: {tuple(v.shape)} {v.dtype}")
elif isinstance(v, list) and v and torch.is_tensor(v[0]):
print(f" {k}: list[{tuple(v[0].shape)}]")
else:
print(f" {k}: {type(v).__name__} {v if isinstance(v, set) else ''}")
# ---- frames 1..TARGET: propagation only ----
for f in range(1, TARGET + 1):
t0 = time.time()
out = model.track_step(
frame_idx=f,
is_init_cond_frame=False,
backbone_features_interactive=None,
backbone_features_propagation=bf_prop(model, f),
image=None,
point_inputs=None,
mask_inputs=None,
gt_masks=None,
frames_to_add_correction_pt=[],
output_dict=output_dict,
num_frames=64,
multiplex_state=ms,
)
output_dict["non_cond_frame_outputs"][f] = out
print(f"frame {f}: {time.time()-t0:.1f}s "
f"maskmem {tuple(out['maskmem_features'].shape)} "
f"obj_ptr {tuple(out['obj_ptr'].shape)}")
# ---- cache what the dense wrapper needs ----
keep = ["maskmem_features", "maskmem_pos_enc", "obj_ptr",
"image_features", "image_pos_enc", "object_score_logits"]
cache = {"cond": {}, "frames": {}, "target": {}, "warmup": {}}
for k in keep:
cache["cond"][k] = out0.get(k)
for f in range(1, TARGET):
cache["frames"][f] = {k: output_dict["non_cond_frame_outputs"][f].get(k)
for k in keep}
def pack_target(frame_idx, bucket):
tgt = output_dict["non_cond_frame_outputs"][frame_idx]
for k in ["pred_masks", "pred_masks_high_res", "object_score_logits",
"maskmem_features", "obj_ptr"]:
bucket[k] = tgt[k]
bucket["ious"] = tgt["multistep_pred_ious"][-1]
pack_target(TARGET, cache["target"])
cache["target_frame"] = TARGET
pack_target(WARMUP, cache["warmup"])
cache["warmup_frame"] = WARMUP
torch.save(cache, "eager_cache.pt")
print("saved eager_cache.pt")
if __name__ == "__main__":
main()