File size: 5,970 Bytes
d70361b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
"""Tuesday baseline (Uday P0): held-out F1/P/R with REAL v3_frozen weights.

Evaluates the delhi_cd test split (never used for training calibration of this
checkpoint) via ``app.evaluation.metrics.binary_metrics`` + delhi_cd GT masks.

Usage:
    python scripts/record_tuesday_baseline.py
"""
from __future__ import annotations

import json
import os
import sys
import time
from datetime import datetime, timezone
from pathlib import Path

import numpy as np
from PIL import Image

ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))

try:
    from dotenv import load_dotenv
    load_dotenv(ROOT / ".env", override=True)
except ImportError:
    pass

CKPT = (ROOT / "models" / "adaptformer_delhi" / "v3_frozen").resolve()
os.environ["ADAPTFORMER_WEIGHTS"] = str(CKPT)
# Model-native operating point for Friday apples-to-apples comparison
os.environ["ADAPTFORMER_THRESHOLD"] = "0.2"
os.environ["DETECTION_DL_THRESHOLD"] = "0.2"
os.environ["DETECTION_TTA"] = "off"
os.environ["DETECTION_FUSION"] = "dl_only"

from app.evaluation.metrics import binary_metrics  # noqa: E402
from app.model_inference import (  # noqa: E402
    get_calibrated_threshold,
    get_loaded_weights_source,
    get_model_status,
    predict_change_mask,
    preload_model,
)

OUT = ROOT / "runs" / "tuesday_baseline_20260728"
TEST_IDS = ROOT / "data" / "delhi_cd" / "test" / "pair_ids.txt"
SPLIT_ROOT = ROOT / "data" / "delhi_cd" / "test"


def _load_rgb(path: Path) -> np.ndarray:
    with Image.open(path) as im:
        return np.asarray(im.convert("RGB"))


def _load_gt(path: Path) -> np.ndarray:
    with Image.open(path) as im:
        return (np.asarray(im.convert("L")) > 127)


def _pair_paths(pair_id: str) -> tuple[Path, Path, Path]:
    # Prefer split folder layout; fall back to flat delhi_cd tiles if present.
    for base in (SPLIT_ROOT, ROOT / "data" / "delhi_cd"):
        b = base / f"{pair_id}_before.png"
        a = base / f"{pair_id}_after.png"
        g = base / f"{pair_id}_gt.png"
        if not g.is_file():
            g = base / f"{pair_id}_mask.png"
        if b.is_file() and a.is_file() and g.is_file():
            return b, a, g
    # delhi_cd/test/manifest.json → library_sources geotiffs + docs labels
    man_path = SPLIT_ROOT / "manifest.json"
    if man_path.is_file():
        man = json.loads(man_path.read_text(encoding="utf-8"))
        for p in man.get("pairs", []):
            if p.get("pair_id") == pair_id:
                before = ROOT / p["before_path"]
                after = ROOT / p["after_path"]
                gt = ROOT / p["gt_mask"]
                if before.is_file() and after.is_file() and gt.is_file():
                    return before, after, gt
    raise FileNotFoundError(f"Missing assets for {pair_id}")


def _load_pair_arrays(before_p: Path, after_p: Path, gt_p: Path):
    from app.evaluation.delhi_eval import _load_label, _load_rgb as _er
    before = _er(before_p)
    after = _er(after_p)
    gt = _load_label(gt_p)
    return before, after, gt

def main() -> int:
    if not (CKPT / "model.safetensors").is_file():
        print(f"MISSING weights: {CKPT / 'model.safetensors'}")
        print("Run: python scripts/export_adaptformer_delhi.py --src models/adaptformer_delhi/best_v3")
        return 1

    OUT.mkdir(parents=True, exist_ok=True)
    pair_ids = [ln.strip() for ln in TEST_IDS.read_text(encoding="utf-8").splitlines() if ln.strip()]
    print(f"ckpt={CKPT}")
    print(f"test pairs ({len(pair_ids)}): {pair_ids}")

    ok = preload_model()
    status = get_model_status()
    loaded = get_loaded_weights_source() or status.get("loadedFrom")
    print(f"preload ok={ok} loadedFrom={loaded} thr={get_calibrated_threshold(0.2)}")
    if not ok or (loaded and "v3_frozen" not in str(loaded).replace("\\", "/")):
        print("ERROR: v3_frozen was not loaded — refusing to record baseline")
        return 1

    thr = float(get_calibrated_threshold(0.2) or 0.2)
    rows = []
    t0 = time.time()
    for pid in pair_ids:
        before_p, after_p, gt_p = _pair_paths(pid)
        before, after, gt = _load_pair_arrays(before_p, after_p, gt_p)
        gt = np.asarray(gt) > 127
        pred_mask, _score = predict_change_mask(before, after, threshold=thr)
        pred = np.asarray(pred_mask) > 127
        if pred.shape != gt.shape:
            pred = np.array(
                Image.fromarray((pred.astype(np.uint8) * 255)).resize(
                    (gt.shape[1], gt.shape[0]), Image.NEAREST
                )
            ) > 127
        m = binary_metrics(pred, gt)
        row = {"pair_id": pid, **m, "threshold": thr}
        rows.append(row)
        print(
            f"  {pid}: F1={m['f1']:.4f} P={m['precision']:.4f} "
            f"R={m['recall']:.4f} IoU={m['iou']:.4f}"
        )

    n = max(len(rows), 1)
    summary = {
        "date": datetime.now(timezone.utc).astimezone().isoformat(),
        "role": "Tuesday baseline for Friday comparison (Uday P0)",
        "weights": str(CKPT),
        "loadedFrom": str(loaded),
        "threshold": thr,
        "fusion": "dl_only",
        "tta": "off",
        "split": "data/delhi_cd/test",
        "pair_ids": pair_ids,
        "n_pairs": len(rows),
        "mean_f1": round(sum(r["f1"] for r in rows) / n, 4),
        "mean_precision": round(sum(r["precision"] for r in rows) / n, 4),
        "mean_recall": round(sum(r["recall"] for r in rows) / n, 4),
        "mean_iou": round(sum(r["iou"] for r in rows) / n, 4),
        "elapsed_sec": round(time.time() - t0, 1),
        "model_status": status,
        "per_pair": rows,
    }
    out_path = OUT / "metrics.json"
    out_path.write_text(json.dumps(summary, indent=2), encoding="utf-8")
    print(f"\nWrote {out_path}")
    print(
        f"BASELINE  F1={summary['mean_f1']}  "
        f"P={summary['mean_precision']}  R={summary['mean_recall']}  "
        f"IoU={summary['mean_iou']}"
    )
    return 0


if __name__ == "__main__":
    raise SystemExit(main())