Files
radio-coverage-dl/src/visualize.py

86 lines
3.0 KiB
Python

"""
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()