Files
JSAC_AIRAN/c18_mnist_doppler.py
T
KiHoLee 6c8471ece0 Match figure styling and labels to the submitted manuscript
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.
2026-08-03 12:49:53 +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.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()