Files
JSAC_AIRAN/c19_mnist_epoch.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

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