""" evaluate.py - Run test sweep on a trained cd-model checkpoint. Usage: python evaluate.py --model changemamba --dataset levir_cd python evaluate.py --model all --dataset all python evaluate.py --checkpoint results/bifa/levir_cd/checkpoints/best_model.pth --model bifa --dataset levir_cd """ from __future__ import annotations import argparse import sys from pathlib import Path import torch import yaml from utils.config_loader import load_dataset_config, load_model_config from utils.gpu_utils import print_gpu_diagnostics, resolve_gpu from utils.model_adapters import get_model_adapter from utils.results_writer import append_to_comparison_table from utils.unified_evaluator import evaluate_with_adapter ROOT = Path(__file__).resolve().parent if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) class ModelEvaluationUnavailable(RuntimeError): pass def _load_datasets() -> list[str]: return sorted(path.stem for path in (ROOT / "configs" / "datasets").glob("*.yaml")) def _load_registry() -> dict: with (ROOT / "configs" / "models" / "registry.yaml").open("r", encoding="utf-8") as f: return yaml.safe_load(f)["models"] def _best_checkpoint(model: str, dataset: str) -> Path: ckpt_dir = ROOT / "results" / model / dataset / "checkpoints" direct = ckpt_dir / "best_model.pth" if direct.exists(): return direct candidates = ( sorted(ckpt_dir.glob("*F1*.pth")) + sorted(ckpt_dir.glob("*.pth")) + sorted(ckpt_dir.glob("**/best_ckpt.pth")) + sorted(ckpt_dir.glob("**/best_ckpt.pt")) + sorted(ckpt_dir.glob("**/*best*.pth")) + sorted(ckpt_dir.glob("**/*best*.pt")) ) if not candidates: raise FileNotFoundError(f"No checkpoint found under {ckpt_dir}") return candidates[-1] def _parse_args() -> argparse.Namespace: registry = _load_registry() datasets = _load_datasets() parser = argparse.ArgumentParser(description="Run a test sweep on a trained cd-model checkpoint.") parser.add_argument("--model", required=True, choices=sorted(registry) + ["all"]) parser.add_argument("--dataset", required=True, choices=datasets + ["all"]) parser.add_argument("--checkpoint", default=None) parser.add_argument("--gpu", default="0") parser.add_argument("--dry-run", action="store_true") parser.add_argument("--allow-missing-profilers", action="store_true") return parser.parse_args() def _evaluate_one( model: str, dataset: str, checkpoint: str | None, gpu: str, dry_run: bool, allow_missing_profilers: bool, ) -> int: cfg = load_dataset_config(dataset) gpu_resolution = resolve_gpu(gpu) print_gpu_diagnostics(gpu_resolution) if dry_run: print(f"[EVAL-DRY-RUN] model={model} dataset={dataset} data_root={cfg['data_root']}") return 0 ckpt = Path(checkpoint) if checkpoint else _best_checkpoint(model, dataset) if not ckpt.is_absolute(): ckpt = ROOT / ckpt print(f"[EVAL] model={model} dataset={dataset} checkpoint={ckpt}") adapter = get_model_adapter(model) if not adapter.supports_inprocess_eval: raise ModelEvaluationUnavailable(adapter.notes_or_failure_reason) device = torch.device(gpu_resolution.local_device) _, code = evaluate_with_adapter( model_name=model, dataset_cfg=cfg, model_config=load_model_config(model), adapter=adapter, checkpoint_path=ckpt, device=device, strict_profiling=not allow_missing_profilers, ) return code def main() -> int: args = _parse_args() models = sorted(_load_registry()) if args.model == "all" else [args.model] datasets = _load_datasets() if args.dataset == "all" else [args.dataset] rows = [] for model in models: for dataset in datasets: try: code = _evaluate_one( model, dataset, args.checkpoint, args.gpu, args.dry_run, args.allow_missing_profilers, ) rows.append((model, dataset, "complete" if code == 0 else "failed")) except FileNotFoundError as exc: print(f"[MISSING] {model}/{dataset}: {exc}") rows.append((model, dataset, "missing")) except ModelEvaluationUnavailable as exc: print(f"[UNAVAILABLE] {model}/{dataset}: {exc}") rows.append((model, dataset, "unavailable")) except RuntimeError as exc: print(f"[FAILED] {model}/{dataset}: {exc}") rows.append((model, dataset, "failed")) append_to_comparison_table() print("\nModel | Dataset | Status") print("--- | --- | ---") for model, dataset, status in rows: print(f"{model} | {dataset} | {status}") return 0 if all(status == "complete" for _, _, status in rows) else 1 if __name__ == "__main__": raise SystemExit(main())