"""Anchored short fine-tune of the organizer baseline UNet on the improved synthetic pairs (Synth3). L2-SP keeps weights close to the organizer weights; snapshots every SNAP_MIN minutes. usage: python train5.py MINUTES SNAP_MIN """ import math, sys, time import numpy as np import segmentation_models_pytorch as smp import torch import torch.nn.functional as F from torch.utils.data import DataLoader from train3 import Synth3 MINUTES = float(sys.argv[1]); SNAP = float(sys.argv[2]) 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() model = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3) base = torch.load("/workspace/unet_r18_cd.pt", map_location="cpu", weights_only=True)["state_dict"] model.load_state_dict(base); model = model.cuda().to(memory_format=torch.channels_last) anchor = {n: p.detach().clone() for n, p in model.named_parameters()} opt = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0) dl = DataLoader(Synth3(10**7), batch_size=64, num_workers=30, pin_memory=True, persistent_workers=True, prefetch_factor=4) w = 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 - MEAN) / STD; post = (post.cuda(non_blocking=True).float() / 255 - MEAN) / STD lab = lab.cuda(non_blocking=True) y = torch.zeros(lab.shape[0], lab.shape[2], lab.shape[3], dtype=torch.long, device="cuda") y[lab[:, 1] > 0] = 2; y[lab[:, 0] > 0] = 1 for g in opt.param_groups: g["lr"] = 1e-4 * min(1, (step + 1) / 100) x = torch.cat([pre, post], 1).contiguous(memory_format=torch.channels_last) with torch.autocast("cuda", dtype=torch.bfloat16): out = model(x) 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 l2sp = sum(((q - anchor[n]) ** 2).sum() for n, q in model.named_parameters()) loss = F.cross_entropy(out, y, weight=w) + dl_ + 1e-3 * l2sp opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); step += 1 if step % 100 == 0: print(f"step {step} {time.time() - t0:.0f}s loss {loss.item():.4f} l2sp {l2sp.item():.2f}", flush=True) el = time.time() - t0 if el > nsnap * SNAP * 60: torch.save({"state_dict": model.state_dict()}, f"/workspace/unet_anch_m{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)