Download code/train3.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 9.16 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/train3.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/train3.py
-
curl -L -o train3.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/train3.py
9.16 kB
| """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 = "/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() | |