86 lines
3.0 KiB
Python
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() |