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

Function main

tests/dump_phase6_reference.py:444–489  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

442
443
444def main() -> None:
445 parser = argparse.ArgumentParser()
446 parser.add_argument("--checkpoint", required=True)
447 parser.add_argument("--prephase-ref", default="tests/ref_phase3")
448 parser.add_argument("--cases", default="tests/phase6_cases.tsv")
449 parser.add_argument("--outdir", default="tests/ref_phase6")
450 args = parser.parse_args()
451
452 os.makedirs(args.outdir, exist_ok=True)
453
454 print("Loading checkpoint...")
455 ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
456 if "model" in ckpt:
457 ckpt = ckpt["model"]
458
459 prompt_weights = {
460 k[len("tracker.sam_prompt_encoder."):]: v
461 for k, v in ckpt.items()
462 if k.startswith("tracker.sam_prompt_encoder.")
463 }
464 mask_weights = {
465 k[len("tracker.sam_mask_decoder."):]: v
466 for k, v in ckpt.items()
467 if k.startswith("tracker.sam_mask_decoder.")
468 }
469 no_mem_embed = ckpt["tracker.no_mem_embed"].float()
470
471 neck_trk_0 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_0")).float()
472 neck_trk_1 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_1")).float()
473 neck_trk_2 = load_tensor(os.path.join(args.prephase_ref, "neck_trk_2")).float()
474
475 cases = load_cases(args.cases)
476 for case in cases:
477 print(f"Dumping {case.case_id}...")
478 dump_case(
479 case,
480 prompt_weights,
481 mask_weights,
482 no_mem_embed,
483 neck_trk_0,
484 neck_trk_1,
485 neck_trk_2,
486 os.path.join(args.outdir, case.case_id),
487 )
488
489 print(f"All tensors saved to: {args.outdir}")
490
491
492if __name__ == "__main__":

Callers 1

Calls 3

dump_caseFunction · 0.85
load_tensorFunction · 0.70
load_casesFunction · 0.70

Tested by

no test coverage detected