Download code/train.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/train.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/train.py
-
curl -L -o train.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/train.py
10.7 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 | |
| class Synth(Dataset): | |
| def __init__(self, n): | |
| self.n = n | |
| self.fx = np.load(f"{P}/flair_x.npy", mmap_mode="r"); self.fy = np.load(f"{P}/flair_y.npy", mmap_mode="r") | |
| 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 = np.load(f"{P}/flair_y.npy") | |
| 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 | |
| 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) | |
| 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 ---------- | |
| 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("/workspace/unet_r18_cd.pt", 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 | |
| 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"] = 4e-4 * (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()}, "/workspace/unet_synth.pt") | |
| if frac >= 1: | |
| break | |
| print("TRAIN_DONE", step, flush=True) | |
| if __name__ == "__main__": | |
| main() | |