File size: 1,397 Bytes
4409fdb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
#!/usr/bin/env python3
"""Re-evaluate a selected checkpoint without ever using test data for selection."""
from __future__ import annotations

import argparse, json
from pathlib import Path
import sys
import torch, yaml
from torch.utils.data import DataLoader
sys.path.insert(0, str(Path(__file__).resolve().parent))
from train_detector import evaluate
from detector_lib import LockedLODDataset, build_model, collate

def main():
    p = argparse.ArgumentParser(); p.add_argument("--config", type=Path, required=True); p.add_argument("--dataset-root", type=Path, required=True); p.add_argument("--labels-root", type=Path, required=True); p.add_argument("--checkpoint", type=Path, required=True); p.add_argument("--output", type=Path, required=True); a=p.parse_args()
    cfg=yaml.safe_load(a.config.read_text()); device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model, processor, kind=build_model(cfg); model.load_state_dict(torch.load(a.checkpoint,map_location="cpu")["model"]); model.to(device)
    results={}
    for manifest in cfg["test_manifests"]:
        ds=LockedLODDataset(a.config.parent.parent/manifest,a.dataset_root,a.labels_root)
        results[Path(manifest).stem]=evaluate(model,processor,kind,DataLoader(ds,batch_size=1,num_workers=0,collate_fn=collate),device)
    a.output.write_text(json.dumps(results,indent=2)+"\n")
if __name__ == "__main__": main()