Code and stored results for AI-native multi-user semantic communications (JSAC submission)
This commit is contained in:
Executable
+181
@@ -0,0 +1,181 @@
|
||||
#!/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 multiple access",
|
||||
"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="--")}
|
||||
fig = plt.figure(figsize=(5.2, 3.9))
|
||||
ax = fig.add_axes([0.14, 0.125, 0.835, 0.845])
|
||||
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=4, lw=1.3, **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=7.5, 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()
|
||||
Reference in New Issue
Block a user