v2: distance-to-transmitter channel — fixes far-field square cutoff

This commit is contained in:
2026-06-08 07:34:55 +08:00
parent 0749daa2f9
commit d25fd5eaa8
9 changed files with 58 additions and 30 deletions

View File

@@ -31,14 +31,17 @@ def main():
device = "cuda" if torch.cuda.is_available() else "cpu"
ckpt = torch.load(args.ckpt, map_location=device, weights_only=False)
base = ckpt.get("args", {}).get("base", 64)
model = UNet(in_channels=2, out_channels=1, base=base).to(device)
cargs = ckpt.get("args", {})
base = cargs.get("base", 64)
add_dist = bool(cargs.get("distance", False))
in_ch = 3 if add_dist else 2
model = UNet(in_channels=in_ch, out_channels=1, base=base).to(device)
model.load_state_dict(ckpt["model_state"])
model.eval()
print(f"Loaded {args.ckpt} (epoch {ckpt.get('epoch', '?')}, "
f"val MSE {ckpt.get('val_loss', float('nan')):.5f})")
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) # unseen cities
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)