260 lines
12 KiB
Python
Executable File
260 lines
12 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 multiple access"),
|
|
}
|
|
|
|
|
|
@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="-")}
|
|
fig = plt.figure(figsize=(5.2, 3.9))
|
|
ax = fig.add_axes([0.14, 0.125, 0.835, 0.845])
|
|
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=3.5, lw=1.2,
|
|
color="gray", marker="^", ls="--")
|
|
if "digital_pilot" in floors:
|
|
ax.semilogy(xgrid, [floors["digital_pilot"]] * len(xgrid),
|
|
label="Digital chain pilot CSI", ms=3.5, lw=1.2,
|
|
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=3.5, lw=1.3, **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])
|
|
pts = pre_pts + meta_pts
|
|
xs, ys = zip(*pts)
|
|
ax.semilogy(xs, ys, label="Proposed signed MAML", ms=3.5, lw=1.3,
|
|
color="tab:green", marker="D", ls="--")
|
|
ax.axvline(warm_end, color="gray", ls=":", lw=1.0)
|
|
ax.set_xlabel("Training step")
|
|
ax.set_ylabel("SER")
|
|
if "digital_genie" in floors:
|
|
ax.set_ylim(bottom=floors["digital_genie"] * 0.5)
|
|
ax.grid(True, which="both", alpha=0.35)
|
|
ax.legend(fontsize=7.5, loc="center left", bbox_to_anchor=(0.02, 0.32))
|
|
out = os.path.join(args.fig_dir, "mnist_epoch_convergence.pdf")
|
|
fig.savefig(out)
|
|
print("saved", out)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|