Download main_method/code/train_spatial_temporal.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 38 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_spatial_temporal.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/train_spatial_temporal.py
-
curl -L -o train_spatial_temporal.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/train_spatial_temporal.py
38 kB
| """ | |
| T-SRDA finetune v2 — improved recipe for Kaggle (T4 x 2). | |
| The v1 run (official STCLN protocol: fixed 32x32 diagonal crops, no | |
| augmentation, flat LR 1e-4, wd=0) peaked at val mIoU 0.4759 around epoch 37 | |
| and then DECLINED to 0.4207 by epoch 99 — textbook overfitting to the 152 | |
| fixed training crops. v2 attacks exactly that, while keeping everything that | |
| made v1 comparable intact. | |
| NEW in v2 | |
| --------- | |
| 1. RANDOM CROPS anywhere in the 128x128 patch (default 64x64, 2 per patch | |
| per epoch) instead of the two fixed 32x32 diagonal crops (0,0)/(1,1). | |
| Every epoch sees different windows, and 64x64 gives 4x the spatial | |
| context of the official 32x32 plus more parcel boundaries per crop. | |
| 2. AUGMENTATION (v1/official: none): | |
| - dihedral-8 (h/v flip + rot90) always sampled | |
| - temporal subsample (drop up to 20% of frames, order kept) p=0.5 | |
| - frame dropout (zero 1-2 whole frames ~ cloudy acquisition) p=0.25 | |
| - per-band gain/bias jitter + small gaussian noise always | |
| 3. AdamW with weight decay 0.02 on >1-dim params (v1 used wd=0 == Adam). | |
| 4. DISCRIMINATIVE LR: the pretrained encoder trains at 0.5x the head LR, | |
| so 43 h of pretraining is not washed out by the randomly-init heads. | |
| 5. WARMUP (3 ep) + COSINE decay to 1% (v1: flat 1e-4, no schedule). | |
| 6. EMA of weights (decay 0.999) — raw AND EMA are validated every epoch and | |
| each keeps its own best checkpoint. EMA is usually the better one. | |
| 7. Label smoothing 0.05 on the CE terms. | |
| 8. Best-val checkpointing + early stopping (--patience), and a resumable | |
| latest.tar every epoch (Kaggle sessions can die — --resume auto). | |
| KEPT IDENTICAL to v1 (so numbers stay comparable to the 0.4124/0.4433 run) | |
| -------------------------------------------------------------------------- | |
| - model architecture (UTAE_ASPP + UTAEClassificationDual); the pretrained | |
| encoder is loaded strict=True from the .utae.tar checkpoint | |
| - loss family: CE + 0.75*Lovász on the sem & refined heads + 0.5*BCE on | |
| the boundary head; uniform class weights with 0/19 ignored (the | |
| inverse-frequency weighting experiment is NOT repeated — it was a | |
| catastrophic regression in this project) | |
| - splits: the same 76 train / 76 val hardcoded patch IDs, fold-4 test | |
| - index positions (arange(T)), AMP, grad clip 5.0, seed protocol | |
| - same optimizer-step budget: 2 patches/batch x 2 crops = 76 steps/epoch | |
| DDP: launch with torchrun to use both T4s (one run, half the wall time): | |
| torchrun --standalone --nproc_per_node=2 finetune_v2.py --pretrain_pth ... | |
| """ | |
| import argparse | |
| import copy | |
| import json | |
| import math | |
| import os | |
| import random | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.distributed as dist | |
| from torch.utils.data import DataLoader, DistributedSampler | |
| from torch.cuda.amp import autocast, GradScaler | |
| from torch.nn.parallel import DistributedDataParallel as DDP | |
| from sklearn.metrics import confusion_matrix, cohen_kappa_score | |
| import config as C | |
| from dataset import (PASTISPatchDataset, train_patch_ids, val_patch_ids, | |
| pad_collate, crop_ij) | |
| from model import UTAE_ASPP, UTAEClassificationDual, compute_boundary_target | |
| from losses import LovaszSoftmaxLoss | |
| PASTIS_NAMES = C.PASTIS_CLASSES # single source of truth — see config.py | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # Distributed helpers (same contract as v1: torchrun env vars, else single GPU) | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def setup_distributed(): | |
| if "RANK" in os.environ: | |
| dist.init_process_group(backend="nccl") | |
| local_rank = int(os.environ["LOCAL_RANK"]) | |
| rank = int(os.environ["RANK"]) | |
| world = int(os.environ["WORLD_SIZE"]) | |
| torch.cuda.set_device(local_rank) | |
| return rank, world, local_rank, torch.device(f"cuda:{local_rank}") | |
| return 0, 1, 0, torch.device("cuda:0") | |
| def cleanup_distributed(): | |
| if dist.is_initialized(): | |
| dist.destroy_process_group() | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # Metrics — identical to v1 finetune.py | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def compute_all_metrics(preds_np, labels_np, num_classes=20): | |
| class_ids = list(range(1, num_classes - 1)) # 1..18 scored | |
| valid = (labels_np > 0) & (labels_np < num_classes - 1) | |
| pv = preds_np[valid] | |
| lv = labels_np[valid] | |
| if len(lv) == 0: | |
| return dict(miou=0.0, oa=0.0, mf1=0.0, kappa=0.0, | |
| per_iou={}, per_f1={}) | |
| cm = confusion_matrix(lv, pv, labels=class_ids) | |
| diag = np.diag(cm).astype(np.float64) | |
| sum0 = cm.sum(0).astype(np.float64) | |
| sum1 = cm.sum(1).astype(np.float64) | |
| iou = diag / (sum0 + sum1 - diag + 1e-8) | |
| f1 = (2 * diag) / (sum0 + sum1 + 1e-8) | |
| miou = float(iou.mean()) | |
| mf1 = float(f1.mean()) | |
| oa = float((pv == lv).mean()) | |
| kappa = float(cohen_kappa_score(lv, pv)) | |
| return dict(miou=miou, oa=oa, mf1=mf1, kappa=kappa, | |
| per_iou={c: float(iou[i]) for i, c in enumerate(class_ids)}, | |
| per_f1={c: float(f1[i]) for i, c in enumerate(class_ids)}) | |
| def print_metrics(metrics, tag): | |
| print(f" [{tag}] mIoU={metrics['miou']:.4f} OA={metrics['oa']:.4f} " | |
| f"mF1={metrics['mf1']:.4f} Kappa={metrics['kappa']:.4f}", flush=True) | |
| print(f" {'Cls':>4} {'Name':<28} {'IoU':>8} {'F1':>8}", flush=True) | |
| print(f" {'-'*44}", flush=True) | |
| for c in range(1, C.N_CLASSES - 1): | |
| iv = metrics['per_iou'].get(c, 0.0) | |
| fv = metrics['per_f1'].get(c, 0.0) | |
| name = PASTIS_NAMES.get(c, f"cls{c}") | |
| flag = " **DEAD**" if iv < 0.001 else (" *rare*" if iv < 0.10 else "") | |
| print(f" {c:>4} {name:<28} {iv:>8.4f} {fv:>8.4f}{flag}", flush=True) | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # v2 augmentations. All operate on one batch of crops: | |
| # xc (B, T, C, h, w) normalized S2 | |
| # yc (B, h, w) labels | |
| # pos (B, T) index positions (arange(T) upstream) | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def aug_dihedral(xc, yc): | |
| """One dihedral-8 transform per crop-batch, image and label together.""" | |
| if torch.rand(()) > 0.5: | |
| xc, yc = xc.flip(-2), yc.flip(-2) | |
| if torch.rand(()) > 0.5: | |
| xc, yc = xc.flip(-1), yc.flip(-1) | |
| k = int(torch.randint(0, 4, ())) | |
| if k: | |
| xc, yc = torch.rot90(xc, k, (-2, -1)), torch.rot90(yc, k, (-2, -1)) | |
| return xc.contiguous(), yc.contiguous() | |
| def aug_temporal_subsample(xc, pos, p=0.5, drop_max=0.2, keep_min=24): | |
| """Drop up to `drop_max` of the frames, order preserved. Positions keep | |
| their original index values -> emulates missing acquisitions of an | |
| otherwise identical season. No-op with probability 1-p.""" | |
| if torch.rand(()) > p: | |
| return xc, pos | |
| T = xc.shape[1] | |
| keep = int(round(T * (1.0 - random.uniform(0.0, drop_max)))) | |
| keep = max(keep_min, min(T, keep)) | |
| if keep >= T: | |
| return xc, pos | |
| idx = torch.randperm(T)[:keep].sort().values.to(xc.device) | |
| return (xc.index_select(1, idx).contiguous(), | |
| pos.index_select(1, idx).contiguous()) | |
| def aug_frame_dropout(xc, p=0.25, max_frames=2): | |
| """Zero 1-2 whole frames (0 == dataset mean after normalisation) — | |
| emulates cloudy/lost acquisitions.""" | |
| if torch.rand(()) > p: | |
| return xc | |
| T = xc.shape[1] | |
| n = random.randint(1, max_frames) | |
| idx = torch.randperm(T)[:n].to(xc.device) | |
| xc = xc.clone() | |
| xc[:, idx] = 0.0 | |
| return xc | |
| def aug_spectral(xc, gain_std=0.05, bias_std=0.03, noise_std=0.01): | |
| """Per-sample, per-band gain/bias jitter + small gaussian noise (data is | |
| z-normalised, so these are fractions of one std).""" | |
| B, T, Ch, H, W = xc.shape | |
| dev = xc.device | |
| gain = 1.0 + gain_std * torch.randn(B, 1, Ch, 1, 1, device=dev) | |
| bias = bias_std * torch.randn(B, 1, Ch, 1, 1, device=dev) | |
| xc = xc * gain + bias | |
| if noise_std > 0: | |
| xc = xc + noise_std * torch.randn_like(xc) | |
| return xc.contiguous() | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # EMA of the raw (unwrapped) network | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| class ModelEMA: | |
| def __init__(self, model, decay=0.999): | |
| self.decay = decay | |
| self.updates = 0 | |
| raw = model.module if isinstance(model, DDP) else model | |
| self.module = copy.deepcopy(raw).eval() | |
| for p in self.module.parameters(): | |
| p.requires_grad_(False) | |
| def update(self, model): | |
| self.updates += 1 | |
| d = min(self.decay, (1.0 + self.updates) / (10.0 + self.updates)) | |
| raw = model.module if isinstance(model, DDP) else model | |
| msd = raw.state_dict() | |
| for k, v in self.module.state_dict().items(): | |
| if v.dtype.is_floating_point: | |
| v.lerp_(msd[k].detach().float(), 1.0 - d) | |
| else: | |
| v.copy_(msd[k]) | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # Optimizer param groups: encoder vs heads x decay vs no-decay | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def build_param_groups(net_raw, base_lr, encoder_mult, wd): | |
| buckets = {"enc_dec": [], "enc_nd": [], "head_dec": [], "head_nd": []} | |
| for name, p in net_raw.named_parameters(): | |
| if not p.requires_grad: | |
| continue | |
| is_enc = name.startswith("utae.") | |
| no_dec = p.ndim <= 1 # LN/BN weights + all biases | |
| buckets[("enc_" if is_enc else "head_") + ("nd" if no_dec else "dec") | |
| ].append(p) | |
| groups = [ | |
| {"params": buckets["enc_dec"], "lr0": base_lr * encoder_mult, | |
| "weight_decay": wd}, | |
| {"params": buckets["enc_nd"], "lr0": base_lr * encoder_mult, | |
| "weight_decay": 0.0}, | |
| {"params": buckets["head_dec"], "lr0": base_lr, | |
| "weight_decay": wd}, | |
| {"params": buckets["head_nd"], "lr0": base_lr, | |
| "weight_decay": 0.0}, | |
| ] | |
| return [g for g in groups if len(g["params"]) > 0] | |
| def lr_factor(step, total_steps, warmup_steps, min_ratio): | |
| """Linear warmup then cosine decay to min_ratio.""" | |
| if step < warmup_steps: | |
| return (step + 1) / max(1, warmup_steps) | |
| p = (step - warmup_steps) / max(1, total_steps - warmup_steps) | |
| p = min(1.0, p) | |
| cos = 0.5 * (1.0 + math.cos(math.pi * p)) | |
| return min_ratio + (1.0 - min_ratio) * cos | |
| def set_lrs(optimizer, factor): | |
| for g in optimizer.param_groups: | |
| g["lr"] = g["lr0"] * factor | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| # Validation on the two fixed diagonal crops of the val patches | |
| # (grid = PATCH_SIZE / crop_size; for 64 -> (0,0),(1,1) of a 2x2 grid, | |
| # for 32 -> exactly the official v1 validation crops) | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def validate(net, val_dl, device, crop_size, limit_batches=0): | |
| net.eval() | |
| grid = max(1, C.PATCH_SIZE // crop_size) | |
| positions = [(0, 0), (1, 1)] if grid >= 2 else [(0, 0)] | |
| all_preds, all_labels = [], [] | |
| for bi, ((xv, pv_, _dv), yv) in enumerate(val_dl): | |
| if limit_batches and bi >= limit_batches: | |
| break | |
| xv = xv.to(device, non_blocking=True) | |
| pv_ = pv_.to(device, non_blocking=True) | |
| yv = yv.to(device, non_blocking=True) | |
| for (i, j) in positions: | |
| xc, yc = crop_ij(xv, yv, i, j, grid) | |
| with autocast(): | |
| out = net(xc, pv_) | |
| pred = out["refined_logits"].float().argmax(dim=1) | |
| all_preds.append(pred.cpu().numpy().ravel()) | |
| all_labels.append(yc.cpu().numpy().ravel()) | |
| return compute_all_metrics(np.concatenate(all_preds), | |
| np.concatenate(all_labels), C.N_CLASSES) | |
| def pack_checkpoint(state_dict, epoch, metrics, seed, which): | |
| return {"epoch": epoch, | |
| "model_state_dict": state_dict, | |
| "val_miou": metrics["miou"], | |
| "val_oa": metrics["oa"], | |
| "val_mf1": metrics["mf1"], | |
| "val_kappa": metrics["kappa"], | |
| "val_per_iou": metrics["per_iou"], | |
| "seed": seed, "which": which, "recipe": "v2"} | |
| # ───────────────────────────────────────────────────────────────────────────── | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--pretrain_pth", type=str, required=True) | |
| ap.add_argument("--seed", type=int, default=C.SEED) | |
| ap.add_argument("--tag", type=str, default="32X32_SAM_refinement") | |
| ap.add_argument("--out_dir", type=str, default="", | |
| help="default: <EXP_ROOT>/checkpoints/finetune_v2/<tag>") | |
| ap.add_argument("--epochs", type=int, default=C.FT_EPOCHS) | |
| ap.add_argument("--batch", type=int, default=C.FT_BATCH) | |
| # --- the v2 recipe knobs ------------------------------------------------- | |
| ap.add_argument("--crop_size", type=int, default=32, | |
| help="random-crop side; 32 = official crop size") | |
| ap.add_argument("--crops_per_patch", type=int, default=2, | |
| help="random crops per patch per epoch (2 keeps the " | |
| "official 76 steps/epoch)") | |
| ap.add_argument("--official_crops", action="store_true", | |
| help="control run: fixed 32x32 diagonal crops, no augs " | |
| "(= v1 recipe with v2 plumbing)") | |
| ap.add_argument("--fixed_crops", action="store_true", default=True, | |
| help="fixed diagonal crops at --crop_size; keeps augmentation settings") | |
| ap.add_argument("--freeze_encoder_epochs", type=int, default=0) | |
| ap.add_argument("--no_aug", action="store_true") | |
| ap.add_argument("--no_temporal_aug", action="store_true") | |
| ap.add_argument("--mask_ce_invalid", action="store_true") | |
| ap.add_argument("--boundary_weight", type=float, default=0.5) | |
| ap.add_argument("--init_finetuned", type=str, default="", | |
| help="initialize all weights from fine-tuned checkpoint; reset optimizer and schedule") | |
| ap.add_argument("--base_lr", type=float, default=1e-4) | |
| ap.add_argument("--encoder_lr_mult", type=float, default=0.5) | |
| ap.add_argument("--min_lr_ratio", type=float, default=0.01) | |
| ap.add_argument("--warmup_epochs", type=float, default=3.0) | |
| ap.add_argument("--wd", type=float, default=0.02) | |
| ap.add_argument("--label_smoothing", type=float, default=0.05) | |
| ap.add_argument("--ema_decay", type=float, default=0.999) | |
| ap.add_argument("--patience", type=int, default=30, | |
| help="early stopping on val mIoU; 0 = run all epochs") | |
| # --- housekeeping -------------------------------------------------------- | |
| ap.add_argument("--print_every", type=int, default=5) | |
| ap.add_argument("--limit_batches", type=int, default=0, | |
| help=">0: truncate train/val loops (smoke test)") | |
| ap.add_argument("--resume", type=str, default="auto", | |
| help="'auto' resumes from <out_dir>/latest.tar if present; " | |
| "'none' starts fresh; or an explicit .tar path") | |
| ap.add_argument('--research', choices=['l2sp','sam','contrast','consistency','swa'], required=True) | |
| args = ap.parse_args() | |
| if args.research == 'swa' and args.resume != 'none': | |
| ap.error('SWA requires --resume none until averaging state supports resume') | |
| if args.official_crops: | |
| args.crop_size = 32 | |
| args.no_aug = True | |
| if args.fixed_crops and args.crop_size not in (32, 64): | |
| ap.error("--fixed_crops requires --crop_size 32 or 64") | |
| out_dir = (Path(args.out_dir) if args.out_dir else | |
| C.EXP_ROOT / "main_method" / "runs" / args.tag) | |
| rank, world, local_rank, device = setup_distributed() | |
| # Seed BEFORE model creation so all ranks build identical weights... | |
| torch.backends.cuda.matmul.allow_tf32 = True | |
| torch.backends.cudnn.allow_tf32 = True | |
| torch.set_num_threads(min(8, os.cpu_count() or 1)) | |
| torch.manual_seed(args.seed) | |
| torch.cuda.manual_seed_all(args.seed) | |
| train_ids = train_patch_ids() | |
| val_ids = val_patch_ids() | |
| train_ds = PASTISPatchDataset(train_ids) | |
| val_ds = PASTISPatchDataset(val_ids) | |
| sampler = (DistributedSampler(train_ds, num_replicas=world, rank=rank, | |
| shuffle=True) | |
| if world > 1 else None) | |
| train_dl = DataLoader( | |
| train_ds, batch_size=args.batch, sampler=sampler, | |
| shuffle=(sampler is None), num_workers=2, pin_memory=True, | |
| persistent_workers=True, collate_fn=pad_collate) | |
| val_dl = DataLoader( | |
| val_ds, batch_size=args.batch, shuffle=False, | |
| num_workers=2, pin_memory=True, collate_fn=pad_collate) | |
| n_crops = len(C.FT_CROP_IJ) if (args.official_crops or args.fixed_crops) else args.crops_per_patch | |
| steps_per_epoch = len(train_dl) * n_crops | |
| total_steps = steps_per_epoch * args.epochs | |
| warmup_steps = int(args.warmup_epochs * steps_per_epoch) | |
| if rank == 0: | |
| print(f"[FT-v2] train={len(train_ids)} patches ({len(set(train_ids))} unique) " | |
| f"val={len(val_ids)} patches ({len(set(val_ids))} unique)", flush=True) | |
| print(f"[FT-v2] world={world} batch={args.batch} patches x {n_crops} crops " | |
| f"-> {steps_per_epoch} steps/epoch/rank", flush=True) | |
| print(f"[FT-v2] crop_size={args.crop_size} " | |
| f"({'fixed diagonal' if (args.official_crops or args.fixed_crops) else 'RANDOM'}) " | |
| f"aug={'off' if args.no_aug else 'dihedral+temporal+spectral'}", | |
| flush=True) | |
| print(f"[FT-v2] base_lr={args.base_lr} enc_mult={args.encoder_lr_mult} " | |
| f"wd={args.wd} ls={args.label_smoothing} ema={args.ema_decay} " | |
| f"warmup_ep={args.warmup_epochs} min_lr_ratio={args.min_lr_ratio} " | |
| f"patience={args.patience}", flush=True) | |
| print(f"[FT-v2] seed={args.seed} epochs={args.epochs} " | |
| f"checkpoints -> {out_dir}", flush=True) | |
| # ---- model + pretrained encoder ---------------------------------------- | |
| encoder = UTAE_ASPP( | |
| n_channels=C.N_CHANNELS, d_model=C.D_MODEL, n_heads=C.N_HEADS, | |
| T_pad=C.T_PAD, windows=C.TSE_WINDOWS, shifts=C.TSE_SHIFTS) | |
| if rank == 0: | |
| print(f"[FT-v2] loading encoder: {args.pretrain_pth}", flush=True) | |
| ckpt = torch.load(args.pretrain_pth, map_location="cpu", weights_only=False) | |
| sd = ckpt["model_state_dict"] | |
| if not any(k.startswith("spatial_encoder.") for k in sd): | |
| # full pretrain graph (latest.tar) -> strip the "utae." prefix | |
| sd = {k[len("utae."):]: v for k, v in sd.items() if k.startswith("utae.")} | |
| if rank == 0: | |
| print("[FT-v2] (extracted encoder weights from full pretrain " | |
| "checkpoint)", flush=True) | |
| encoder.load_state_dict(sd, strict=True) | |
| anchor = {n: p.detach().clone().to(device) for n,p in encoder.named_parameters()} | |
| network = UTAEClassificationDual( | |
| encoder, d_model=C.D_MODEL, num_classes=C.N_CLASSES).to(device) | |
| net_ref = network | |
| if world > 1: | |
| network = DDP(network, device_ids=[local_rank]) | |
| net_ref = network.module | |
| # ...and only AFTER model construction diverge the per-rank RNGs used for | |
| # shuffling/augmentation, so the two GPUs see different crops/transforms. | |
| torch.manual_seed(args.seed + 1000 * rank) | |
| random.seed(args.seed + 1000 * rank) | |
| if args.init_finetuned: | |
| initial = torch.load(args.init_finetuned, map_location="cpu", weights_only=False) | |
| net_ref.load_state_dict(initial["model_state_dict"], strict=True) | |
| if rank == 0: | |
| print(f"[INIT] fine-tuned weights from {args.init_finetuned}; optimizer/schedule start fresh", flush=True) | |
| # ---- losses: uniform weights (inv-freq was a catastrophic regression) --- | |
| class_weights = torch.ones(C.N_CLASSES, device=device) | |
| class_weights[0] = 0.0 | |
| class_weights[C.N_CLASSES - 1] = 0.0 | |
| ce_loss = nn.CrossEntropyLoss(weight=class_weights, | |
| label_smoothing=args.label_smoothing) | |
| if args.mask_ce_invalid: | |
| def masked_ce(logits, target): | |
| valid = (target > 0) & (target < C.N_CLASSES - 1) | |
| if not valid.any(): | |
| return logits.sum() * 0.0 | |
| return torch.nn.functional.cross_entropy( | |
| logits.permute(0, 2, 3, 1)[valid], target[valid], | |
| weight=class_weights, label_smoothing=args.label_smoothing) | |
| ce_loss = masked_ce | |
| lovasz_loss = LovaszSoftmaxLoss(ignore_indices=[0, C.N_CLASSES - 1]) | |
| bce_loss = nn.BCEWithLogitsLoss() | |
| # ---- optimizer / EMA ---------------------------------------------------- | |
| optimizer = torch.optim.AdamW( | |
| build_param_groups(net_ref, args.base_lr, args.encoder_lr_mult, | |
| args.wd), | |
| lr=args.base_lr, betas=(0.9, 0.95)) | |
| scaler = GradScaler() | |
| ema = ModelEMA(network, decay=args.ema_decay) | |
| from research_methods import pixel_contrast | |
| features = {} | |
| if args.research == 'contrast': | |
| net_ref.sta.register_forward_hook(lambda module, inputs, output: features.update(value=output)) | |
| swa = torch.optim.swa_utils.AveragedModel(net_ref) if args.research == 'swa' else None | |
| best_swa = -1.0 | |
| def training_loss(x, positions, target, boundary): | |
| with autocast(dtype=torch.bfloat16 if args.research == 'contrast' else torch.float16): | |
| outputs = network(x, positions) | |
| sem = outputs['sem_logits'].float() | |
| refined = outputs['refined_logits'].float() | |
| loss = (ce_loss(sem,target) + .75*lovasz_loss(sem,target) | |
| + ce_loss(refined,target) + .75*lovasz_loss(refined,target) | |
| + args.boundary_weight*bce_loss(outputs['bnd_logits'].float().squeeze(1),boundary)) | |
| if args.research == 'l2sp': | |
| # Sum of squared distance, encoder only; original pretrained anchor. | |
| loss = loss + .001*sum((p-anchor[n]).square().sum() for n,p in net_ref.utae.named_parameters()) | |
| elif args.research == 'contrast': | |
| loss = loss + .05*pixel_contrast(features['value'], target) | |
| elif args.research == 'consistency' and epoch >= 3: | |
| with torch.no_grad(), autocast(): | |
| teacher = ema.module(aug_spectral(x, gain_std=.02, bias_std=.01, noise_std=.005),positions)['refined_logits'].float().softmax(1) | |
| valid = (target>0)&(target<C.N_CLASSES-1)&(teacher.max(1).values>=.8) | |
| if valid.any(): | |
| kl = torch.nn.functional.kl_div(refined.log_softmax(1),teacher,reduction='none').sum(1) | |
| loss = loss + .1*min(1.,(epoch-2)/5)*kl[valid].mean() | |
| return loss | |
| if rank == 0: | |
| n_params = sum(p.numel() for p in network.parameters() if p.requires_grad) | |
| n_enc = sum(p.numel() for n_, p in net_ref.named_parameters() | |
| if n_.startswith("utae.")) | |
| print(f"[FT-v2] params={n_params/1e6:.2f}M (encoder {n_enc/1e6:.2f}M)", | |
| flush=True) | |
| # ---- resume -------------------------------------------------------------- | |
| os.makedirs(out_dir, exist_ok=True) | |
| latest_path = out_dir / "latest.tar" | |
| start_epoch, best_raw, best_ema = 0, {"miou": -1.0}, {"miou": -1.0} | |
| best_monitor, patience_counter, history = -1.0, 0, [] | |
| resume_from = None | |
| if args.resume == "auto": | |
| resume_from = latest_path if latest_path.is_file() else None | |
| elif args.resume not in ("none", ""): | |
| resume_from = Path(args.resume) | |
| if not resume_from.is_file(): | |
| raise FileNotFoundError(f"--resume: no such file {resume_from}") | |
| if resume_from is not None: | |
| ck = torch.load(resume_from, map_location="cpu", weights_only=False) | |
| net_ref.load_state_dict(ck["model_state_dict"]) | |
| ema.module.load_state_dict(ck["ema_state_dict"]) | |
| ema.updates = ck.get("ema_updates", 0) | |
| optimizer.load_state_dict(ck["optimizer_state_dict"]) | |
| scaler.load_state_dict(ck["scaler_state_dict"]) | |
| best_raw = ck["best_raw"]; best_ema = ck["best_ema"] | |
| best_monitor = ck.get("best_monitor", -1.0) | |
| patience_counter = ck.get("patience_counter", 0) | |
| history = ck.get("history", []) | |
| start_epoch = int(ck["epoch"]) + 1 | |
| if rank == 0: | |
| print(f"[FT-v2] RESUMED from {resume_from} (finished epoch " | |
| f"{ck['epoch']}) -> epoch {start_epoch}/{args.epochs}", | |
| flush=True) | |
| if start_epoch >= args.epochs: | |
| print("[FT-v2] already complete; nothing to do.", flush=True) | |
| cleanup_distributed() | |
| return | |
| elif rank == 0: | |
| print(f"[FT-v2] starting fresh (--resume={args.resume})", flush=True) | |
| if rank == 0: | |
| print(f"[HYPOTHESIS] temporal_aug={not args.no_temporal_aug} " | |
| f"mask_ce_invalid={args.mask_ce_invalid} boundary_weight={args.boundary_weight} " | |
| f"ema_decay={args.ema_decay}", flush=True) | |
| if args.init_finetuned and start_epoch == 0: | |
| m_init = validate(net_ref, val_dl, device, args.crop_size, args.limit_batches) | |
| if rank == 0: | |
| with open(out_dir / "initial_metrics.json", "w") as f: | |
| json.dump(m_init, f, indent=2) | |
| print(f"[INITIAL] mIoU={m_init['miou']:.6f}", flush=True) | |
| # ---- training loop ------------------------------------------------------- | |
| gstep = start_epoch * steps_per_epoch | |
| hi = C.PATCH_SIZE - args.crop_size | |
| for epoch in range(start_epoch, args.epochs): | |
| if sampler is not None: | |
| sampler.set_epoch(epoch) | |
| network.train() | |
| frozen = epoch < args.freeze_encoder_epochs | |
| for parameter in net_ref.utae.parameters(): | |
| parameter.requires_grad_(not frozen) | |
| if frozen: | |
| net_ref.utae.eval() # freeze BatchNorm statistics and encoder dropout | |
| if rank == 0: | |
| print(f"[ENCODER] frozen={frozen}", flush=True) | |
| ep_start = time.time() | |
| ep_loss, n_steps = 0.0, 0 | |
| stop_flag = False | |
| for it, ((xp, pos, _days), yp) in enumerate(train_dl): | |
| if args.limit_batches and it >= args.limit_batches: | |
| break | |
| xp = xp.to(device, non_blocking=True) # (B,T,10,128,128) | |
| pos = pos.to(device, non_blocking=True) | |
| yp = yp.to(device, non_blocking=True) # (B,128,128) | |
| if args.official_crops or args.fixed_crops: | |
| crop_list = [crop_ij(xp, yp, i, j, C.PATCH_SIZE // args.crop_size) | |
| for (i, j) in [tuple(x) for x in C.FT_CROP_IJ]] | |
| else: | |
| crop_list = [] | |
| for _ in range(n_crops): | |
| i = int(torch.randint(0, hi + 1, ())) | |
| j = int(torch.randint(0, hi + 1, ())) | |
| crop_list.append( | |
| (xp[:, :, :, i:i + args.crop_size, | |
| j:j + args.crop_size], | |
| yp[:, i:i + args.crop_size, j:j + args.crop_size])) | |
| for xc, yc in crop_list: | |
| pos_c = pos | |
| if not args.no_aug: | |
| xc, yc = aug_dihedral(xc, yc) | |
| if not args.no_temporal_aug: | |
| xc, pos_c = aug_temporal_subsample(xc, pos_c) | |
| xc = aug_frame_dropout(xc) | |
| xc = aug_spectral(xc) | |
| bnd_gt = compute_boundary_target(yc) | |
| optimizer.zero_grad(set_to_none=True) | |
| loss = training_loss(xc,pos_c,yc,bnd_gt) | |
| if not torch.isfinite(loss): | |
| raise RuntimeError('Nonfinite research loss') | |
| scaler.scale(loss).backward() | |
| if args.research == 'sam': | |
| parameters = [p for p in net_ref.parameters() if p.grad is not None] | |
| norm = torch.linalg.vector_norm(torch.stack([p.grad.float().norm()/scaler.get_scale() for p in parameters])) | |
| if not torch.isfinite(norm): | |
| # AMP overflow: let GradScaler skip this update and lower | |
| # its scale before the next crop. No SAM perturbation made. | |
| scaler.unscale_(optimizer) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| optimizer.zero_grad(set_to_none=True) | |
| print('[SAM AMP] skipped overflowing first pass; scale=', scaler.get_scale(), flush=True) | |
| continue | |
| offsets = [] | |
| with torch.no_grad(): | |
| for parameter in parameters: | |
| offset = parameter.grad / scaler.get_scale() * (.02/(norm+1e-12)) | |
| parameter.add_(offset) | |
| offsets.append(offset) | |
| optimizer.zero_grad(set_to_none=True) | |
| bn = [m for m in net_ref.modules() if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)] | |
| states = [m.training for m in bn] | |
| for m in bn: m.eval() | |
| try: | |
| second_loss = training_loss(xc,pos_c,yc,bnd_gt) | |
| scaler.scale(second_loss).backward() | |
| finally: | |
| with torch.no_grad(): | |
| for parameter,offset in zip(parameters,offsets): parameter.sub_(offset) | |
| for m,state in zip(bn,states): m.train(state) | |
| scaler.unscale_(optimizer) | |
| torch.nn.utils.clip_grad_norm_(network.parameters(),max_norm=5.0) | |
| factor = 1.0 if args.research == 'swa' else lr_factor(gstep,total_steps,warmup_steps,args.min_lr_ratio) | |
| set_lrs(optimizer,factor) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| ema.update(network) | |
| gstep += 1 | |
| ep_loss += loss.item() | |
| n_steps += 1 | |
| if rank == 0 and it % 10 == 0: | |
| print(f" ep {epoch:3d} batch {it:3d}/{len(train_dl)} " | |
| f"step {n_steps:3d}/{steps_per_epoch} " | |
| f"loss={loss.item():.4f} " | |
| f"lr={optimizer.param_groups[-1]['lr']:.2e} " | |
| f"t={time.time()-ep_start:.0f}s", flush=True) | |
| if swa is not None and epoch >= 5: | |
| swa.update_parameters(net_ref) | |
| # Calibrate BN only on the same two fixed training crops. Dropout off. | |
| swa.module.eval() | |
| bn = [m for m in swa.module.modules() if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)] | |
| momenta = [m.momentum for m in bn] | |
| for m in bn: | |
| m.reset_running_stats(); m.momentum = None; m.train() | |
| with torch.no_grad(): | |
| for bi, ((xx,pp,_), yy) in enumerate(train_dl): | |
| if args.limit_batches and bi >= args.limit_batches: break | |
| xx,pp,yy = xx.to(device),pp.to(device),yy.to(device) | |
| for i,j in C.FT_CROP_IJ: | |
| crop,_ = crop_ij(xx,yy,i,j,C.PATCH_SIZE//args.crop_size) | |
| with autocast(): swa.module(crop,pp) | |
| for m,momentum in zip(bn,momenta): m.momentum=momentum; m.eval() | |
| # ---- validation: raw + EMA ------------------------------------------ | |
| if rank == 0: | |
| m_raw = validate(net_ref, val_dl, device, args.crop_size, | |
| args.limit_batches) | |
| m_ema = validate(ema.module, val_dl, device, args.crop_size, | |
| args.limit_batches) | |
| m_swa = validate(swa.module,val_dl,device,args.crop_size,args.limit_batches) if swa is not None and epoch >= 5 else None | |
| if m_swa is not None and m_swa['miou'] > best_swa: | |
| best_swa = m_swa['miou'] | |
| torch.save(pack_checkpoint(swa.module.state_dict(),epoch,m_swa,args.seed,'swa'),out_dir/'model_best_swa.tar') | |
| ep_sec = time.time() - ep_start | |
| mean_loss = ep_loss / max(n_steps, 1) | |
| history.append({ | |
| "epoch": epoch, "train_loss": mean_loss, | |
| "lr": optimizer.param_groups[-1]["lr"], | |
| "raw": {k: m_raw[k] for k in ("miou", "oa", "mf1", "kappa")}, | |
| "ema": {k: m_ema[k] for k in ("miou", "oa", "mf1", "kappa")}, | |
| "epoch_sec": ep_sec, "swa": m_swa, "research": args.research, | |
| }) | |
| with open(out_dir / "history.json", "w") as f: | |
| json.dump(history, f, indent=1) | |
| improved = False | |
| if m_raw["miou"] > best_raw["miou"]: | |
| best_raw = {"miou": m_raw["miou"], "epoch": epoch} | |
| torch.save(pack_checkpoint(net_ref.state_dict(), epoch, | |
| m_raw, args.seed, "raw"), | |
| out_dir / "model_best_raw.tar") | |
| improved = True | |
| if m_ema["miou"] > best_ema["miou"]: | |
| best_ema = {"miou": m_ema["miou"], "epoch": epoch} | |
| torch.save(pack_checkpoint(ema.module.state_dict(), epoch, | |
| m_ema, args.seed, "ema"), | |
| out_dir / "model_best_ema.tar") | |
| improved = True | |
| monitor = max(m_raw["miou"], m_ema["miou"]) | |
| if monitor > best_monitor: | |
| best_monitor = monitor | |
| patience_counter = 0 | |
| else: | |
| patience_counter += 1 | |
| print(f"[ep {epoch:3d} DONE] loss={mean_loss:.4f} " | |
| f"raw_mIoU={m_raw['miou']:.4f} ema_mIoU={m_ema['miou']:.4f} " | |
| f"best_raw={best_raw['miou']:.4f}@{best_raw['epoch']} " | |
| f"best_ema={best_ema['miou']:.4f}@{best_ema['epoch']} " | |
| f"patience={patience_counter}/{args.patience} " | |
| f"t={ep_sec:.0f}s", flush=True) | |
| if improved or (epoch % args.print_every == 0): | |
| print_metrics(m_raw, f"ep {epoch} RAW") | |
| print_metrics(m_ema, f"ep {epoch} EMA") | |
| # resumable latest.tar every epoch | |
| tmp = out_dir / "latest.tar.tmp" | |
| torch.save({ | |
| "epoch": epoch, | |
| "model_state_dict": net_ref.state_dict(), | |
| "ema_state_dict": ema.module.state_dict(), | |
| "ema_updates": ema.updates, | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scaler_state_dict": scaler.state_dict(), | |
| "best_raw": best_raw, "best_ema": best_ema, | |
| "best_monitor": best_monitor, | |
| "patience_counter": patience_counter, | |
| "history": history, | |
| "args": vars(args), | |
| }, tmp) | |
| os.replace(tmp, latest_path) | |
| if args.patience > 0 and patience_counter >= args.patience: | |
| print(f"[EARLY STOP] best_raw={best_raw} best_ema={best_ema}", | |
| flush=True) | |
| stop_flag = True | |
| # keep ranks in step and share rank 0's early-stop decision | |
| if world > 1: | |
| flag = torch.tensor([1 if stop_flag else 0], device=device) | |
| dist.broadcast(flag, src=0) | |
| stop_flag = bool(flag.item()) | |
| if stop_flag: | |
| break | |
| if rank == 0: | |
| print(f"\n[FT-v2 COMPLETE] best_raw={best_raw} best_ema={best_ema}", | |
| flush=True) | |
| print(f"[FT-v2] checkpoints in {out_dir}", flush=True) | |
| cleanup_distributed() | |
| if __name__ == "__main__": | |
| main() | |