File size: 3,591 Bytes
48a2d52 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | """Continue-train a Siamese (SiamCD) or Hybrid (HybridCD) model on Synth4 (EarthView negatives), constant LR,
snapshots every SNAP minutes (for trajectory soups).
usage: python train9.py KIND MINUTES SNAP P_EV OUTP SIAM_INIT [UNET_INIT]
"""
import sys, time
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
KIND, MINUTES, SNAP, P_EV, OUTP, SIAM = sys.argv[1], float(sys.argv[2]), float(sys.argv[3]), float(sys.argv[4]), sys.argv[5], sys.argv[6]
UNET = sys.argv[7] if len(sys.argv) > 7 else None
sys.argv = sys.argv[:1] + ["60", "32"] # train2/train3 parse argv at import time
from ev_synth import Synth4
from train3 import dice
from model_v2 import SiamCD
from model_v3 import HybridCD
model = SiamCD() if KIND == "siam" else HybridCD()
own = model.state_dict()
sd = torch.load(SIAM, map_location="cpu", weights_only=True)["state_dict"]
sd = {k: v for k, v in sd.items() if k in own and own[k].shape == v.shape}
print("loaded", len(sd), "of", len(own), flush=True)
model.load_state_dict(sd, strict=False)
if KIND == "hybrid":
model.unet.load_state_dict(torch.load(UNET, map_location="cpu", weights_only=True)["state_dict"])
model = model.cuda()
if KIND == "siam":
base = {"enc": 3e-5, "new": 1e-4, "rest": 1e-4, "unet": 0}
else:
base = {"enc": 3e-5, "new": 5e-4, "rest": 1e-4, "unet": 2e-5}
groups = {k: [] for k in base}
for n, p in model.named_parameters():
k = "unet" if n.startswith("unet.") else "enc" if n.startswith("enc.") else \
("new" if KIND == "hybrid" and n.startswith(("head_b", "head_t", "pres")) else "rest")
groups[k].append(p)
opt = torch.optim.AdamW([{"params": v, "lr": base[k], "base": base[k]} for k, v in groups.items() if v], weight_decay=1e-4)
dl = DataLoader(Synth4(10**7, P_EV), batch_size=32, num_workers=28, 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; 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)
for g in opt.param_groups:
g["lr"] = g["base"] * min(1, (step + 1) / 100)
with torch.autocast("cuda", dtype=torch.bfloat16):
out = model(pre, post)
ch, pr, sem_p, sem_q = out[0].float(), out[1].float(), out[2], out[3]
loss = F.binary_cross_entropy_with_logits(ch, lab, pos_weight=pw) + dice(ch, lab) \
+ 0.5 * (F.cross_entropy(sem_p.float(), sp, ignore_index=255) + F.cross_entropy(sem_q.float(), sq, ignore_index=255)) \
+ 0.3 * F.binary_cross_entropy_with_logits(pr, pres)
if KIND == "hybrid":
lab3 = torch.zeros_like(sp); lab3[lab[:, 1] > 0] = 2; lab3[lab[:, 0] > 0] = 1
loss = loss + 0.3 * F.cross_entropy(out[4].float(), lab3, weight=torch.tensor([1.0, 2.0, 2.0], device=lab.device))
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} ips {step * 32 / (time.time() - t0):.0f}", 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)
|