""" visualize.py — render model predictions vs ground truth on unseen test maps. Loads checkpoints/best.pt and saves comparison panels to outputs/. Usage: python src\\visualize.py python src\\visualize.py --n 6 """ import argparse from pathlib import Path import numpy as np import torch import matplotlib.pyplot as plt from dataset import RadioMapSeerDataset from model import UNet def main(): p = argparse.ArgumentParser() p.add_argument("--ckpt", default="checkpoints/best.pt") p.add_argument("--n", type=int, default=6, help="number of samples to render") p.add_argument("--img-size", type=int, default=256) p.add_argument("--out-dir", default="outputs") p.add_argument("--seed", type=int, default=0) args = p.parse_args() device = "cuda" if torch.cuda.is_available() else "cpu" ckpt = torch.load(args.ckpt, map_location=device, weights_only=False) cargs = ckpt.get("args", {}) base = cargs.get("base", 64) 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) model.load_state_dict(ckpt["model_state"]) model.eval() print(f"Loaded {args.ckpt} | in_channels={in_ch} | distance={add_dist} | " f"val MSE {ckpt.get('val_loss', float('nan')):.5f}") ds = RadioMapSeerDataset(split="test", img_size=args.img_size, add_distance=add_dist) rng = np.random.default_rng(args.seed) idxs = rng.choice(len(ds), size=args.n, replace=False) out_dir = Path(args.out_dir) out_dir.mkdir(exist_ok=True) rmses = [] for k, idx in enumerate(idxs): x, y = ds[int(idx)] with torch.no_grad(): pred = model(x.unsqueeze(0).to(device)).cpu().squeeze().numpy() buildings, antenna = x[0].numpy(), x[1].numpy() truth = y.squeeze().numpy() err = np.abs(pred - truth) rmse = float(np.sqrt(np.mean((pred - truth) ** 2))) rmses.append(rmse) fig, axes = plt.subplots(1, 5, figsize=(18, 4)) panels = [ (buildings, "City map (buildings)", "gray", 1.0), (antenna, "Transmitter", "gray", 1.0), (truth, "Ground-truth coverage","viridis", 1.0), (pred, "Predicted coverage", "viridis", 1.0), (err, "Absolute error", "magma", float(err.max())), ] for ax, (img, title, cmap, vmax) in zip(axes, panels): ax.imshow(img, cmap=cmap, vmin=0, vmax=vmax) ax.set_title(title, fontsize=11) ax.axis("off") fig.suptitle(f"Test sample {idx} — RMSE {rmse:.4f}", fontsize=13) fig.tight_layout() out_path = out_dir / f"prediction_{k:02d}.png" fig.savefig(out_path, dpi=110, bbox_inches="tight") plt.close(fig) print(f" saved {out_path} (RMSE {rmse:.4f})") print(f"\nDone. {args.n} panels in {out_dir}/ | mean RMSE {np.mean(rmses):.4f}") if __name__ == "__main__": main()