File size: 3,721 Bytes
eea5f0e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
66
67
68
69
70
71
72
73
"""Fine-tune HybridCD (Satlas Siamese + baseline-UNet expert) on synthetic pairs.

usage: python train4.py MINUTES SIAM_CKPT UNET_CKPT
"""
import math
import os
import sys
import time

import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader

from model_v3 import HybridCD
from train3 import Synth3, dice

MINUTES = float(sys.argv[1]); SIAM = sys.argv[2]; UNET = sys.argv[3]
OUT = "/workspace/hybrid_v3"


def main():
    torch.backends.cudnn.benchmark = True
    model = 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("siam keys loaded", len(sd), "missing", len([k for k in own if k not in sd and not k.startswith("unet.")]))
    model.load_state_dict(sd, strict=False)
    model.unet.load_state_dict(torch.load(UNET, map_location="cpu", weights_only=True)["state_dict"])
    model = model.cuda()
    groups = {"unet": [], "enc": [], "new": [], "rest": []}
    for n, p in model.named_parameters():
        k = "unet" if n.startswith("unet.") else "enc" if n.startswith("enc.") else \
            "new" if n.startswith(("head_b", "head_t", "pres")) else "rest"
        groups[k].append(p)
    base = {"unet": 3e-5, "enc": 5e-5, "new": 6e-4, "rest": 2e-4}
    opt = torch.optim.AdamW([{"params": v, "lr": base[k], "base": base[k]} for k, v in groups.items()], weight_decay=1e-4)
    dl = DataLoader(Synth3(10**7), batch_size=32, num_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) / 200)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            ch, pr, sem_p, sem_q, ul = model(pre, post)
        ch, pr = ch.float(), pr.float()
        lab3 = torch.zeros_like(sp); lab3[lab[:, 1] > 0] = 2; lab3[lab[:, 0] > 0] = 1
        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_u = F.cross_entropy(ul.float(), lab3, weight=torch.tensor([1.0, 2.0, 2.0], device=ul.device))
        loss = loss_ch + 0.5 * loss_sem + 0.3 * loss_pr + 0.3 * loss_u
        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} u {loss_u.item():.3f} "
                  f"pr {loss_pr.item():.3f} ips {step * 32 / (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 > 1200:
            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()