import torch
from sam2.build_sam import build_sam2_video_predictor
import numpy as np
predictor = build_sam2_video_predictor(
"configs/sam2.1/sam2.1_hiera_l.yaml",
"./checkpoints/sam2.1_hiera_large.pt",
device="cuda"
)
video_path = "./video_frames"
inference_state = predictor.init_state(video_path=video_path)
# 跟踪多个对象
objects_to_track = [
{"id": 1, "point": [200, 150], "frame": 0}, # 人员 1
{"id": 2, "point": [400, 200], "frame": 0}, # 人员 2
{"id": 3, "point": [600, 300], "frame": 0}, # 球
]
for obj in objects_to_track:
predictor.add_new_points_or_box(
inference_state=inference_state,
frame_idx=obj["frame"],
obj_id=obj["id"],
points=np.array([obj["point"]], dtype=np.float32),
labels=np.array([1], dtype=np.int32)
)
# 传播所有对象
all_masks = {}
for frame_idx, obj_ids, mask_logits in predictor.propagate_in_video(inference_state):
all_masks[frame_idx] = {}
for i, obj_id in enumerate(obj_ids):
all_masks[frame_idx][obj_id] = (mask_logits[i] > 0.0).cpu().numpy()