diff --git a/checkpoints_v4/best.pt b/checkpoints_v4/best.pt new file mode 100644 index 0000000..934517f Binary files /dev/null and b/checkpoints_v4/best.pt differ diff --git a/checkpoints_v4/history.csv b/checkpoints_v4/history.csv new file mode 100644 index 0000000..6a0750a --- /dev/null +++ b/checkpoints_v4/history.csv @@ -0,0 +1,21 @@ +epoch,train_mse,val_mse +1,0.000914,0.000824 +2,0.000637,0.000939 +3,0.000570,0.000672 +4,0.000504,0.000629 +5,0.000455,0.000595 +6,0.000416,0.000640 +7,0.000391,0.000564 +8,0.000365,0.000544 +9,0.000344,0.000527 +10,0.000327,0.000515 +11,0.000309,0.000534 +12,0.000297,0.000504 +13,0.000284,0.000517 +14,0.000271,0.000462 +15,0.000263,0.000468 +16,0.000250,0.000450 +17,0.000243,0.000460 +18,0.000234,0.000448 +19,0.000227,0.000446 +20,0.000220,0.000443 diff --git a/outputs_v4/prediction_00.png b/outputs_v4/prediction_00.png new file mode 100644 index 0000000..e634ab6 Binary files /dev/null and b/outputs_v4/prediction_00.png differ diff --git a/outputs_v4/prediction_01.png b/outputs_v4/prediction_01.png new file mode 100644 index 0000000..75774fc Binary files /dev/null and b/outputs_v4/prediction_01.png differ diff --git a/outputs_v4/prediction_02.png b/outputs_v4/prediction_02.png new file mode 100644 index 0000000..3e56dbb Binary files /dev/null and b/outputs_v4/prediction_02.png differ diff --git a/outputs_v4/prediction_03.png b/outputs_v4/prediction_03.png new file mode 100644 index 0000000..c2f0c02 Binary files /dev/null and b/outputs_v4/prediction_03.png differ diff --git a/outputs_v4/prediction_04.png b/outputs_v4/prediction_04.png new file mode 100644 index 0000000..b35ad3e Binary files /dev/null and b/outputs_v4/prediction_04.png differ diff --git a/outputs_v4/prediction_05.png b/outputs_v4/prediction_05.png new file mode 100644 index 0000000..eba665c Binary files /dev/null and b/outputs_v4/prediction_05.png differ diff --git a/src/export_onnx.py b/src/export_onnx.py index a09899b..bc25519 100644 --- a/src/export_onnx.py +++ b/src/export_onnx.py @@ -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, diff --git a/src/model.py b/src/model.py index ee87e63..7b1bb9b 100644 --- a/src/model.py +++ b/src/model.py @@ -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") \ No newline at end of file + print(f"Peak VRAM (WNet forward, batch 2): {torch.cuda.max_memory_allocated()/1e9:.2f} GB") \ No newline at end of file diff --git a/src/train.py b/src/train.py index 3936b77..ac4b685 100644 --- a/src/train.py +++ b/src/train.py @@ -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__": diff --git a/src/visualize.py b/src/visualize.py index f85544b..d6546c2 100644 --- a/src/visualize.py +++ b/src/visualize.py @@ -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} | " diff --git a/web/index.html b/web/index.html index 68a261c..3346790 100644 --- a/web/index.html +++ b/web/index.html @@ -56,7 +56,8 @@