()
| 733 | |
| 734 | |
| 735 | def 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) |
no test coverage detected