fnruha0921's picture
satlas 4-seed EV train11.py
cfbf904 verified
Raw History Blame Contribute Delete
2.42 kB
"""Satlas Siamese (SiamCD) from Satlas Aerial init on Synth4 (EarthView negatives), cosine LR.
usage: python train11.py MINUTES P_EV OUT
"""
import math, sys, time
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
MINUTES, P_EV, OUT = float(sys.argv[1]), float(sys.argv[2]), sys.argv[3]
sys.argv = sys.argv[:1] + ["60", "32"]
from ev_synth import Synth4
from train3 import dice
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)
dl = DataLoader(Synth4(10**7, P_EV), batch_size=32, num_workers=17, 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; snapped = False
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"] * 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 = 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)
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} ips {step * 32 / (time.time() - t0):.0f}", flush=True)
if frac >= 0.75 and not snapped:
torch.save({"state_dict": model.state_dict()}, OUT + "_m75.pt"); snapped = True
if frac >= 1:
break
torch.save({"state_dict": model.state_dict()}, OUT + ".pt")
print("TRAIN_DONE", step, flush=True)