"""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()