"""Train SiamCD (model_v2) on synthetic pairs with per-image semantics and satellite-like degradation. usage: python train3.py MINUTES [BATCH] """ import math import os import sys import time import cv2 import numpy as np import torch import torch.nn.functional as F from torch.utils.data import DataLoader from model_v2 import SiamCD from train2 import CROP, Synth MINUTES = float(sys.argv[1]) if len(sys.argv) > 1 else 90 BATCH = int(sys.argv[2]) if len(sys.argv) > 2 else 32 OUT = os.environ.get("OUT", "/workspace/siam_v2") def degrade(img, rng): """Make crisp downsampled aerial imagery look like display-stretched pansharpened satellite imagery.""" x = img if rng.random() < 0.8: # resolution loss f = rng.uniform(1.0, 1.5) small = cv2.resize(x, (int(CROP / f), int(CROP / f)), interpolation=cv2.INTER_AREA) x = cv2.resize(small, (CROP, CROP), interpolation=[cv2.INTER_LINEAR, cv2.INTER_CUBIC][rng.integers(2)]) if rng.random() < 0.4: # pansharpening / sharpening halos bl = cv2.GaussianBlur(x, (0, 0), rng.uniform(0.8, 1.6)) x = cv2.addWeighted(x, 1 + rng.uniform(0.3, 1.0), bl, -rng.uniform(0.3, 1.0), 0) if rng.random() < 0.5: # display stretch (percentile clip per image) lo, hi = np.percentile(x, rng.uniform(0.5, 3)), np.percentile(x, 100 - rng.uniform(0.5, 3)) x = np.clip((x.astype(np.float32) - lo) * 255.0 / max(hi - lo, 1), 0, 255).astype(np.uint8) if rng.random() < 0.3: # haze a = rng.uniform(0.05, 0.25) x = (x * (1 - a) + a * rng.uniform(150, 230)).astype(np.uint8) if rng.random() < 0.5: ok, enc = cv2.imencode(".jpg", x, [cv2.IMWRITE_JPEG_QUALITY, int(rng.integers(40, 92))]) x = cv2.imdecode(enc, cv2.IMREAD_UNCHANGED) return x class Synth3(Synth): 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.30 else ("t" if r < 0.60 else ("bt" if r < 0.68 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() sem = np.zeros(y.shape, np.uint8); sem[y == 1] = 1; sem[y == 2] = 2 sem_pre, sem_post = sem.copy(), sem.copy() lb = np.zeros(y.shape, np.uint8); lt = 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: sel = rng.choice(comps, size=rng.integers(1, len(comps) + 1), replace=False) region = np.isin(cc, sel).astype(np.uint8) if rng.random() < 0.2: region = region & self.blob(rng, region.shape, region.sum() * 0.5) 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) sem_pre[grow.astype(bool)] = 0 lb[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) sem_post[region.astype(bool)] = 0 lt[(region & tmask).astype(bool)] = 1 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 mode == "neg" and rng.random() < 0.3: # field / ground appearance change (not a target) g = (y == 3) post = post.astype(np.int16); post[g] += rng.integers(-40, 40, 3).astype(np.int16) post = np.clip(post, 0, 255).astype(np.uint8) if (lb.any() or lt.any()) and rng.random() < 0.2: # demolition / regrowth -> not a target pre, post = post, pre sem_pre, sem_post = sem_post, sem_pre lb[:] = 0; lt[:] = 0 season = sem == 2 pre = degrade(self.photo(pre, rng, tmask if rng.random() < 0.7 else season), rng) post = degrade(self.photo(post, rng, tmask if rng.random() < 0.7 else season), rng) if rng.random() < 0.8: # misregistration / parallax of pre (affine) a = np.deg2rad(rng.uniform(-1.5, 1.5)); s = rng.uniform(0.98, 1.02) M = cv2.getRotationMatrix2D((CROP / 2, CROP / 2), np.rad2deg(a), s) M[:, 2] += rng.uniform(-4, 4, 2) pre = cv2.warpAffine(pre, M, (CROP, CROP), borderMode=cv2.BORDER_REFLECT) sem_pre = cv2.warpAffine(sem_pre, M, (CROP, CROP), flags=cv2.INTER_NEAREST, borderMode=cv2.BORDER_REFLECT) ignore = np.zeros(y.shape, bool) if rng.random() < 0.15: yy, xx = np.mgrid[:CROP, :CROP] for im, sm in ((pre, sem_pre), (post, sem_post)): if rng.random() < 0.6: th = rng.uniform(0, 2 * np.pi); d = rng.uniform(-CROP * 0.7, -CROP * 0.15) nd = (xx - CROP / 2) * np.cos(th) + (yy - CROP / 2) * np.sin(th) < d im[nd] = 0; sm[nd] = 255; ignore |= nd lb[ignore] = 0; lt[ignore] = 0 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, lb, lt, sem_pre, sem_post = map(d4, (pre, post, lb, lt, sem_pre, sem_post)) pres = np.array([lb.sum() >= 20, lt.sum() >= 20], np.float32) t = lambda a: torch.from_numpy(a.transpose(2, 0, 1).copy()) return (t(pre), t(post), torch.from_numpy(np.stack([lb, lt]).astype(np.float32)), torch.from_numpy(pres), torch.from_numpy(sem_pre.astype(np.int64)), torch.from_numpy(sem_post.astype(np.int64))) def dice(logit, y): p = logit.sigmoid() inter = (p * y).sum((0, 2, 3)); den = p.sum((0, 2, 3)) + y.sum((0, 2, 3)) return (1 - (2 * inter + 1) / (den + 1)).mean() def main(): torch.backends.cudnn.benchmark = True model = SiamCD("/workspace/satlas/aerial_swinb_si.pth").cuda() enc = [p for n, p in model.named_parameters() if n.startswith("enc.")] rest = [p for n, p in model.named_parameters() if not n.startswith("enc.")] opt = torch.optim.AdamW([{"params": enc, "lr": 1e-4, "base": 1e-4}, {"params": rest, "lr": 6e-4, "base": 6e-4}], weight_decay=1e-4) dl = DataLoader(Synth3(10**7), batch_size=BATCH, num_workers=int(os.environ.get("WORKERS", 30)), pin_memory=True, persistent_workers=True, prefetch_factor=4) pw = torch.tensor([2.0, 2.0]).cuda().view(1, 2, 1, 1) t0 = time.time(); step = 0; budget = MINUTES * 60; last_snap = t0 for pre, post, lab, pres, sp, sq in dl: pre = pre.cuda(non_blocking=True).float() / 255; post = post.cuda(non_blocking=True).float() / 255 lab, pres, sp, sq = lab.cuda(non_blocking=True), pres.cuda(non_blocking=True), sp.cuda(non_blocking=True), sq.cuda(non_blocking=True) frac = min((time.time() - t0) / budget, 1.0) for g in opt.param_groups: g["lr"] = g["base"] * 0.5 * (1 + math.cos(math.pi * frac)) * min(1, (step + 1) / 300) with torch.autocast("cuda", dtype=torch.bfloat16): ch, pr, sem_p, sem_q = model(pre, post) ch, pr = ch.float(), pr.float() loss_ch = F.binary_cross_entropy_with_logits(ch, lab, pos_weight=pw) + dice(ch, lab) loss_sem = F.cross_entropy(sem_p.float(), sp, ignore_index=255) + F.cross_entropy(sem_q.float(), sq, ignore_index=255) loss_pr = F.binary_cross_entropy_with_logits(pr, pres) loss = loss_ch + 0.5 * loss_sem + 0.3 * loss_pr opt.zero_grad(set_to_none=True); loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step(); step += 1 if step % 100 == 0: print(f"step {step} {time.time() - t0:.0f}s loss {loss.item():.4f} ch {loss_ch.item():.3f} sem {loss_sem.item():.3f} " f"pr {loss_pr.item():.3f} ips {step * BATCH / (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 > 1800: last_snap = time.time() torch.save({"state_dict": model.state_dict()}, OUT + f"_m{int((last_snap - t0) // 60)}.pt") if frac >= 1: break print("TRAIN_DONE", step, flush=True) if __name__ == "__main__": main()