fnruha0921's picture
add code
eea5f0e verified
Raw History Blame Contribute Delete
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()