"""Synthetic sanity check: does the CHMv2 height-drop cue improve the 0.5776 ensemble's tree_removal score? usage: python chm_eval.py N (run in /workspace with pkg/assets populated) """ import glob import sys import cv2 import numpy as np import torch N = int(sys.argv[1]) if len(sys.argv) > 1 else 600 sys.argv = sys.argv[:1] sys.path.insert(0, "/workspace/pkg/assets") import segmentation_models_pytorch as smp # noqa: E402 from model_v2 import SiamCD # noqa: E402 from chm_fuse import CHM, cue_abs, cue_diff # noqa: E402 from train3b import Synth3 # noqa: E402 dev = "cuda" MEAN = np.array([0.485, 0.456, 0.406], np.float32); STD = np.array([0.229, 0.224, 0.225], np.float32) u = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3) u.load_state_dict(torch.load("/workspace/pkg/assets/model/unet_synth.pt", map_location="cpu", weights_only=True)["state_dict"]) u = u.to(dev).eval() sm = [] for f in sorted(glob.glob("/workspace/pkg/assets/model/siam_*.pt")): s = SiamCD(); s.load_state_dict(torch.load(f, map_location="cpu", weights_only=True)["state_dict"]); sm.append(s.to(dev).eval()) print("satlas models", len(sm), flush=True) chms = {s: CHM("/workspace/pkg/assets/chmv2", dev, scale=s) for s in (1.0, 1.5)} @torch.no_grad() def ens(pres, posts): a = torch.from_numpy(np.stack(pres)).permute(0, 3, 1, 2).to(dev).float() / 255 b = torch.from_numpy(np.stack(posts)).permute(0, 3, 1, 2).to(dev).float() / 255 m = torch.tensor(MEAN, device=dev).view(1, 3, 1, 1); sd = torch.tensor(STD, device=dev).view(1, 3, 1, 1) prob = 0 for dims in ([], [3], [2], [2, 3]): ai, bi = (torch.flip(a, dims), torch.flip(b, dims)) if dims else (a, b) pu = torch.softmax(u(torch.cat([(ai - m) / sd, (bi - m) / sd], 1)), 1)[:, 1:3] cs = 0 for sk in sm: with torch.autocast("cuda", dtype=torch.float16): ch, _, _, _ = sk(ai, bi) cs = cs + ch.float().sigmoid() / len(sm) ci = 0.5 * pu.float() + 0.5 * cs prob = prob + (torch.flip(ci, dims) if dims else ci) / 4 return prob ds = Synth3(10**7) P, HP, HQ, HP2, HQ2, LB, LT, SEM, ND = ([] for _ in range(9)) for s in range(0, N, 16): items = [ds[i] for i in range(s, min(s + 16, N))] pres = [it[0].numpy().transpose(1, 2, 0).copy() for it in items] posts = [it[1].numpy().transpose(1, 2, 0).copy() for it in items] P.append(ens(pres, posts).cpu().numpy()) HP.append(chms[1.0].heights(pres)); HQ.append(chms[1.0].heights(posts)) HP2.append(chms[1.5].heights(pres)); HQ2.append(chms[1.5].heights(posts)) LB += [it[2][0].numpy() > 0.5 for it in items]; LT += [it[2][1].numpy() > 0.5 for it in items] SEM += [it[4].numpy() for it in items] ND += [(a.max(2) == 0) | (b.max(2) == 0) for a, b in zip(pres, posts)] P = np.concatenate(P); HP = torch.cat(HP); HQ = torch.cat(HQ); HP2 = torch.cat(HP2); HQ2 = torch.cat(HQ2); ND = np.stack(ND) P[:, 0][ND] = 0; P[:, 1][ND] = 0 sem = np.stack(SEM); lt = np.stack(LT); hp = HP.cpu().numpy(); hq = HQ.cpu().numpy() print(f"CHM metres pre tree px {hp[sem == 2].mean():.2f} pre non-tree px {hp[(sem != 2) & (sem != 255)].mean():.2f} " f"removed px pre {hp[lt].mean():.2f} post {hq[lt].mean():.2f} (x1.5: tree {HP2.cpu().numpy()[sem == 2].mean():.2f} removed pre {HP2.cpu().numpy()[lt].mean():.2f} post {HQ2.cpu().numpy()[lt].mean():.2f}) overall p50/p90/p99 {np.percentile(hp, [50, 90, 99]).round(2).tolist()}", flush=True) def clean(m, amin=30): n, cc, st, _ = cv2.connectedComponentsWithStats(m.astype(np.uint8), 8) keep = np.zeros(n, bool); keep[1:] = st[1:, cv2.CC_STAT_AREA] >= amin return keep[cc] def score(prob, gts): k3 = np.ones((3, 3), np.uint8); tp = fp = fn = tn = 0; shp = [] for p, g in zip(prob, gts): pm = clean(p > 0.5); pp, gp = pm.sum() >= 20, g.sum() >= 20 tp += pp and gp; fp += pp and not gp; fn += gp and not pp; tn += (not pp) and (not gp) if gp: if not pp: shp.append(0.0); continue gd = cv2.dilate(g.astype(np.uint8), k3).astype(bool); pd = cv2.dilate(pm.astype(np.uint8), k3).astype(bool) pr = (pm & gd).sum() / pm.sum(); rc = (g & pd).sum() / g.sum() shp.append(2 * pr * rc / (pr + rc + 1e-9)) f1p = 2 * tp / max(2 * tp + fp + fn, 1); f1n = 2 * tn / max(2 * tn + fn + fp, 1) return 0.5 * (f1p + f1n) / 2 + 0.5 * (float(np.mean(shp)) if shp else 0.0), (tp, fp, fn, tn) np.savez("/workspace/eval_arr.npz", pt=P[:, 1].astype(np.float16), hp=HP.cpu().numpy().astype(np.float16), hq=HQ.cpu().numpy().astype(np.float16), hp2=HP2.cpu().numpy().astype(np.float16), hq2=HQ2.cpu().numpy().astype(np.float16), lt=lt, nd=ND) print("SAVED", flush=True) sb, cb = score(P[:, 0], LB) print(f"building {sb:.4f} {cb}", flush=True) cues = {"abs1.0": cue_abs(HP, HQ), "diff1.0": cue_diff(HP, HQ), "abs1.5": cue_abs(HP2, HQ2), "diff1.5": cue_diff(HP2, HQ2)} res = {} for name, c in cues.items(): c = c.cpu().numpy(); c[ND] = 0 for a in (0.0, 0.2, 0.3, 0.4, 0.5, 1.0): st, ct = score((1 - a) * P[:, 1] + a * c, LT) res[(name, a)] = st print(f"tree cue={name} a={a:.1f} {st:.4f} tp/fp/fn/tn {ct}", flush=True) print("EVAL_DONE", flush=True)