Detector / scripts /eval_detector.py
Benxelua's picture
Add reproducible six-run detector-domain training package
4409fdb verified
Raw History Blame Contribute Delete
1.4 kB
#!/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()