fnruha0921's picture
Upload code/chm/chm_eval.py with huggingface_hub
13e8ef9 verified
Raw History Blame Contribute Delete
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)}
@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)