Files
JSAC_AIRAN/c18_mnist_doppler.py
T
KiHoLee 8fe5f499b8 Widen the axes margin and fix the clipped y label of Fig. 8
The BERT study spans less than one decade, so its log axis carried the
wide "6 x 10^-1" tick labels, which pushed the y label off the canvas.
Give that axis plain decimal ticks and widen the axes margin of every
result figure by the same amount, which keeps the axes box identical
across figures at the 4:3 ratio.
2026-08-26 19:09:49 +09:00

189 lines
8.6 KiB
Python
Executable File

#!/usr/bin/env python3
# ============================================================
# c18_mnist_doppler.py
#
# MNIST real-data SER versus normalized Doppler at a fixed SNR,
# reusing the trained checkpoints of c13_mnist.py. The signed
# MAML receiver re-adapts its decoder-side parameters for every
# (SNR, Doppler) task.
# ============================================================
import argparse
import csv
import os
import numpy as np
import torch
import c11_doppler_csi as base
import c13_mnist as m13
@torch.no_grad()
def run_sweep(args, device):
base.set_seed(args.seed + 3)
_, te = m13.get_datasets(args.data_root)
x, y = m13.tensorize(te)
xs_tr, ys_tr = m13.tensorize(m13.get_datasets(args.data_root)[0])
chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device)
X_pilot = base.make_pilot(args.nfft, device)
sem = m13.MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
sem.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_semantic.pt"),
map_location=device))
sem.eval()
mm = m13.MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
mm.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_maml.pt"),
map_location=device))
mm.eval()
tfm = m13.MnistTransformerMA(args.users, args.dim, args.hidden).to(device)
tfm.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_tf.pt"),
map_location=device))
tfm.eval()
aem = m13.MnistPerUserAE(args.users, args.dim, args.hidden).to(device)
aem.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_ae.pt"),
map_location=device))
aem.eval()
txc = m13.TxClassifier(args.dim).to(device)
txc.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_txcls.pt"),
map_location=device))
txc.eval()
snr = args.eval_snr
noise_var = 10 ** (-snr / 10.0)
rng = np.random.default_rng(args.seed + 3)
csv_path = os.path.join(args.save_dir, "mnist_doppler.csv")
with open(csv_path, "w", newline="") as f:
w = csv.writer(f)
w.writerow(["snr_db", "fd_norm", "method", "ser"])
for fd in args.eval_fds:
args.eval_fd = fd
errs = {"digital_genie": 0, "digital_pilot": 0, "semantic": 0,
"semantic_tf": 0, "semantic_ae": 0, "semantic_maml": 0}
fast_mm = m13.adapt_mnist(mm, xs_tr, ys_tr, chan, X_pilot,
args, snr, device)
total = 0
for _ in range(args.eval_nb):
imgs, labels = m13.sample_frames(x, y, args.eval_batch,
args.users, device, rng)
g_p, g_d = chan.sample(args.eval_batch, fd, args.delta)
logits = m13.forward_frames(sem, imgs, chan, g_p, g_d,
X_pilot, noise_var, fd, args.delta)
errs["semantic"] += (logits.argmax(-1) != labels).sum().item()
lg2 = m13.forward_frames(tfm, imgs, chan, g_p, g_d, X_pilot,
noise_var, fd, args.delta)
errs["semantic_tf"] += (lg2.argmax(-1) != labels).sum().item()
lg4 = m13.forward_frames(aem, imgs, chan, g_p, g_d, X_pilot,
noise_var, fd, args.delta)
errs["semantic_ae"] += (lg4.argmax(-1) != labels).sum().item()
lg3 = m13.forward_frames_fast(mm, fast_mm, imgs, chan, g_p,
g_d, X_pilot, noise_var, fd,
args.delta)
errs["semantic_maml"] += (lg3.argmax(-1) != labels).sum().item()
B = imgs.shape[0]
pred_cls = txc(imgs.reshape(B * args.users, 1, 28, 28)).argmax(-1)
pred_cls = pred_cls.view(B, args.users)
X_d = m13.digital_tx(pred_cls, args.users, args.nfft, device)
Y_d = chan.transmit(X_d, g_d, noise_var)
Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(B, -1), g_p,
noise_var)
H_ls = Y_p * torch.conj(X_pilot).unsqueeze(0)
H_true = chan.genie_H(g_d)
for name, H in [("digital_genie", H_true),
("digital_pilot", H_ls)]:
rec = m13.digital_detect(Y_d, H, args.users, args.nfft)
errs[name] += (rec != labels).sum().item()
total += labels.numel()
for name, e in errs.items():
w.writerow([snr, fd, name, e / total])
f.flush()
print(f"fd={fd:7.4f} | " +
" ".join(f"{k}={v / total:.4e}" for k, v in errs.items()),
flush=True)
print("saved", csv_path)
def make_fig(args):
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
rows = list(csv.DictReader(open(os.path.join(args.save_dir,
"mnist_doppler.csv"))))
LAB = {"digital_genie": "Digital chain genie CSI",
"digital_pilot": "Digital chain pilot CSI",
"semantic_tf": "Transformer SE separation",
"semantic_ae": "Per-user AE",
"semantic": "Proposed signed joint",
"semantic_maml": "Proposed signed MAML"}
STY = {"digital_genie": dict(color="gray", marker="^", ls="--"),
"digital_pilot": dict(color="k", marker="v", ls="-"),
"semantic_tf": dict(color="tab:purple", marker="P", ls="-"),
"semantic_ae": dict(color="tab:brown", marker="X", ls="-"),
"semantic": dict(color="tab:red", marker="o", ls="-"),
"semantic_maml": dict(color="tab:green", marker="D", ls="--")}
plt.rcParams.update({"font.size": 13, "axes.labelsize": 13,
"xtick.labelsize": 12, "ytick.labelsize": 12,
"axes.linewidth": 1.1, "grid.linewidth": 0.8,
"xtick.major.width": 1.1, "ytick.major.width": 1.1,
"xtick.minor.width": 0.8, "ytick.minor.width": 0.8,
"xtick.major.size": 4.5, "ytick.major.size": 4.5})
fig = plt.figure(figsize=(5.2, 3.9))
ax = fig.add_axes([0.185, 0.145, 0.79, 0.79])
for mkey in LAB:
pts = sorted([(float(r["fd_norm"]), float(r["ser"]))
for r in rows if r["method"] == mkey])
pts = [(a, b) for a, b in pts if b > 0]
if not pts:
continue
xs, ys = zip(*pts)
ax.semilogy(xs, ys, label=LAB[mkey], ms=5, lw=1.8, **STY[mkey])
ax.set_xlabel(r"Normalized Doppler $f_D T_{\mathrm{sym}}$")
ax.set_ylabel("SER")
ax.grid(True, which="both", alpha=0.35)
ax.legend(fontsize=9, framealpha=1.0, labelspacing=0.3,
handlelength=1.8, loc="center right", bbox_to_anchor=(0.985, 0.72))
out = os.path.join(args.fig_dir,
f"mnist_ser_vs_doppler_snr{int(args.eval_snr)}.pdf")
fig.savefig(out)
print("saved", out)
def main():
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["run", "fig"], required=True)
p.add_argument("--users", type=int, default=8)
p.add_argument("--dim", type=int, default=128)
p.add_argument("--hidden", type=int, default=256)
p.add_argument("--nfft", type=int, default=128)
p.add_argument("--cp", type=int, default=16)
p.add_argument("--taps", type=int, default=8)
p.add_argument("--delta", type=int, default=6)
p.add_argument("--eval-snr", type=float, default=15.0)
p.add_argument("--eval-fds", type=float, nargs="+",
default=[0.001, 0.002, 0.005, 0.01, 0.02, 0.03, 0.05,
0.0567, 0.07, 0.085, 0.1])
p.add_argument("--eval-inner-steps", type=int, default=5)
p.add_argument("--eval-inner-lr", type=float, default=0.01)
p.add_argument("--support", type=int, default=32)
p.add_argument("--eval-batch", type=int, default=64)
p.add_argument("--eval-nb", type=int, default=60)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--save-dir", type=str, default="results_mnist")
p.add_argument("--fig-dir", type=str, default="fig")
p.add_argument("--data-root", type=str, default="data_mnist")
args = p.parse_args()
args.eval_fd = args.eval_fds[0]
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device:", device)
if args.mode == "run":
run_sweep(args, device)
else:
make_fig(args)
if __name__ == "__main__":
main()