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.
278 lines
13 KiB
Python
Executable File
278 lines
13 KiB
Python
Executable File
#!/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()
|