Download code/chm/gate_eval.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 2.38 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/chm/gate_eval.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/chm/gate_eval.py
-
curl -L -o gate_eval.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/chm/gate_eval.py
2.38 kB
| """Component-level CHM gating on saved synthetic arrays.""" | |
| import cv2 | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| import sys | |
| sys.argv = sys.argv[:1] | |
| from chm_eval_score import score # noqa: E402 | |
| d = np.load("/workspace/eval_arr.npz") | |
| pt, lt, nd = d["pt"].astype(np.float32), d["lt"], d["nd"] | |
| def sm(h, k=5): | |
| return F.avg_pool2d(torch.from_numpy(h.astype(np.float32))[:, None], k, 1, k // 2, count_include_pad=False)[:, 0].numpy() | |
| def gate(pt, hp, hq, drop_min, pre_min, add=None): | |
| out = [] | |
| for p, a, b in zip(pt, hp, hq): | |
| m = (p > 0.5).astype(np.uint8) | |
| n, cc, st, _ = cv2.connectedComponentsWithStats(m, 8) | |
| keep = m.astype(bool).copy() | |
| for i in range(1, n): | |
| r = cc == i | |
| if (a[r] - b[r]).mean() < drop_min and a[r].mean() < pre_min: | |
| keep[r] = False | |
| if add is not None: # add strong CHM removal blobs where the ensemble is at least weakly positive | |
| hmin, gmax, pmin, amin = add | |
| c = ((a > hmin) & (b < gmax) & (p > pmin)).astype(np.uint8) | |
| n2, cc2, st2, _ = cv2.connectedComponentsWithStats(c, 8) | |
| for i in range(1, n2): | |
| if st2[i, cv2.CC_STAT_AREA] >= amin: | |
| keep |= cc2 == i | |
| out.append(keep.astype(np.float32)) | |
| return np.stack(out) | |
| for tag, hp, hq in (("x1.0", d["hp"], d["hq"]), ("x1.5", d["hp2"], d["hq2"])): | |
| hp, hq = sm(hp), sm(hq) | |
| print(tag, "base", round(score(pt, lt)[0], 4), flush=True) | |
| for dm, pm in ((0.5, 2.0), (1.0, 2.5), (1.5, 3.0), (2.0, 4.0)): | |
| s, c = score(gate(pt, hp, hq, dm, pm), lt) | |
| print(f"{tag} drop-gate drop<{dm} & pre<{pm}: {s:.4f} {c}", flush=True) | |
| for add in ((5, 2, 0.3, 40), (5, 2, 0.2, 60), (6, 1.5, 0.15, 80), (4, 2, 0.35, 30)): | |
| s, c = score(gate(pt, hp, hq, 1.0, 2.5, add), lt) | |
| print(f"{tag} drop-gate(1.0,2.5)+add{add}: {s:.4f} {c}", flush=True) | |
| # how many GT-negative pairs would a pure CHM blob detector flag | |
| fp = sum(((a > 5) & (b < 2)).sum() >= 60 for a, b, g in zip(hp, hq, lt) if g.sum() < 20) | |
| tp = sum(((a > 5) & (b < 2)).sum() >= 60 for a, b, g in zip(hp, hq, lt) if g.sum() >= 20) | |
| print(f"{tag} CHM-only blob(>5m -> <2m, >=60px): flags {tp}/{int(sum(g.sum() >= 20 for g in lt))} pos, {fp}/{int(sum(g.sum() < 20 for g in lt))} neg", flush=True) | |
| print("GATE_DONE") | |