()
| 442 | |
| 443 | |
| 444 | def 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 | |
| 492 | if __name__ == "__main__": |
no test coverage detected