File size: 5,037 Bytes
ce209f5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | """
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())
|