| """ |
| 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()) |
|
|