CD-Models / evaluate.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
5.04 kB
"""
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())