File size: 4,366 Bytes
2b046b1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
74
75
76
77
78
79
80
81
82
83
84
85
"""Anchored UNet fine-tune (as train5) on Synth3 + real EarthView same-place/different-date 'no change' pairs.
usage: python train7.py MINUTES SNAP_MIN P_EV OUT_PREFIX
"""
import os
import sys
import time

import cv2
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 train2 import CROP
from train3 import Synth3, degrade

MINUTES, SNAP, P_EV = float(sys.argv[1]), float(sys.argv[2]), float(sys.argv[3])
OUTP = sys.argv[4] if len(sys.argv) > 4 else "/workspace/unet_ev_m"


class Synth4(Synth3):
    def __init__(self, n):
        super().__init__(n)
        self.ev = np.load("/workspace/prep/ev_imgs.npy", mmap_mode="r")
        print("earthview pairs", self.ev.shape, flush=True)

    def __getitem__(self, idx):
        rng = np.random.default_rng((idx * 7919 + os.getpid() * 104729 + time.time_ns()) % 2**32)
        if rng.random() >= P_EV:
            return super().__getitem__(idx)
        img = self.ev[rng.integers(len(self.ev))]
        H = img.shape[0]
        size = int(CROP * rng.uniform(0.55, 0.7))           # 1 m -> 0.55-0.7 m/px after resize
        r, c = rng.integers(0, H - size + 1, 2)
        crop = cv2.resize(np.ascontiguousarray(img[r:r + size, c:c + size]), (CROP, CROP), interpolation=cv2.INTER_CUBIC)
        pre, post = degrade(self.photo(crop, rng), rng), degrade(self.photo(crop.copy(), rng), rng)
        if rng.random() < 0.8:  # parallax / residual misregistration
            M = cv2.getRotationMatrix2D((CROP / 2, CROP / 2), rng.uniform(-1.5, 1.5), rng.uniform(0.98, 1.02))
            M[:, 2] += rng.uniform(-4, 4, 2)
            pre = cv2.warpAffine(pre, M, (CROP, CROP), borderMode=cv2.BORDER_REFLECT)
        k = rng.integers(8)
        def d4(a):
            a = np.rot90(a, k % 4)
            return np.ascontiguousarray(a[:, ::-1] if k >= 4 else a)
        pre, post = d4(pre), d4(post)
        z = np.zeros((CROP, CROP), np.uint8); ig = np.full((CROP, CROP), 255, np.int64)
        t = lambda a: torch.from_numpy(a.transpose(2, 0, 1).copy())
        return (t(pre), t(post), torch.zeros(2, CROP, CROP), torch.zeros(2), torch.from_numpy(ig), torch.from_numpy(ig.copy()))


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)
model.load_state_dict(torch.load("/workspace/unet_r18_cd.pt", map_location="cpu", weights_only=True)["state_dict"])
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(Synth4(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"{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)