Enlarge the in-canvas fonts and line weights of the result figures so that they stay legible at the printed column width, split the Fig. 5 convergence curve into a pre-meta adaptation entry and the proposed MAML entry, and rename the autoencoder legend to match the table row. Add the analytic complexity replot behind Fig. 4, which was missing from the repository, and correct the table numbering in the README.
189 lines
8.6 KiB
Python
Executable File
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.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()
|