#!/usr/bin/env python3 # ============================================================ # c19_mnist_epoch.py # # Real-data convergence study: evaluation SER versus training # step for the proposed signed receiver, the Transformer SE # separator, and the per-user AE, all trained on the MNIST # TDL chain with the identical protocol. The digital chain # floors (genie / aged pilot) are drawn as horizontal # references since they involve no training. # ============================================================ import argparse import csv import os import numpy as np import torch import c11_doppler_csi as base import c13_mnist as m13 MODELS = { "semantic": (m13.MnistSemanticMA, "Proposed signed joint"), "semantic_tf": (m13.MnistTransformerMA, "Transformer SE separation"), "semantic_ae": (m13.MnistPerUserAE, "Per-user AE"), } @torch.no_grad() def eval_ser(model, x, y, chan, X_pilot, args, device, rng): model.eval() noise_var = 10 ** (-args.eval_snr / 10.0) errs, total = 0, 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, args.eval_fd, args.delta) logits = m13.forward_frames(model, imgs, chan, g_p, g_d, X_pilot, noise_var, args.eval_fd, args.delta) errs += (logits.argmax(-1) != labels).sum().item() total += labels.numel() model.train() return errs / total @torch.no_grad() def eval_ser_adapted(model, fast, x, y, chan, X_pilot, args, device): noise_var = 10 ** (-args.eval_snr / 10.0) rng = np.random.default_rng(args.seed + 3) errs, total = 0, 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, args.eval_fd, args.delta) logits = m13.forward_frames_fast(model, fast, imgs, chan, g_p, g_d, X_pilot, noise_var, args.eval_fd, args.delta) errs += (logits.argmax(-1) != labels).sum().item() total += labels.numel() return errs / total def run_maml_track(args, device): """Retrace the signed joint training and record the task-adapted SER at every checkpoint, appending method 'semantic_maml_pre'.""" tr, te = m13.get_datasets(args.data_root) x_tr, y_tr = m13.tensorize(tr) x_te, y_te = m13.tensorize(te) chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = base.make_pilot(args.nfft, device) base.set_seed(args.seed) model = m13.MnistSemanticMA(args.users, args.dim, args.hidden).to(device) opt = torch.optim.Adam(model.parameters(), lr=args.lr) rng = np.random.default_rng(args.seed) csv_path = os.path.join(args.save_dir, "mnist_epoch.csv") with open(csv_path, "a", newline="") as f: w = csv.writer(f) fast = m13.adapt_mnist(model, x_tr, y_tr, chan, X_pilot, args, args.eval_snr, device) ser0 = eval_ser_adapted(model, fast, x_te, y_te, chan, X_pilot, args, device) w.writerow(["semantic_maml_pre", 0, ser0]) for step in range(1, args.steps + 1): snr = float(rng.choice(args.train_snrs)) fd = float(rng.choice(args.train_fds)) imgs, labels = m13.sample_frames(x_tr, y_tr, args.batch, args.users, device, rng) g_p, g_d = chan.sample(args.batch, fd, args.delta) noise_var = 10 ** (-snr / 10.0) logits = m13.forward_frames(model, imgs, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta) loss = torch.nn.functional.cross_entropy( logits.reshape(-1, 10), labels.reshape(-1)) opt.zero_grad(set_to_none=True) loss.backward() opt.step() if step % args.ckpt_every == 0: fast = m13.adapt_mnist(model, x_tr, y_tr, chan, X_pilot, args, args.eval_snr, device) ser = eval_ser_adapted(model, fast, x_te, y_te, chan, X_pilot, args, device) w.writerow(["semantic_maml_pre", step, ser]) f.flush() print(f"[maml-pre {step}/{args.steps}] ser={ser:.4e}", flush=True) print("appended semantic_maml_pre to", csv_path) def main(): p = argparse.ArgumentParser() p.add_argument("--mode", choices=["run", "run-maml", "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("--train-snrs", type=float, nargs="+", default=[0, 5, 10, 15, 20, 25]) p.add_argument("--train-fds", type=float, nargs="+", default=[0.002, 0.005, 0.01, 0.02, 0.05, 0.1]) p.add_argument("--steps", type=int, default=10000) p.add_argument("--ckpt-every", type=int, default=500) p.add_argument("--batch", type=int, default=64) p.add_argument("--lr", type=float, default=1e-3) p.add_argument("--eval-snr", type=float, default=15.0) p.add_argument("--eval-fd", type=float, default=0.05) 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=30) 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() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print("device:", device) csv_path = os.path.join(args.save_dir, "mnist_epoch.csv") if args.mode == "run-maml": run_maml_track(args, device) return if args.mode == "run": tr, te = m13.get_datasets(args.data_root) x_tr, y_tr = m13.tensorize(tr) x_te, y_te = m13.tensorize(te) chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = base.make_pilot(args.nfft, device) with open(csv_path, "w", newline="") as f: w = csv.writer(f) w.writerow(["method", "step", "ser"]) for key, (cls, _) in MODELS.items(): base.set_seed(args.seed) model = cls(args.users, args.dim, args.hidden).to(device) opt = torch.optim.Adam(model.parameters(), lr=args.lr) rng = np.random.default_rng(args.seed) erng = np.random.default_rng(args.seed + 3) ser0 = eval_ser(model, x_te, y_te, chan, X_pilot, args, device, np.random.default_rng(args.seed + 3)) w.writerow([key, 0, ser0]) for step in range(1, args.steps + 1): snr = float(rng.choice(args.train_snrs)) fd = float(rng.choice(args.train_fds)) imgs, labels = m13.sample_frames(x_tr, y_tr, args.batch, args.users, device, rng) g_p, g_d = chan.sample(args.batch, fd, args.delta) noise_var = 10 ** (-snr / 10.0) logits = m13.forward_frames(model, imgs, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta) loss = torch.nn.functional.cross_entropy( logits.reshape(-1, 10), labels.reshape(-1)) opt.zero_grad(set_to_none=True) loss.backward() opt.step() if step % args.ckpt_every == 0: ser = eval_ser(model, x_te, y_te, chan, X_pilot, args, device, np.random.default_rng(args.seed + 3)) w.writerow([key, step, ser]) f.flush() print(f"[{key} {step}/{args.steps}] ser={ser:.4e}", flush=True) print("saved", csv_path) return # ---- fig ---- import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt rows = list(csv.DictReader(open(csv_path))) dig = list(csv.DictReader(open(os.path.join(args.save_dir, "mnist_results.csv")))) floors = {r["method"]: float(r["ser"]) for r in dig if float(r["snr_db"]) == args.eval_snr and r["method"] in ("digital_genie", "digital_pilot")} STY = {"semantic": dict(color="tab:red", marker="o", ls="-"), "semantic_tf": dict(color="tab:purple", marker="P", ls="-"), "semantic_ae": dict(color="tab:brown", marker="X", 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]) xgrid = sorted({int(r["step"]) for r in rows if r["method"] == "semantic"}) if "digital_genie" in floors: ax.semilogy(xgrid, [floors["digital_genie"]] * len(xgrid), label="Digital chain genie CSI", ms=5, lw=1.7, color="gray", marker="^", ls="--") if "digital_pilot" in floors: ax.semilogy(xgrid, [floors["digital_pilot"]] * len(xgrid), label="Digital chain pilot CSI", ms=5, lw=1.7, color="k", marker="v", ls="-") for key in ["semantic_tf", "semantic_ae", "semantic"]: label = MODELS[key][1] pts = sorted([(int(r["step"]), float(r["ser"])) for r in rows if r["method"] == key]) xs, ys = zip(*pts) ax.semilogy(xs, ys, label=label, ms=5, lw=1.8, **STY[key]) # signed MAML: the deployed receiver applies the five-step # task-conditional adaptation at every checkpoint. During the # warm-start phase the adapted SER of the evolving joint model is # shown, and right of the dotted line decoder-side meta-training # continues within the same total step budget. maml_path = os.path.join(args.save_dir, "mnist_maml_epoch.csv") if os.path.exists(maml_path): mrows = list(csv.DictReader(open(maml_path))) meta_pts = sorted([(int(r["step"]), float(r["ser"])) for r in mrows]) warm_end = meta_pts[0][0] pre_pts = sorted([(int(r["step"]), float(r["ser"])) for r in rows if r["method"] == "semantic_maml_pre" and int(r["step"]) < warm_end]) if not pre_pts: pre_pts = sorted([(int(r["step"]), float(r["ser"])) for r in rows if r["method"] == "semantic" and int(r["step"]) < warm_end]) # Left of the dotted line no meta-training has happened yet, so # the curve is the joint model with the same five-step # adaptation applied. It is drawn dashed with open markers and # carries its own legend entry, since it is the control # condition rather than the proposed MAML receiver. pre_seg = pre_pts + meta_pts[:1] xs, ys = zip(*pre_seg) ax.semilogy(xs, ys, label="Adaptation from joint model", ms=5, lw=1.8, color="tab:green", marker="D", ls="--", markerfacecolor="none") xs, ys = zip(*meta_pts) ax.semilogy(xs, ys, label="Proposed signed MAML", ms=5, lw=1.8, color="tab:green", marker="D", ls="-") ax.axvline(warm_end, color="gray", ls=":", lw=1.4) ax.set_xlabel("Training step") ax.set_ylabel("SER") if "digital_genie" in floors: # extra headroom below the genie floor so that the seven-entry # legend sits in free space instead of over the curves ax.set_ylim(bottom=floors["digital_genie"] * 0.5) ax.grid(True, which="both", alpha=0.35) ax.legend(fontsize=9, loc="center left", bbox_to_anchor=(0.02, 0.34), framealpha=1.0, labelspacing=0.3, handlelength=1.8) out = os.path.join(args.fig_dir, "mnist_epoch_convergence.pdf") fig.savefig(out) print("saved", out) if __name__ == "__main__": main()