Download code/chm/chm_eval.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 5.27 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/chm/chm_eval.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/chm/chm_eval.py
-
curl -L -o chm_eval.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/chm/chm_eval.py
5.27 kB
| """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)} | |
| 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) | |