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.
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.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()
|