#!/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.155, 0.145, 0.82, 0.82]) 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()