MCPcopy Create free account
hub / github.com/PABannier/sam3.cpp / main

Function main

tests/dump_phase7_reference.py:735–832  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

733
734
735def main() -> None:
736 parser = argparse.ArgumentParser()
737 parser.add_argument("--checkpoint", required=True)
738 parser.add_argument("--prephase-ref", default="tests/ref_phase3")
739 parser.add_argument("--phase6-ref", default="tests/ref_phase6")
740 parser.add_argument("--cases", default="tests/phase7_cases.tsv")
741 parser.add_argument("--outdir", default="tests/ref_phase7")
742 args = parser.parse_args()
743
744 os.makedirs(args.outdir, exist_ok=True)
745
746 ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
747 if "model" in ckpt:
748 ckpt = ckpt["model"]
749
750 tracker_weights = {
751 k[len("tracker."):]: v
752 for k, v in ckpt.items()
753 if k.startswith("tracker.")
754 }
755 prompt_weights = {
756 k[len("tracker.sam_prompt_encoder."):]: v
757 for k, v in ckpt.items()
758 if k.startswith("tracker.sam_prompt_encoder.")
759 }
760 mask_weights = {
761 k[len("tracker.sam_mask_decoder."):]: v
762 for k, v in ckpt.items()
763 if k.startswith("tracker.sam_mask_decoder.")
764 }
765
766 neck_trk_0 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_0")).float()
767 neck_trk_1 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_1")).float()
768 neck_trk_2 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_2")).float()
769
770 cases = load_cases(args.cases)
771 for case in cases:
772 case_dir = os.path.join(args.outdir, case.case_id)
773 os.makedirs(case_dir, exist_ok=True)
774 write_meta(case_dir, case)
775
776 save_tensor(os.path.join(case_dir, "input_curr_neck_trk_0"), neck_trk_0)
777 save_tensor(os.path.join(case_dir, "input_curr_neck_trk_1"), neck_trk_1)
778 save_tensor(os.path.join(case_dir, "input_curr_neck_trk_2"), neck_trk_2)
779
780 mem_outputs = []
781 obj_ptrs = []
782 for idx, mem_case in enumerate(case.mem_cases):
783 masks = load_ggml_masks(os.path.join(args.phase6_ref, mem_case, "sam_dec_masks"), 288, 288).float()
784 low_res_mask = masks[:, :1, :, :].contiguous()
785 sam_token = load_ggml_bd(os.path.join(args.phase6_ref, mem_case, "sam_dec_sam_token")).float()
786 obj_score = load_ggml_bd(os.path.join(args.phase6_ref, mem_case, "sam_dec_obj_score")).float()
787
788 save_tensor(os.path.join(case_dir, f"input_mem_mask_logits_{idx}"), low_res_mask)
789 save_tensor(os.path.join(case_dir, f"input_mem_sam_token_{idx}"), sam_token)
790 save_tensor(os.path.join(case_dir, f"input_mem_obj_score_{idx}"), obj_score)
791
792 mem_out = run_memory_encoder(tracker_weights, neck_trk_2, low_res_mask, obj_score)

Callers 1

Calls 15

write_metaFunction · 0.85
load_ggml_masksFunction · 0.85
load_ggml_bdFunction · 0.85
run_memory_encoderFunction · 0.85
run_obj_ptrFunction · 0.85
build_prompt_and_posFunction · 0.85
run_memory_attentionFunction · 0.85
run_sam_decoderFunction · 0.85
load_tensorFunction · 0.70
load_casesFunction · 0.70
save_tensorFunction · 0.70
save_ggml_nchwFunction · 0.70

Tested by

no test coverage detected