fnruha0921's picture
HOT+greenhouse train13.py
2b10290 verified
Raw History Blame Contribute Delete
3.97 kB
"""Train on Synth6 (HOT buildings + vinyl greenhouses [+ EarthView negatives]).
KIND unet: from organizer baseline, L2-SP, const LR, snapshots every SNAP min. KIND siam: from Satlas init, cosine (siam_v2 recipe).
usage: python train13.py KIND MINUTES SNAP P_EV OUTP
"""
import math, sys, time
import torch, torch.nn.functional as F
from torch.utils.data import DataLoader
KIND, MINUTES, SNAP, P_EV, OUTP = sys.argv[1], float(sys.argv[2]), float(sys.argv[3]), float(sys.argv[4]), sys.argv[5]
sys.argv = sys.argv[:1] + ["60", "32"]
from synth6 import Synth6
from train3 import dice
if KIND == "unet":
import segmentation_models_pytorch as smp
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()
anchor = {n: p.detach().clone() for n, p in model.named_parameters()}
opt = torch.optim.AdamW([{"params": list(model.parameters()), "base": 1e-4}], lr=1e-4, weight_decay=0)
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()
bs = 64
else:
from model_v2 import SiamCD
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, "base": 1e-4}, {"params": rest, "base": 6e-4}], lr=1e-4, weight_decay=1e-4)
bs = 32
dl = DataLoader(Synth6(10**7, P_EV), batch_size=bs, num_workers=19, pin_memory=True, persistent_workers=True, prefetch_factor=4)
pw = torch.tensor([2.0, 2.0]).cuda().view(1, 2, 1, 1); w3 = 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; 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) / (MINUTES * 60), 1.0)
for g in opt.param_groups:
g["lr"] = g["base"] * (min(1, (step + 1) / 100) if KIND == "unet" else 0.5 * (1 + math.cos(math.pi * frac)) * min(1, (step + 1) / 300))
if KIND == "unet":
y = torch.zeros_like(sp); y[lab[:, 1] > 0] = 2; y[lab[:, 0] > 0] = 1
with torch.autocast("cuda", dtype=torch.bfloat16):
out = model(torch.cat([(pre - MEAN) / STD, (post - MEAN) / STD], 1))
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
l2 = sum(((q - anchor[n]) ** 2).sum() for n, q in model.named_parameters())
loss = F.cross_entropy(out, y, weight=w3) + dl_ + 1e-3 * l2
else:
with torch.autocast("cuda", dtype=torch.bfloat16):
ch, pr, sem_p, sem_q = model(pre, post)
ch, pr = ch.float(), pr.float()
loss = F.binary_cross_entropy_with_logits(ch, lab, pos_weight=pw) + dice(ch, lab) + 0.3 * F.binary_cross_entropy_with_logits(pr, pres) \
+ 0.5 * (F.cross_entropy(sem_p.float(), sp, ignore_index=255) + F.cross_entropy(sem_q.float(), sq, ignore_index=255))
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 % 200 == 0: print(f"step {step} {time.time() - t0:.0f}s loss {loss.item():.4f}", flush=True)
el = time.time() - t0
if KIND == "unet" and el > nsnap * SNAP * 60:
torch.save({"state_dict": model.state_dict()}, f"{OUTP}_m{int(nsnap * SNAP)}.pt"); nsnap += 1; print("snapshot", int((nsnap - 1) * SNAP), flush=True)
if frac >= 1: break
torch.save({"state_dict": model.state_dict()}, OUTP + ".pt")
print("TRAIN_DONE", step, flush=True)