v4:Two-stage cascade (WNet)

This commit is contained in:
2026-06-16 17:47:20 +08:00
parent daf3d68247
commit 11fea3d054
13 changed files with 154 additions and 55 deletions

View File

@@ -1,14 +1,16 @@
"""
export_onnx.py — export a trained U-Net to ONNX, auto-detecting channels (v1=2, v2=3).
export_onnx.py — export a trained model to ONNX, auto-detecting config from the checkpoint.
python src\\export_onnx.py --ckpt checkpoints\\best.pt --out web\\radio_unet_v1.onnx
python src\\export_onnx.py --ckpt checkpoints_v2\\best.pt --out web\\radio_unet_v2.onnx
python src\\export_onnx.py --ckpt checkpoints\\best.pt --out web\\radio_unet_v1.onnx
python src\\export_onnx.py --ckpt checkpoints_v2\\best.pt --out web\\radio_unet_v2.onnx
python src\\export_onnx.py --ckpt checkpoints_v3\\best.pt --out web\\radio_unet_v3.onnx
python src\\export_onnx.py --ckpt checkpoints_v4\\best.pt --out web\\radio_unet_v4.onnx
"""
import argparse
import numpy as np
import torch
from model import UNet
from model import UNet, WNet
from dataset import RadioMapSeerDataset
@@ -22,12 +24,16 @@ def main():
cargs = ckpt.get("args", {})
base = cargs.get("base", 64)
add_dist = bool(cargs.get("distance", False))
up_mode = cargs.get("up_mode", "deconv")
arch = cargs.get("arch", "unet")
in_ch = 3 if add_dist else 2
model = UNet(in_channels=in_ch, out_channels=1, base=base)
Net = WNet if arch == "wnet" else UNet
model = Net(in_channels=in_ch, out_channels=1, base=base, up_mode=up_mode)
model.load_state_dict(ckpt["model_state"])
model.eval()
print(f"Loaded {args.ckpt} | in_channels={in_ch} | distance={add_dist}")
print(f"Loaded {args.ckpt} | arch={arch} | in_channels={in_ch} | "
f"distance={add_dist} | up_mode={up_mode}")
dummy = torch.randn(1, in_ch, 256, 256)
torch.onnx.export(model, dummy, args.out,

View File

@@ -89,17 +89,44 @@ class UNet(nn.Module):
x = self.up4(x, x1)
return torch.sigmoid(self.outc(x))
class WNet(nn.Module):
"""Two cascaded U-Nets (RadioUNet-style): a coarse predictor + a refiner.
unet1: input channels -> coarse coverage map (1 ch)
unet2: input channels + coarse -> refined coverage map (1 ch)
The refiner sees a global coverage estimate as an extra input channel,
so every pixel has long-range context -> larger effective receptive field.
forward() returns the refined map by default (so eval/ONNX stay single-output);
pass return_coarse=True during training for deep supervision.
"""
def __init__(self, in_channels=3, out_channels=1, base=64, up_mode="resize"):
super().__init__()
self.unet1 = UNet(in_channels, out_channels, base=base, up_mode=up_mode)
self.unet2 = UNet(in_channels + 1, out_channels, base=base, up_mode=up_mode)
def forward(self, x, return_coarse=False):
coarse = self.unet1(x)
refined = self.unet2(torch.cat([x, coarse], dim=1))
if return_coarse:
return refined, coarse
return refined
if __name__ == "__main__":
device = "cuda" if torch.cuda.is_available() else "cpu"
model = UNet(in_channels=2, out_channels=1, base=64).to(device)
n_params = sum(p.numel() for p in model.parameters())
print(f"Device: {device}")
print(f"Parameters: {n_params:,}")
x = torch.randn(4, 2, 256, 256, device=device)
y = model(x)
print("Input :", tuple(x.shape))
print("Output:", tuple(y.shape), "| range:", round(float(y.min()), 3), "-", round(float(y.max()), 3))
unet = UNet(in_channels=3, out_channels=1, base=64, up_mode="resize").to(device)
print(f"UNet params: {sum(p.numel() for p in unet.parameters()):,}")
wnet = WNet(in_channels=3, out_channels=1, base=64, up_mode="resize").to(device)
print(f"WNet params: {sum(p.numel() for p in wnet.parameters()):,}")
x = torch.randn(2, 3, 256, 256, device=device)
refined, coarse = wnet(x, return_coarse=True)
print("Input :", tuple(x.shape))
print("Refined:", tuple(refined.shape),
"| range:", round(float(refined.min()), 3), "-", round(float(refined.max()), 3))
print("Coarse :", tuple(coarse.shape))
if device == "cuda":
print(f"Peak VRAM (forward, batch 4): {torch.cuda.max_memory_allocated()/1e9:.2f} GB")
print(f"Peak VRAM (WNet forward, batch 2): {torch.cuda.max_memory_allocated()/1e9:.2f} GB")

View File

@@ -1,14 +1,14 @@
"""
train.py — train the U-Net on RadioMapSeer coverage maps.
train.py — train the radio-coverage models on RadioMapSeer.
v1: 2-channel input (buildings, transmitter).
v2: add --distance for a 3rd channel (distance-to-Tx) that fixes the square cutoff.
v1: 2-channel UNet. v2: --distance (3-channel). v3: --up-mode resize.
v4: --arch wnet (two cascaded U-Nets: coarse + refine, deep supervision).
Examples:
# v1
python src\\train.py --epochs 30 --batch-size 32 --num-workers 8
# v2 (separate checkpoint dir so v1 stays intact)
python src\\train.py --distance --epochs 30 --batch-size 32 --num-workers 8 --ckpt-dir checkpoints_v2
# v3 (single U-Net, resize-conv)
python src\\train.py --distance --up-mode resize --epochs 30 --batch-size 32 --num-workers 8 --ckpt-dir checkpoints_v3
# v4 (WNet), warm-starting stage 1 from the trained v3
python src\\train.py --arch wnet --distance --up-mode resize --init-unet1-from checkpoints_v3\\best.pt --epochs 25 --batch-size 16 --num-workers 8 --ckpt-dir checkpoints_v4
"""
import argparse
@@ -21,12 +21,13 @@ import torch.nn as nn
from torch.utils.data import DataLoader
from dataset import RadioMapSeerDataset
from model import UNet
from model import UNet, WNet
torch.backends.cudnn.benchmark = True # autotune convs (input size is fixed)
def evaluate(net, loader, criterion, device, amp):
"""Validation MSE on the final output (WNet returns the refined map by default)."""
net.eval()
total, n = 0.0, 0
with torch.no_grad():
@@ -50,16 +51,25 @@ def main():
p.add_argument("--num-workers", type=int, default=4)
p.add_argument("--ckpt-dir", default="checkpoints")
p.add_argument("--distance", action="store_true",
help="v2: add a distance-to-transmitter input channel")
help="v2+: add a distance-to-transmitter input channel")
p.add_argument("--compile", action="store_true",
help="wrap model in torch.compile (optional; flaky on Windows)")
p.add_argument("--up-mode", default="deconv", choices=["deconv", "resize"],
help="v3: 'resize' = bilinear upsample + conv (removes checkerboard)")
help="v3+: 'resize' = bilinear upsample + conv (removes checkerboard)")
p.add_argument("--arch", default="unet", choices=["unet", "wnet"],
help="v4: 'wnet' = two cascaded U-Nets (coarse + refine)")
p.add_argument("--aux-weight", type=float, default=0.4,
help="WNet deep-supervision weight on the coarse output")
p.add_argument("--init-from", default=None,
help="warm-start: load full model weights from a checkpoint")
p.add_argument("--init-unet1-from", default=None,
help="WNet: warm-start stage-1 (unet1) from a single-UNet checkpoint (e.g. v3)")
args = p.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
amp = (device == "cuda")
print(f"Device: {device} | AMP: {amp} | distance: {args.distance} | compile: {args.compile}")
print(f"Device: {device} | AMP: {amp} | arch: {args.arch} | distance: {args.distance} | "
f"up_mode: {args.up_mode} | compile: {args.compile}")
train_ds = RadioMapSeerDataset(split="train", img_size=args.img_size,
max_samples=args.max_samples, add_distance=args.distance)
@@ -74,8 +84,26 @@ def main():
val_loader = DataLoader(val_ds, shuffle=False, **dl_kwargs)
in_ch = 3 if args.distance else 2
model = UNet(in_channels=in_ch, out_channels=1, base=args.base, up_mode=args.up_mode).to(device)
if args.arch == "wnet":
model = WNet(in_channels=in_ch, out_channels=1, base=args.base, up_mode=args.up_mode).to(device)
else:
model = UNet(in_channels=in_ch, out_channels=1, base=args.base, up_mode=args.up_mode).to(device)
# WNet: warm-start stage 1 from an existing single-UNet checkpoint (same in_ch + up_mode)
if args.arch == "wnet" and args.init_unet1_from:
ck1 = torch.load(args.init_unet1_from, map_location=device, weights_only=False)
model.unet1.load_state_dict(ck1["model_state"])
print(f"Warm-started unet1 from {args.init_unet1_from} "
f"(epoch {ck1.get('epoch','?')}, val {ck1.get('val_loss','?')})")
net = torch.compile(model) if args.compile else model
if args.init_from:
ckpt0 = torch.load(args.init_from, map_location=device, weights_only=False)
model.load_state_dict(ckpt0["model_state"])
print(f"Warm-started from {args.init_from} "
f"(epoch {ckpt0.get('epoch','?')}, val {ckpt0.get('val_loss','?')})")
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
scaler = torch.amp.GradScaler(enabled=amp)
@@ -96,20 +124,26 @@ def main():
x, y = x.to(device), y.to(device)
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device, dtype=torch.float16, enabled=amp):
loss = criterion(net(x), y)
if args.arch == "wnet":
refined, coarse = net(x, return_coarse=True)
main_loss = criterion(refined, y) # the real metric
loss = main_loss + args.aux_weight * criterion(coarse, y) # + deep supervision
else:
main_loss = criterion(net(x), y)
loss = main_loss
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
running += loss.item() * x.size(0)
running += main_loss.item() * x.size(0) # log refined MSE (comparable to val)
n += x.size(0)
if i % 50 == 0:
print(f" epoch {epoch} step {i}/{len(train_loader)} | loss {running/n:.5f}")
print(f" epoch {epoch} step {i}/{len(train_loader)} | loss {running/n:.6f}")
train_loss = running / n
val_loss = evaluate(net, val_loader, criterion, device, amp)
dt = time.time() - t0
print(f"[epoch {epoch}/{args.epochs}] train MSE {train_loss:.5f} | "
f"val MSE {val_loss:.5f} | val RMSE {val_loss**0.5:.5f} | {dt:.0f}s")
print(f"[epoch {epoch}/{args.epochs}] train MSE {train_loss:.6f} | "
f"val MSE {val_loss:.6f} | val RMSE {val_loss**0.5:.5f} | {dt:.0f}s")
with open(history_path, "a", newline="") as f:
csv.writer(f).writerow([epoch, f"{train_loss:.6f}", f"{val_loss:.6f}"])
@@ -118,13 +152,13 @@ def main():
torch.save({"epoch": epoch, "model_state": model.state_dict(),
"val_loss": val_loss, "args": vars(args)},
ckpt_dir / "best.pt")
print(f" -> saved best.pt (val MSE {val_loss:.5f})")
print(f" -> saved best.pt (val MSE {val_loss:.6f})")
if device == "cuda":
print(f" peak VRAM: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")
torch.cuda.reset_peak_memory_stats()
print(f"\nDone. Best val MSE: {best_val:.5f} (RMSE {best_val**0.5:.5f})")
print(f"\nDone. Best val MSE: {best_val:.6f} (RMSE {best_val**0.5:.5f})")
if __name__ == "__main__":

View File

@@ -16,7 +16,8 @@ import torch
import matplotlib.pyplot as plt
from dataset import RadioMapSeerDataset
from model import UNet
from model import UNet, WNet
def main():
@@ -36,7 +37,13 @@ def main():
add_dist = bool(cargs.get("distance", False))
up_mode = cargs.get("up_mode", "deconv")
in_ch = 3 if add_dist else 2
model = UNet(in_channels=in_ch, out_channels=1, base=base, up_mode=up_mode).to(device)
arch = cargs.get("arch", "unet")
if arch == "wnet":
model = WNet(in_channels=in_ch, out_channels=1, base=base, up_mode=up_mode).to(device)
else:
arch = cargs.get("arch", "unet")
Net = WNet if arch == "wnet" else UNet
model = Net(in_channels=in_ch, out_channels=1, base=base, up_mode=up_mode).to(device)
model.load_state_dict(ckpt["model_state"])
model.eval()
print(f"Loaded {args.ckpt} | in_channels={in_ch} | distance={add_dist} | "