Code and stored results for AI-native multi-user semantic communications (JSAC submission)
This commit is contained in:
Executable
+155
@@ -0,0 +1,155 @@
|
||||
#!/usr/bin/env python3
|
||||
# ============================================================
|
||||
# c22_maml_epoch.py
|
||||
#
|
||||
# Meta-training trajectory for the signed MAML receiver on the
|
||||
# MNIST TDL chain: warm-started from the converged joint model,
|
||||
# the adapted SER at (15 dB, nu = 0.05) is recorded every
|
||||
# ckpt-every meta-steps, extending the convergence figure of
|
||||
# c19 beyond the joint-training phase.
|
||||
# ============================================================
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import os
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import c11_doppler_csi as base
|
||||
import c13_mnist as m13
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def eval_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 main():
|
||||
p = argparse.ArgumentParser()
|
||||
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("--meta-steps", type=int, default=2500)
|
||||
p.add_argument("--pretrain-steps", type=int, default=0)
|
||||
p.add_argument("--ckpt-every", type=int, default=250)
|
||||
p.add_argument("--meta-batch", type=int, default=4)
|
||||
p.add_argument("--inner-lr", type=float, default=0.02)
|
||||
p.add_argument("--batch", type=int, default=64)
|
||||
p.add_argument("--lr", type=float, default=1e-3)
|
||||
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-snr", type=float, default=15.0)
|
||||
p.add_argument("--eval-fd", type=float, default=0.05)
|
||||
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("--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)
|
||||
|
||||
base.set_seed(args.seed)
|
||||
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)
|
||||
model = m13.MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
|
||||
rng = np.random.default_rng(args.seed)
|
||||
if args.pretrain_steps > 0:
|
||||
# retrace the joint training for the warm-start phase so that
|
||||
# the meta phase continues the same trajectory as the joint curve
|
||||
popt = torch.optim.Adam(model.parameters(), lr=1e-3)
|
||||
for step in range(1, args.pretrain_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 = F.cross_entropy(logits.reshape(-1, 10),
|
||||
labels.reshape(-1))
|
||||
popt.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
popt.step()
|
||||
if step % 2000 == 0:
|
||||
print(f"[pretrain {step}/{args.pretrain_steps}]", flush=True)
|
||||
else:
|
||||
model.load_state_dict(torch.load(
|
||||
os.path.join(args.save_dir, "mnist_semantic.pt"),
|
||||
map_location=device))
|
||||
opt = torch.optim.Adam(model.parameters(), lr=args.lr)
|
||||
|
||||
def task_loss(fast, snr, fd):
|
||||
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_fast(model, fast, imgs, chan, g_p, g_d,
|
||||
X_pilot, noise_var, fd, args.delta)
|
||||
return F.cross_entropy(logits.reshape(-1, 10), labels.reshape(-1))
|
||||
|
||||
csv_path = os.path.join(args.save_dir, "mnist_maml_epoch.csv")
|
||||
with open(csv_path, "w", newline="") as f:
|
||||
w = csv.writer(f)
|
||||
w.writerow(["step", "ser"])
|
||||
fast0 = m13.adapt_mnist(model, x_tr, y_tr, chan, X_pilot, args,
|
||||
args.eval_snr, device)
|
||||
ser0 = eval_adapted(model, fast0, x_te, y_te, chan, X_pilot, args,
|
||||
device)
|
||||
w.writerow([args.pretrain_steps, ser0])
|
||||
print(f"[meta 0] ser={ser0:.4e}", flush=True)
|
||||
for step in range(1, args.meta_steps + 1):
|
||||
meta_loss = 0.0
|
||||
for _ in range(args.meta_batch):
|
||||
snr, fd = base.sample_task(args, rng)
|
||||
fast = {k: v for k, v in model.named_parameters()
|
||||
if k.startswith(m13.MNIST_DECODER_KEYS_PREFIX)}
|
||||
loss_sup = task_loss(fast, snr, fd)
|
||||
grads = torch.autograd.grad(loss_sup, list(fast.values()))
|
||||
fast = {k: p - args.inner_lr * g.detach()
|
||||
for (k, p), g in zip(fast.items(), grads)}
|
||||
meta_loss = meta_loss + task_loss(fast, snr, fd)
|
||||
meta_loss = meta_loss / args.meta_batch
|
||||
opt.zero_grad(set_to_none=True)
|
||||
meta_loss.backward()
|
||||
opt.step()
|
||||
if step % args.ckpt_every == 0:
|
||||
fastc = m13.adapt_mnist(model, x_tr, y_tr, chan, X_pilot,
|
||||
args, args.eval_snr, device)
|
||||
ser = eval_adapted(model, fastc, x_te, y_te, chan, X_pilot,
|
||||
args, device)
|
||||
w.writerow([args.pretrain_steps + step, ser])
|
||||
f.flush()
|
||||
print(f"[meta {step}/{args.meta_steps}] ser={ser:.4e}",
|
||||
flush=True)
|
||||
print("saved", csv_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user