"""Anchored UNet fine-tune (as train5) on Synth3 + real EarthView same-place/different-date 'no change' pairs. usage: python train7.py MINUTES SNAP_MIN P_EV OUT_PREFIX """ 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 from train2 import CROP from train3 import Synth3, degrade MINUTES, SNAP, P_EV, P_PAIR = float(sys.argv[1]), float(sys.argv[2]), float(sys.argv[3]), float(sys.argv[4]) OUTP = sys.argv[5] if len(sys.argv) > 5 else "/workspace/unet_ev2_m" class Synth4(Synth3): def __init__(self, n): super().__init__(n) self.ev = np.load("/workspace/prep/ev_imgs.npy", mmap_mode="r") self.evp = np.load("/workspace/prep/ev_pairs.npy", mmap_mode="r") print("earthview singles", self.ev.shape, "real pairs", self.evp.shape, flush=True) def __getitem__(self, idx): rng = np.random.default_rng((idx * 7919 + os.getpid() * 104729 + time.time_ns()) % 2**32) u = rng.random() if u >= P_EV + (P_PAIR if len(self.evp) else 0): return super().__getitem__(idx) size = int(CROP * rng.uniform(0.55, 0.7)) # 1 m -> 0.55-0.7 m/px after resize r, c = rng.integers(0, 384 - size + 1, 2) rs = lambda a: cv2.resize(np.ascontiguousarray(a[r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_CUBIC) if u < P_EV: # single-date tile, two independent 'acquisitions' crop = rs(self.ev[rng.integers(len(self.ev))]) pre, post = degrade(self.photo(crop, rng), rng), degrade(self.photo(crop.copy(), rng), rng) else: # real different-date pair (light extra aug only) pr = self.evp[rng.integers(len(self.evp))] pre, post = rs(pr[0]), rs(pr[1]) if rng.random() < 0.5: pre, post = post, pre if rng.random() < 0.5: pre, post = degrade(pre, rng), degrade(post, rng) if rng.random() < 0.8: # parallax / residual misregistration M = cv2.getRotationMatrix2D((CROP / 2, CROP / 2), rng.uniform(-1.5, 1.5), rng.uniform(0.98, 1.02)) M[:, 2] += rng.uniform(-4, 4, 2) pre = cv2.warpAffine(pre, M, (CROP, CROP), borderMode=cv2.BORDER_REFLECT) k = rng.integers(8) def d4(a): a = np.rot90(a, k % 4) return np.ascontiguousarray(a[:, ::-1] if k >= 4 else a) pre, post = d4(pre), d4(post) z = np.zeros((CROP, CROP), np.uint8); ig = np.full((CROP, CROP), 255, np.int64) t = lambda a: torch.from_numpy(a.transpose(2, 0, 1).copy()) return (t(pre), t(post), torch.zeros(2, CROP, CROP), torch.zeros(2), torch.from_numpy(ig), torch.from_numpy(ig.copy())) MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1).cuda(); STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1).cuda() model = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3) model.load_state_dict(torch.load("/workspace/unet_r18_cd.pt", map_location="cpu", weights_only=True)["state_dict"]) model = model.cuda().to(memory_format=torch.channels_last) anchor = {n: p.detach().clone() for n, p in model.named_parameters()} opt = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0) dl = DataLoader(Synth4(10**7), batch_size=64, num_workers=30, pin_memory=True, persistent_workers=True, prefetch_factor=4) w = torch.tensor([1.0, 2.0, 2.0]).cuda() t0 = time.time(); step = 0; nsnap = 1 for pre, post, lab, pres, sp, sq in dl: pre = (pre.cuda(non_blocking=True).float() / 255 - MEAN) / STD; post = (post.cuda(non_blocking=True).float() / 255 - MEAN) / STD lab = lab.cuda(non_blocking=True) y = torch.zeros(lab.shape[0], lab.shape[2], lab.shape[3], dtype=torch.long, device="cuda") y[lab[:, 1] > 0] = 2; y[lab[:, 0] > 0] = 1 for g in opt.param_groups: g["lr"] = 1e-4 * min(1, (step + 1) / 100) x = torch.cat([pre, post], 1).contiguous(memory_format=torch.channels_last) with torch.autocast("cuda", dtype=torch.bfloat16): out = model(x) out = out.float(); p = out.softmax(1) dl_ = sum(1 - (2 * (p[:, c] * (y == c)).sum() + 1) / (p[:, c].sum() + (y == c).sum() + 1) for c in (1, 2)) / 2 l2sp = sum(((q - anchor[n]) ** 2).sum() for n, q in model.named_parameters()) loss = F.cross_entropy(out, y, weight=w) + dl_ + 1e-3 * l2sp 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} l2sp {l2sp.item():.2f}", flush=True) el = time.time() - t0 if el > nsnap * SNAP * 60: torch.save({"state_dict": model.state_dict()}, f"{OUTP}{int(nsnap * SNAP)}.pt"); nsnap += 1 print("snapshot", int((nsnap - 1) * SNAP), flush=True) if el > MINUTES * 60: break print("TRAIN_DONE", step, flush=True)