#!/usr/bin/env python3 # ============================================================ # c11_doppler_csi.py # # TWC/TCOM revision experiments: # Graceful degradation & CSI-aging robustness under a # time-varying frequency-selective (TDL-like) OFDM channel. # # Frame structure (per frame): # [ pilot OFDM symbol ] ... gap (DELTA-1 symbols) ... [ data OFDM symbol ] # - Channel taps evolve continuously (Jakes sum-of-sinusoids) # across the frame -> pilot CSI is OUTDATED at the data symbol. # - Time-varying convolution within each OFDM symbol -> genuine ICI. # # Methods: # qpsk_genie : OFDMA-QPSK, comb allocation (N/U subcarriers/user, # repetition + MRC), perfect CSI at data time (bound) # qpsk_pilot : same, but LS pilot CSI (aged) -> cliff effect # joint : proposed user-wise attention semantic demux, # joint training over (SNR, fD) grid, zero-shot eval # joint_ft : joint + test-time fine-tuning (same budget as MAML) # maml : SNR/Doppler-aware MAML (decoder-side inner loop) # + test-time task-conditional adaptation # # The proposed receiver uses the SAME pilot (MMSE equalization with # the aged LS estimate), so pilot overhead is identical to baseline. # Difference is isolated to the demapping: fixed coherent QPSK vs # learned semantic demultiplexing. # # Tasks tau = (SNR, fD_norm), fD_norm = f_D * T_sym (N samples). # # Usage: # python3 c11_doppler_csi.py --mode train-joint # python3 c11_doppler_csi.py --mode train-maml # python3 c11_doppler_csi.py --mode eval # python3 c11_doppler_csi.py --mode fig # ============================================================ import argparse import math import os import csv import numpy as np import torch import torch.nn as nn import torch.nn.functional as F try: from torch.func import functional_call except Exception: from torch.nn.utils.stateless import functional_call try: from scipy.special import j0 as bessel_j0 except Exception: def bessel_j0(x): """Abramowitz & Stegun 9.4.1 / 9.4.3 polynomial approximation.""" x = abs(float(x)) if x < 3.0: t = (x / 3.0) ** 2 return (1.0 - 2.2499997 * t + 1.2656208 * t ** 2 - 0.3163866 * t ** 3 + 0.0444479 * t ** 4 - 0.0039444 * t ** 5 + 0.0002100 * t ** 6) t = 3.0 / x f0 = (0.79788456 - 0.00000077 * t - 0.00552740 * t ** 2 - 0.00009512 * t ** 3 + 0.00137237 * t ** 4 - 0.00072805 * t ** 5 + 0.00014476 * t ** 6) th = (x - 0.78539816 - 0.04166397 * t - 0.00003954 * t ** 2 + 0.00262573 * t ** 3 - 0.00054125 * t ** 4 - 0.00029333 * t ** 5 + 0.00013558 * t ** 6) return f0 * math.cos(th) / math.sqrt(x) # ------------------------------------------------------------ # Reproducibility # ------------------------------------------------------------ def set_seed(seed: int = 0): torch.manual_seed(seed) np.random.seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) # ------------------------------------------------------------ # Time-varying TDL channel (Jakes sum-of-sinusoids) # ------------------------------------------------------------ class TDLChannel: """ Exponential PDP, L taps, per-tap Jakes Doppler via sum-of-sinusoids. fd_norm is the Doppler normalized to the useful OFDM symbol duration (N samples): fd_norm = f_D * N * T_s. """ def __init__(self, N=128, cp=16, L=8, pdp_decay=3.0, n_sin=16, device="cpu"): self.N = N self.cp = cp self.L = L self.S = N + cp self.n_sin = n_sin self.device = device p = torch.exp(-torch.arange(L, dtype=torch.float32) / pdp_decay) self.pdp = (p / p.sum()).to(device) # (L,) def sample(self, B, fd_norm, delta): """ Generate tap gains at pilot-symbol samples and data-symbol samples. Pilot occupies samples [0, S); data occupies [delta*S, delta*S + S). Returns g_p, g_d: (B, L, S) complex tap trajectories. """ dev = self.device S, L, Ns = self.S, self.L, self.n_sin fd_samp = fd_norm / self.N # per-sample normalized Doppler n_p = torch.arange(S, device=dev, dtype=torch.float32) n_d = n_p + delta * S n_all = torch.cat([n_p, n_d]) # (2S,) theta = 2 * math.pi * torch.rand(B, L, 1, Ns, device=dev) phi = 2 * math.pi * torch.rand(B, L, 1, Ns, device=dev) omega = 2 * math.pi * fd_samp * torch.cos(theta) # (B,L,1,Ns) ph = omega * n_all.view(1, 1, -1, 1) + phi # (B,L,2S,Ns) g = torch.exp(1j * ph).sum(dim=-1) / math.sqrt(Ns) # (B,L,2S) g = g * torch.sqrt(self.pdp).view(1, L, 1) return g[:, :, :S], g[:, :, S:] def transmit(self, X_freq, g, noise_var): """ One OFDM symbol through the time-varying channel. X_freq: (B, N) complex subcarrier vector (unitary convention) g: (B, L, S) tap gains over the symbol (incl. CP samples) Returns Y: (B, N) complex received subcarrier vector. """ B = X_freq.shape[0] N, cp, S, L = self.N, self.cp, self.S, self.L x = torch.fft.ifft(X_freq, dim=1) * math.sqrt(N) # (B,N) x_cp = torch.cat([x[:, -cp:], x], dim=1) # (B,S) y = torch.zeros(B, S, dtype=torch.complex64, device=x.device) for l in range(L): if l == 0: y = y + g[:, 0, :] * x_cp else: y[:, l:] = y[:, l:] + g[:, l, l:] * x_cp[:, :S - l] n = (torch.randn(B, S, device=x.device) + 1j * torch.randn(B, S, device=x.device)) * math.sqrt(noise_var / 2.0) y = y + n y_data = y[:, cp:] # discard CP Y = torch.fft.fft(y_data, dim=1) / math.sqrt(N) # (B,N) return Y def genie_H(self, g): """Effective per-subcarrier channel: DFT of time-averaged taps.""" g_bar = g[:, :, self.cp:].mean(dim=2) # (B,L) h = torch.zeros(g.shape[0], self.N, dtype=torch.complex64, device=g.device) h[:, :self.L] = g_bar return torch.fft.fft(h, dim=1) # (B,N) def make_pilot(N, device, seed=1234): """Fixed pseudo-random QPSK pilot, |X_p[k]| = 1.""" gen = torch.Generator(device="cpu").manual_seed(seed) idx = torch.randint(0, 4, (N,), generator=gen) ang = math.pi / 4 + idx.to(torch.float32) * math.pi / 2 return torch.exp(1j * ang).to(device) # ------------------------------------------------------------ # OFDMA-QPSK baseline (comb allocation + repetition + MRC) # ------------------------------------------------------------ QPSK = None # filled at runtime def qpsk_constellation(device): ang = math.pi / 4 + torch.arange(4, device=device, dtype=torch.float32) * math.pi / 2 return torch.exp(1j * ang) # (4,) def qpsk_tx(labels, U, N, device): """labels: (B,U) in {0..3} -> X: (B,N), comb allocation.""" B = labels.shape[0] const = qpsk_constellation(device) sym = const[labels] # (B,U) X = torch.zeros(B, N, dtype=torch.complex64, device=device) for u in range(U): ks = torch.arange(u, N, U, device=device) X[:, ks] = sym[:, u].unsqueeze(1) return X def qpsk_detect(Y, H_hat, U, N): """MRC over each user's comb subcarriers with (possibly aged) CSI.""" B = Y.shape[0] const = qpsk_constellation(Y.device) # (4,) preds = torch.zeros(B, U, dtype=torch.long, device=Y.device) for u in range(U): ks = torch.arange(u, N, U, device=Y.device) Z = (torch.conj(H_hat[:, ks]) * Y[:, ks]).sum(dim=1) # (B,) metric = torch.real(torch.conj(const).view(1, 4) * Z.unsqueeze(1)) preds[:, u] = metric.argmax(dim=1) return preds # ------------------------------------------------------------ # Proposed model: masked superposition + user-wise attention # operating on the MMSE-equalized received subcarrier vector # (real/imag stacked features, masks applied per subcarrier) # ------------------------------------------------------------ class UserWiseAttentionSem(nn.Module): def __init__(self, U, V, d, hidden=256): super().__init__() self.U, self.V, self.d = U, V, d self.codebook = nn.Parameter(torch.randn(V, d)) self.masks = nn.Parameter(torch.randn(U, d)) self.query = nn.Parameter(torch.randn(U, 2 * d)) self.key = nn.Linear(2 * d, 2 * d, bias=False) self.val = nn.Linear(2 * d, 2 * d, bias=False) self.cls = nn.Sequential( nn.Linear(2 * d, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, V), ) # ---- transmitter ---- def tx(self, labels, params=None): cb = self.codebook if params is None else params["codebook"] mk = self.masks if params is None else params["masks"] e = F.embedding(labels, cb) # (B,U,d) m = F.normalize(mk, dim=1) # (U,d) y = (e * m.unsqueeze(0)).sum(dim=1) # (B,d) y = y / torch.sqrt(torch.mean(y ** 2, dim=1, keepdim=True) + 1e-12) return y, m # ---- receiver ---- def rx(self, Yeq, m, params=None): B = Yeq.shape[0] R = Yeq.unsqueeze(1) * m.unsqueeze(0) # (B,U,d) complex phi = torch.cat([R.real, R.imag], dim=-1) # (B,U,2d) if params is None: K = self.key(phi) Vv = self.val(phi) q = self.query else: K = F.linear(phi, params["key.weight"]) Vv = F.linear(phi, params["val.weight"]) q = params["query"] scores = torch.einsum("ud,bid->bui", q, K) / math.sqrt(2 * self.d) attn = F.softmax(scores, dim=-1) # (B,U,U) z = torch.einsum("bui,bid->bud", attn, Vv) + phi # residual (B,U,2d) if params is None: logits = self.cls(z) else: h1 = F.linear(z, params["cls.0.weight"], params["cls.0.bias"]) h1 = F.relu(h1) logits = F.linear(h1, params["cls.2.weight"], params["cls.2.bias"]) return logits # (B,U,V) class TransformerSem(UserWiseAttentionSem): """ SOTA baseline: shared-embedding TX (same masking) + Transformer encoder separation treating users as tokens (JSAC'26 [21]-style). """ def __init__(self, U, V, d, hidden=256, n_heads=4, n_layers=2): super().__init__(U, V, d, hidden) self.inp = nn.Linear(2 * d, hidden) layer = nn.TransformerEncoderLayer(d_model=hidden, nhead=n_heads, dim_feedforward=2 * hidden, batch_first=True) self.encoder = nn.TransformerEncoder(layer, num_layers=n_layers) self.out = nn.Linear(hidden, V) def rx(self, Yeq, m, params=None): R = Yeq.unsqueeze(1) * m.unsqueeze(0) # (B,U,d) phi = torch.cat([R.real, R.imag], dim=-1) # (B,U,2d) z = self.encoder(self.inp(phi)) # (B,U,hidden) return self.out(z) # (B,U,V) class PerUserAESem(nn.Module): """ SOTA baseline: DeepMA-style per-user autoencoder multiple access. Each user has a dedicated codebook (no shared masking) and a dedicated decoder head operating on the equalized observation. """ def __init__(self, U, V, d, hidden=256): super().__init__() self.U, self.V, self.d = U, V, d self.codebooks = nn.Parameter(torch.randn(U, V, d)) self.heads = nn.ModuleList([ nn.Sequential(nn.Linear(2 * d, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, V)) for _ in range(U) ]) def tx(self, labels, params=None): B = labels.shape[0] idx = labels + torch.arange(self.U, device=labels.device).view(1, -1) * self.V e = self.codebooks.view(self.U * self.V, self.d)[idx] # (B,U,d) y = e.sum(dim=1) y = y / torch.sqrt(torch.mean(y ** 2, dim=1, keepdim=True) + 1e-12) return y, None def rx(self, Yeq, m, params=None): phi = torch.cat([Yeq.real, Yeq.imag], dim=-1) # (B,2d) return torch.stack([h(phi) for h in self.heads], dim=1) # (B,U,V) class SignedAttentionSem(UserWiseAttentionSem): """ Proposed signed user-wise attention. The candidate Gram matrix T = Phi Phi^T / (2d) is mapped by a small score network to signed combining weights W = I + f_theta(T), which can subtract correlated interference, unlike convex nonnegative softmax weights. The query/key/value projections of the parent are unused. """ def __init__(self, U, V, d, hidden=256, score_hidden=64): super().__init__(U, V, d, hidden) self.score = nn.Sequential( nn.Linear(U * U, score_hidden), nn.ReLU(inplace=True), nn.Linear(score_hidden, score_hidden), nn.ReLU(inplace=True), nn.Linear(score_hidden, U * U), ) def rx(self, Yeq, m, params=None): B = Yeq.shape[0] R = Yeq.unsqueeze(1) * m.unsqueeze(0) # (B,U,d) phi = torch.cat([R.real, R.imag], dim=-1) # (B,U,2d) T = torch.bmm(phi, phi.transpose(1, 2)) / (2 * self.d) if params is None: w = self.score(T.reshape(B, -1)).view(B, self.U, self.U) W = torch.eye(self.U, device=phi.device).unsqueeze(0) + w z = torch.bmm(W, phi) return self.cls(z) h = T.reshape(B, -1) h = F.relu(F.linear(h, params["score.0.weight"], params["score.0.bias"])) h = F.relu(F.linear(h, params["score.2.weight"], params["score.2.bias"])) w = F.linear(h, params["score.4.weight"], params["score.4.bias"]) W = torch.eye(self.U, device=phi.device).unsqueeze(0) + w.view(B, self.U, self.U) z = torch.bmm(W, phi) h1 = F.relu(F.linear(z, params["cls.0.weight"], params["cls.0.bias"])) return F.linear(h1, params["cls.2.weight"], params["cls.2.bias"]) MODEL_CLASSES = { "attention": UserWiseAttentionSem, "signed": SignedAttentionSem, "transformer": TransformerSem, "peruser": PerUserAESem, } DECODER_KEYS = ["query", "key.weight", "val.weight", "cls.0.weight", "cls.0.bias", "cls.2.weight", "cls.2.bias"] SIGNED_DECODER_KEYS = ["score.0.weight", "score.0.bias", "score.2.weight", "score.2.bias", "score.4.weight", "score.4.bias", "cls.0.weight", "cls.0.bias", "cls.2.weight", "cls.2.bias"] def decoder_keys_for(model): return SIGNED_DECODER_KEYS if isinstance(model, SignedAttentionSem) else DECODER_KEYS def get_params(model, keys=None): d = dict(model.named_parameters()) if keys is None: return d return {k: d[k] for k in keys} # ------------------------------------------------------------ # End-to-end forward through channel (differentiable) # ------------------------------------------------------------ def aging_rho(fd_norm, delta, N, cp): """Pilot-to-data temporal correlation: rho = J0(2*pi*fd_samp*lag).""" fd_samp = fd_norm / N lag = delta * (N + cp) return float(bessel_j0(2.0 * math.pi * fd_samp * lag)) def semantic_forward(model, labels, chan, g_p, g_d, X_pilot, noise_var, fd_norm, delta, params=None): """ Full pipeline: TX -> time-varying channel (pilot + data symbols) -> LS pilot estimate -> aging-aware LMMSE equalization -> user-wise attention. Aging-aware LMMSE: with H_d = rho*H_p + sqrt(1-rho^2)*innovation and LS estimate Hhat = H_p + e (Var e = sigma^2), the LMMSE predictor of the data-time channel is Htil = (rho/(1+sigma^2))*Hhat with residual variance q = 1 - rho^2/(1+sigma^2). The equalizer stays regularized at high SNR because q > 0 whenever rho < 1. """ full = None if params is not None: full = dict(model.named_parameters()) full.update(params) y_emb, m = model.tx(labels, params=full) X_data = y_emb.to(torch.complex64) # real coeffs, imag=0 Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(labels.shape[0], -1), g_p, noise_var) Y_d = chan.transmit(X_data, g_d, noise_var) H_ls = Y_p * torch.conj(X_pilot).unsqueeze(0) # |X_p|=1 rho = aging_rho(fd_norm, delta, chan.N, chan.cp) H_til = (rho / (1.0 + noise_var)) * H_ls q = 1.0 - (rho ** 2) / (1.0 + noise_var) Yeq = torch.conj(H_til) * Y_d / (H_til.abs() ** 2 + q + noise_var) # per-frame RMS normalization: removes the task-dependent scale of the # LMMSE output (~rho/(1+sigma^2)) so that a single receiver operates on # a scale-invariant representation across (SNR, Doppler) tasks rms = torch.sqrt(torch.mean(Yeq.abs() ** 2, dim=1, keepdim=True) + 1e-12) Yeq = Yeq / rms return model.rx(Yeq, m, params=full) # ------------------------------------------------------------ # Task sampling # ------------------------------------------------------------ def sample_task(args, rng): snr = float(rng.choice(args.train_snrs)) fd = float(rng.choice(args.train_fds)) return snr, fd # ------------------------------------------------------------ # Training # ------------------------------------------------------------ def run_batch_loss(model, chan, X_pilot, args, snr, fd, batch, params=None, device="cpu"): labels = torch.randint(0, args.vocab, (batch, args.users), device=device) g_p, g_d = chan.sample(batch, fd, args.delta) noise_var = 10 ** (-snr / 10.0) logits = semantic_forward(model, labels, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta, params=params) loss = F.cross_entropy(logits.reshape(-1, args.vocab), labels.reshape(-1)) return loss, logits, labels def train_joint(args, device, model_key="attention", ckpt="joint.pt"): set_seed(args.seed) chan = TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = make_pilot(args.nfft, device) model = MODEL_CLASSES[model_key](args.users, args.vocab, args.dim, args.hidden).to(device) opt = torch.optim.Adam(model.parameters(), lr=args.lr) rng = np.random.default_rng(args.seed) for step in range(1, args.steps + 1): snr, fd = sample_task(args, rng) loss, _, _ = run_batch_loss(model, chan, X_pilot, args, snr, fd, args.batch, device=device) opt.zero_grad(set_to_none=True) loss.backward() opt.step() if step % 200 == 0: print(f"[{model_key} {step}/{args.steps}] loss={loss.item():.4f} (snr={snr}, fd={fd})", flush=True) os.makedirs(args.save_dir, exist_ok=True) torch.save(model.state_dict(), os.path.join(args.save_dir, ckpt)) print(f"saved {ckpt}") def train_maml(args, device, model_key="attention", ckpt="maml.pt"): set_seed(args.seed) chan = TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = make_pilot(args.nfft, device) model = MODEL_CLASSES[model_key](args.users, args.vocab, args.dim, args.hidden).to(device) opt = torch.optim.Adam(model.parameters(), lr=args.meta_lr) rng = np.random.default_rng(args.seed) dkeys = decoder_keys_for(model) for step in range(1, args.meta_steps + 1): meta_loss = 0.0 for _ in range(args.meta_batch): snr, fd = sample_task(args, rng) fast = get_params(model, dkeys) for _ in range(args.inner_steps): loss_sup, _, _ = run_batch_loss(model, chan, X_pilot, args, snr, fd, args.batch, params=fast, device=device) grads = torch.autograd.grad(loss_sup, list(fast.values()), create_graph=False) fast = {k: p - args.inner_lr * g.detach() for (k, p), g in zip(fast.items(), grads)} loss_q, _, _ = run_batch_loss(model, chan, X_pilot, args, snr, fd, args.batch, params=fast, device=device) meta_loss = meta_loss + loss_q meta_loss = meta_loss / args.meta_batch opt.zero_grad(set_to_none=True) meta_loss.backward() opt.step() if step % 100 == 0: print(f"[maml {step}/{args.meta_steps}] meta-loss={meta_loss.item():.4f}", flush=True) os.makedirs(args.save_dir, exist_ok=True) torch.save(model.state_dict(), os.path.join(args.save_dir, ckpt)) print(f"saved {ckpt}") # ------------------------------------------------------------ # Evaluation # ------------------------------------------------------------ def adapt(model, chan, X_pilot, args, snr, fd, device): """Test-time task-conditional adaptation (decoder-side, FOMAML-style).""" fast = {k: v.detach().clone() for k, v in get_params(model, decoder_keys_for(model)).items()} for k in fast: fast[k].requires_grad_(True) for _ in range(args.eval_inner_steps): loss, _, _ = run_batch_loss(model, chan, X_pilot, args, snr, fd, args.support, params=fast, device=device) grads = torch.autograd.grad(loss, list(fast.values())) fast = {k: (p - args.eval_inner_lr * g).detach().requires_grad_(True) for (k, p), g in zip(fast.items(), grads)} return {k: v.detach() for k, v in fast.items()} @torch.no_grad() def eval_semantic(model, chan, X_pilot, args, snr, fd, n_batches, device, params=None): errs, total = 0, 0 for _ in range(n_batches): labels = torch.randint(0, args.vocab, (args.eval_batch, args.users), device=device) g_p, g_d = chan.sample(args.eval_batch, fd, args.delta) noise_var = 10 ** (-snr / 10.0) logits = semantic_forward(model, labels, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta, params=params) preds = logits.argmax(dim=-1) errs += (preds != labels).sum().item() total += labels.numel() return errs / total @torch.no_grad() def eval_qpsk(chan, X_pilot, args, snr, fd, n_batches, device, genie=False): errs, total = 0, 0 for _ in range(n_batches): labels = torch.randint(0, args.vocab, (args.eval_batch, args.users), device=device) g_p, g_d = chan.sample(args.eval_batch, fd, args.delta) noise_var = 10 ** (-snr / 10.0) X_d = qpsk_tx(labels, args.users, args.nfft, device) Y_d = chan.transmit(X_d, g_d, noise_var) if genie: H_hat = chan.genie_H(g_d) else: Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(labels.shape[0], -1), g_p, noise_var) H_hat = Y_p * torch.conj(X_pilot).unsqueeze(0) preds = qpsk_detect(Y_d, H_hat, args.users, args.nfft) errs += (preds != labels).sum().item() total += labels.numel() return errs / total def eval_all(args, device): set_seed(args.seed + 1) chan = TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = make_pilot(args.nfft, device) joint = UserWiseAttentionSem(args.users, args.vocab, args.dim, args.hidden).to(device) joint.load_state_dict(torch.load(os.path.join(args.save_dir, "joint.pt"), map_location=device)) joint.eval() maml = UserWiseAttentionSem(args.users, args.vocab, args.dim, args.hidden).to(device) maml.load_state_dict(torch.load(os.path.join(args.save_dir, "maml.pt"), map_location=device)) maml.eval() signed_models = {} for key, ckpt in [("sjoint", "signed_joint.pt"), ("smaml", "signed_maml.pt")]: path = os.path.join(args.save_dir, ckpt) if os.path.exists(path): mdl = SignedAttentionSem(args.users, args.vocab, args.dim, args.hidden).to(device) mdl.load_state_dict(torch.load(path, map_location=device)) mdl.eval() signed_models[key] = mdl # SOTA comparison models (zero-shot joint training), if available sota = {} for key, ckpt in [("transformer", "transformer.pt"), ("peruser", "peruser.pt")]: path = os.path.join(args.save_dir, ckpt) if os.path.exists(path): mdl = MODEL_CLASSES[key](args.users, args.vocab, args.dim, args.hidden).to(device) mdl.load_state_dict(torch.load(path, map_location=device)) mdl.eval() sota[key] = mdl sweeps = [] # Sweep A: SER vs Doppler at fixed SNR for fd in args.eval_fds: sweeps.append(("doppler", args.eval_snr_fixed, fd)) # Sweep B: SER vs SNR at fixed (high) Doppler for snr in args.eval_snrs: sweeps.append(("snr", snr, args.eval_fd_fixed)) os.makedirs(args.save_dir, exist_ok=True) csv_path = os.path.join(args.save_dir, "doppler_csi_results.csv") with open(csv_path, "w", newline="") as f: w = csv.writer(f) w.writerow(["sweep", "snr_db", "fd_norm", "method", "ser"]) for sweep, snr, fd in sweeps: row = {} row["qpsk_genie"] = eval_qpsk(chan, X_pilot, args, snr, fd, args.eval_nb, device, genie=True) row["qpsk_pilot"] = eval_qpsk(chan, X_pilot, args, snr, fd, args.eval_nb, device, genie=False) for key, mdl in sota.items(): row[key] = eval_semantic(mdl, chan, X_pilot, args, snr, fd, args.eval_nb, device) row["joint"] = eval_semantic(joint, chan, X_pilot, args, snr, fd, args.eval_nb, device) if "sjoint" in signed_models: row["sjoint"] = eval_semantic(signed_models["sjoint"], chan, X_pilot, args, snr, fd, args.eval_nb, device) # adaptation-based methods: average over independent trials adapt_list = [("joint_ft", joint), ("maml", maml)] if "smaml" in signed_models: adapt_list.append(("smaml", signed_models["smaml"])) for name, mdl in adapt_list: sers = [] for _ in range(args.adapt_trials): with torch.enable_grad(): fast = adapt(mdl, chan, X_pilot, args, snr, fd, device) sers.append(eval_semantic(mdl, chan, X_pilot, args, snr, fd, args.eval_nb_adapt, device, params=fast)) row[name] = float(np.mean(sers)) for method, ser in row.items(): w.writerow([sweep, snr, fd, method, ser]) f.flush() print(f"[{sweep}] snr={snr:5.1f} fd={fd:6.3f} | " + " ".join(f"{k}={v:.4e}" for k, v in row.items()), flush=True) print(f"saved {csv_path}") # ------------------------------------------------------------ # Figures # ------------------------------------------------------------ LABELS = { "qpsk_genie": "OFDMA-QPSK genie CSI", "qpsk_pilot": "OFDMA-QPSK pilot CSI", "transformer": "Transformer SE separation", "peruser": "Per-user AE multiple access", "joint": "Proposed softmax joint", "joint_ft": "Proposed softmax fine-tuned", "maml": "Proposed softmax MAML", "sjoint": "Proposed signed joint", "smaml": "Proposed signed MAML", } STYLES = { "qpsk_genie": dict(color="gray", marker="^", ls="--"), "qpsk_pilot": dict(color="k", marker="v", ls="-"), "transformer": dict(color="tab:purple", marker="P", ls="-"), "peruser": dict(color="tab:brown", marker="X", ls="-"), "joint": dict(color="tab:blue", marker="s", ls="-"), "joint_ft": dict(color="tab:green", marker="D", ls="-"), "maml": dict(color="tab:cyan", marker="d", ls="-"), "sjoint": dict(color="tab:red", marker="o", ls="-"), "smaml": dict(color="tab:green", marker="D", ls="-"), } PLOT_KEYS = ["qpsk_genie", "qpsk_pilot", "transformer", "peruser", "joint", "sjoint", "smaml"] def make_figs(args): import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt csv_path = os.path.join(args.save_dir, "doppler_csi_results.csv") rows = [] with open(csv_path) as f: for r in csv.DictReader(f): rows.append(r) os.makedirs(args.fig_dir, exist_ok=True) # Fig A: SER vs normalized Doppler fig = plt.figure(figsize=(5.2, 3.9)) ax = fig.add_axes([0.14, 0.125, 0.835, 0.845]) for m in PLOT_KEYS: latest = {} for r in rows: if r["sweep"] == "doppler" and r["method"] == m: latest[float(r["fd_norm"])] = float(r["ser"]) pts = sorted((k, v) for k, v in latest.items() if v > 0) if pts: x, y = zip(*pts) ax.semilogy(x, y, label=LABELS[m], ms=4, lw=1.2, **STYLES[m]) ax.set_xscale("log") 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, loc="upper left") fig.savefig(os.path.join(args.fig_dir, f"ser_vs_doppler_snr{int(args.eval_snr_fixed)}.pdf")) print("saved fig A") # Fig B: SER vs SNR at fixed high Doppler fig = plt.figure(figsize=(5.2, 3.9)) ax = fig.add_axes([0.14, 0.125, 0.835, 0.845]) for m in PLOT_KEYS: latest = {} for r in rows: if r["sweep"] == "snr" and r["method"] == m: latest[float(r["snr_db"])] = float(r["ser"]) pts = sorted((k, v) for k, v in latest.items() if v > 0) if pts: x, y = zip(*pts) ax.semilogy(x, y, label=LABELS[m], ms=4, lw=1.2, **STYLES[m]) ax.set_xlabel("SNR (dB)") ax.set_ylabel("SER") ax.grid(True, which="both", alpha=0.35) ax.legend(fontsize=7, loc="lower center") fig.savefig(os.path.join(args.fig_dir, f"ser_vs_snr_fd{args.eval_fd_fixed}.pdf")) print("saved fig B") # ------------------------------------------------------------ # Main # ------------------------------------------------------------ def main(): p = argparse.ArgumentParser() p.add_argument("--mode", choices=["train-joint", "train-maml", "train-transformer", "train-peruser", "train-signed-joint", "train-signed-maml", "eval", "fig"], required=True) # system p.add_argument("--users", type=int, default=8) p.add_argument("--vocab", type=int, default=4) 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, help="pilot-to-data gap (OFDM symbols)") # task grids 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]) # training p.add_argument("--steps", type=int, default=6000) p.add_argument("--batch", type=int, default=256) p.add_argument("--lr", type=float, default=1e-3) # maml p.add_argument("--meta-steps", type=int, default=3000) p.add_argument("--meta-batch", type=int, default=4) p.add_argument("--inner-steps", type=int, default=1) p.add_argument("--inner-lr", type=float, default=0.05) p.add_argument("--meta-lr", type=float, default=1e-3) # evaluation p.add_argument("--eval-batch", type=int, default=256) p.add_argument("--eval-nb", type=int, default=200, help="eval batches (zero-shot)") p.add_argument("--eval-nb-adapt", type=int, default=40, help="eval batches per adapt trial") p.add_argument("--adapt-trials", type=int, default=8) p.add_argument("--support", type=int, default=64, help="support frames per adapt step") p.add_argument("--eval-inner-steps", type=int, default=5) p.add_argument("--eval-inner-lr", type=float, default=0.02) p.add_argument("--eval-fds", type=float, nargs="+", default=[0.001, 0.002, 0.005, 0.01, 0.02, 0.05, 0.1]) p.add_argument("--eval-snr-fixed", type=float, default=15.0) p.add_argument("--eval-snrs", type=float, nargs="+", default=[0, 5, 10, 15, 20, 25, 30]) p.add_argument("--eval-fd-fixed", type=float, default=0.05) # misc p.add_argument("--save-dir", type=str, default="results_doppler") p.add_argument("--fig-dir", type=str, default="fig") p.add_argument("--seed", type=int, default=0) args = p.parse_args() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"device: {device}") if args.mode == "train-joint": train_joint(args, device) elif args.mode == "train-transformer": train_joint(args, device, model_key="transformer", ckpt="transformer.pt") elif args.mode == "train-peruser": train_joint(args, device, model_key="peruser", ckpt="peruser.pt") elif args.mode == "train-signed-joint": train_joint(args, device, model_key="signed", ckpt="signed_joint.pt") elif args.mode == "train-signed-maml": train_maml(args, device, model_key="signed", ckpt="signed_maml.pt") elif args.mode == "train-maml": train_maml(args, device) elif args.mode == "eval": eval_all(args, device) elif args.mode == "fig": make_figs(args) if __name__ == "__main__": main()