fnruha0921's picture
add code
eea5f0e verified
Raw History Blame Contribute Delete
11.3 kB
"""Train the 6-channel UNet (baseline architecture) on synthetic pre/post pairs.
Label map (softmax, same as baseline): 0 background, 1 new_building, 2 tree_removal.
Scene label codes from prep.py: 1 building, 2 tree, 3 ground donor.
"""
import math
import os
import sys
import time
import cv2
import numpy as np
import segmentation_models_pytorch as smp
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
P = "/workspace/prep"
CROP = 192
MEAN = np.array([0.485, 0.456, 0.406], np.float32)
STD = np.array([0.229, 0.224, 0.225], np.float32)
MINUTES = float(sys.argv[1]) if len(sys.argv) > 1 else 25
INIT = sys.argv[2] if len(sys.argv) > 2 else "/workspace/unet_r18_cd.pt"
PEAK = float(sys.argv[3]) if len(sys.argv) > 3 else 4e-4
OUT = sys.argv[4] if len(sys.argv) > 4 else "/workspace/unet_synth2"
class Synth(Dataset):
def __init__(self, n):
self.n = n
fx = [np.load(f"{P}/flair_x.npy")]; fyl = [np.load(f"{P}/flair_y.npy")]
if os.path.exists(f"{P}/flair2_y.npy"):
fx.append(np.load(f"{P}/flair2_x.npy")); fyl.append(np.load(f"{P}/flair2_y.npy"))
self.fx = np.concatenate(fx); self.fy = np.concatenate(fyl)
self.tx = np.load(f"{P}/tcd_x.npy", mmap_mode="r"); self.ty = np.load(f"{P}/tcd_y.npy", mmap_mode="r")
fy = self.fy
self.b_idx = np.flatnonzero((fy == 1).mean((1, 2)) > 0.01)
self.t_idx = np.flatnonzero((fy == 2).mean((1, 2)) > 0.15)
self.g_idx = np.flatnonzero((fy == 3).mean((1, 2)) > 0.6)
print("building scenes", len(self.b_idx), "tree scenes", len(self.t_idx), "ground donors", len(self.g_idx), flush=True)
def __len__(self):
return self.n
# ---------- scene sampling ----------
def scene(self, rng, kind):
if kind == "tcd":
i = rng.integers(len(self.tx)); x, y = self.tx[i], self.ty[i]
s = rng.uniform(0.85, 1.15)
size = int(round(CROP / s)); size = min(size, 320)
r, c = rng.integers(0, 321 - size, 2)
x = cv2.resize(np.ascontiguousarray(x[r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_AREA)
y = cv2.resize(np.ascontiguousarray(y[r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_NEAREST)
return x.copy(), y.copy()
pool = {"b": self.b_idx, "t": self.t_idx, "g": self.g_idx, "any": None}[kind]
i = rng.integers(len(self.fx)) if pool is None else pool[rng.integers(len(pool))]
return np.array(self.fx[i]), np.array(self.fy[i])
def donor(self, rng):
x, _ = self.scene(rng, "g")
k = rng.integers(4)
x = np.rot90(x, k).copy()
return x
@staticmethod
def paste(dst, src, region, rng):
"""Replace region of dst by src pixels with feathered edge and mean colour matching."""
if region.sum() == 0:
return dst
ring = cv2.dilate(region.astype(np.uint8), np.ones((9, 9), np.uint8)) & (~region.astype(bool))
src = src.astype(np.float32)
if ring.sum() > 10 and rng.random() < 0.7:
a = rng.uniform(0.3, 0.8)
shift = dst[ring.astype(bool)].mean(0) - src[region.astype(bool)].mean(0)
src = np.clip(src + a * shift, 0, 255)
alpha = cv2.GaussianBlur(region.astype(np.float32), (0, 0), rng.uniform(0.6, 1.2))
alpha = np.maximum(alpha, region.astype(np.float32) * 0.9)[..., None]
return (dst * (1 - alpha) + src * alpha).astype(np.uint8)
@staticmethod
def blob(rng, shape, area):
m = np.zeros(shape, np.uint8)
cy, cx = rng.integers(0, shape[0]), rng.integers(0, shape[1])
rad = math.sqrt(area / math.pi)
for _ in range(rng.integers(1, 5)):
oy, ox = rng.normal(0, rad * 0.5, 2)
ax = (int(max(3, rad * rng.uniform(0.5, 1.3))), int(max(3, rad * rng.uniform(0.4, 1.1))))
cv2.ellipse(m, (int(cx + ox), int(cy + oy)), ax, rng.uniform(0, 180), 0, 360, 1, -1)
if rng.random() < 0.5: # jagged edge
noise = cv2.GaussianBlur(rng.random(shape).astype(np.float32), (0, 0), 3)
m = ((cv2.GaussianBlur(m.astype(np.float32), (0, 0), 4) + (noise - 0.5) * 0.6) > 0.5).astype(np.uint8)
return m
# ---------- photometric: independent per epoch ----------
@staticmethod
def photo(img, rng, tree=None):
x = img.astype(np.float32)
if tree is not None and rng.random() < 0.3: # seasonal: vegetation toward brown / pale
hsv = cv2.cvtColor(img, cv2.COLOR_RGB2HSV).astype(np.float32)
t = cv2.GaussianBlur(tree.astype(np.float32), (0, 0), 2)
hsv[..., 0] -= t * rng.uniform(5, 25)
hsv[..., 1] *= 1 - t * rng.uniform(0, 0.5)
hsv[..., 0] %= 180
x = cv2.cvtColor(np.clip(hsv, 0, 255).astype(np.uint8), cv2.COLOR_HSV2RGB).astype(np.float32)
x = x * rng.uniform(0.7, 1.3) + rng.uniform(-30, 30) # contrast / brightness
x = x * rng.uniform(0.88, 1.12, 3) + rng.uniform(-12, 12, 3) # colour cast
g = rng.uniform(0.7, 1.4)
x = 255 * (np.clip(x, 0, 255) / 255) ** g
m = x.mean(2, keepdims=True); x = m + (x - m) * rng.uniform(0.6, 1.4) # saturation
if rng.random() < 0.3:
x = cv2.GaussianBlur(x, (0, 0), rng.uniform(0.3, 1.0))
if rng.random() < 0.3:
x = x + rng.normal(0, rng.uniform(1, 6), x.shape)
return np.clip(x, 0, 255).astype(np.uint8)
def __getitem__(self, idx):
rng = np.random.default_rng((idx * 7919 + os.getpid() * 104729 + time.time_ns()) % 2**32)
r = rng.random()
mode = "b" if r < 0.33 else ("t" if r < 0.66 else ("bt" if r < 0.74 else "neg"))
if "t" in mode:
x, y = self.scene(rng, "tcd" if rng.random() < 0.5 else "t")
elif mode == "b":
x, y = self.scene(rng, "b")
else:
x, y = self.scene(rng, ["tcd", "any", "b", "t"][rng.integers(4)])
pre, post = x.copy(), x.copy()
lab = np.zeros(y.shape, np.uint8)
bmask = (y == 1).astype(np.uint8); tmask = (y == 2).astype(np.uint8)
if "b" in mode and bmask.sum() > 20:
n, cc, stats, _ = cv2.connectedComponentsWithStats(bmask, 8)
comps = [i for i in range(1, n) if stats[i, cv2.CC_STAT_AREA] >= 12]
if comps:
k = rng.integers(1, len(comps) + 1)
sel = rng.choice(comps, size=k, replace=False)
region = np.isin(cc, sel).astype(np.uint8)
if rng.random() < 0.2: # partial extension: only part of a building is new
cut = self.blob(rng, region.shape, region.sum() * 0.5)
region = region & cut
if region.sum() >= 12:
grow = cv2.dilate(region, np.ones((3, 3), np.uint8), iterations=int(rng.integers(1, 3)))
pre = self.paste(pre, self.donor(rng), grow, rng)
lab[region.astype(bool)] = 1
if "t" in mode and tmask.sum() > 150:
for _ in range(10):
b = self.blob(rng, tmask.shape, rng.uniform(150, 6000))
region = b & cv2.dilate(tmask, np.ones((3, 3), np.uint8))
if (region & tmask).sum() >= 80:
post = self.paste(post, self.donor(rng), region, rng)
lab[(region & tmask).astype(bool) & (lab == 0)] = 2
break
if mode == "neg" and rng.random() < 0.3 and bmask.sum() > 20: # roof colour change only
post = post.astype(np.int16); post[bmask.astype(bool)] += rng.integers(-60, 60, 3).astype(np.int16)
post = np.clip(post, 0, 255).astype(np.uint8)
if lab.any() and rng.random() < 0.15: # reversed pair (demolition / regrowth) -> no target change
pre, post = post, pre
lab[:] = 0
pre = self.photo(pre, rng, tmask); post = self.photo(post, rng, tmask)
if rng.random() < 0.7: # misregistration of pre (parallax / residual shift)
M = np.float32([[1, 0, rng.uniform(-2.5, 2.5)], [0, 1, rng.uniform(-2.5, 2.5)]])
pre = cv2.warpAffine(pre, M, (CROP, CROP), borderMode=cv2.BORDER_REFLECT)
if rng.random() < 0.1: # no-data black region (straight edge), independent per image
for im in (pre, post) if rng.random() < 0.5 else ((pre,) if rng.random() < 0.5 else (post,)):
yy, xx = np.mgrid[:CROP, :CROP]
th = rng.uniform(0, 2 * np.pi); d = rng.uniform(-CROP * 0.7, -CROP * 0.2)
nd = (xx - CROP / 2) * np.cos(th) + (yy - CROP / 2) * np.sin(th) < d
im[nd] = 0
lab[nd] = 0
k = rng.integers(8) # shared dihedral transform
def d4(a):
a = np.rot90(a, k % 4)
return a[:, ::-1] if k >= 4 else a
pre, post, lab = d4(pre), d4(post), d4(lab)
inp = np.concatenate([(pre / 255.0 - MEAN) / STD, (post / 255.0 - MEAN) / STD], 2).astype(np.float32)
return torch.from_numpy(inp.transpose(2, 0, 1).copy()), torch.from_numpy(lab.copy()).long()
def dice_loss(logits, y):
p = logits.softmax(1)
loss = 0
for c in (1, 2):
t = (y == c).float(); q = p[:, c]
loss += 1 - (2 * (q * t).sum() + 1) / (q.sum() + t.sum() + 1)
return loss / 2
def main():
torch.backends.cudnn.benchmark = True
model = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3)
ck = torch.load(INIT, map_location="cpu", weights_only=True)
model.load_state_dict(ck["state_dict"]) # start from the organizer baseline weights
model = model.cuda().to(memory_format=torch.channels_last)
ds = Synth(10**7)
dl = DataLoader(ds, batch_size=64, num_workers=30, pin_memory=True, persistent_workers=True, prefetch_factor=4)
opt = torch.optim.AdamW(model.parameters(), lr=4e-4, weight_decay=1e-4)
w = torch.tensor([1.0, 2.0, 2.0]).cuda()
t0 = time.time(); step = 0; budget = MINUTES * 60; last_snap = t0
for x, y in dl:
x = x.cuda(non_blocking=True).to(memory_format=torch.channels_last); y = y.cuda(non_blocking=True)
frac = min((time.time() - t0) / budget, 1.0)
for g in opt.param_groups:
g["lr"] = PEAK * (0.5 * (1 + math.cos(math.pi * frac))) * min(1, (step + 1) / 200)
with torch.autocast("cuda", dtype=torch.bfloat16):
out = model(x)
loss = F.cross_entropy(out, y, weight=w) + dice_loss(out.float(), y)
opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); step += 1
if step % 100 == 0:
print(f"step {step} {time.time() - t0:.0f}s loss {loss.item():.4f} ips {step * 64 / (time.time() - t0):.0f}", flush=True)
if step % 1000 == 0 or frac >= 1:
torch.save({"state_dict": model.state_dict()}, OUT + ".pt")
if time.time() - last_snap > 3600:
last_snap = time.time()
torch.save({"state_dict": model.state_dict()}, OUT + f"_h{int((last_snap - t0) // 3600)}.pt")
print("snapshot", flush=True)
if frac >= 1:
break
print("TRAIN_DONE", step, flush=True)
if __name__ == "__main__":
main()