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