Files
JSAC_AIRAN/c11_doppler_csi.py
T

805 lines
34 KiB
Python
Executable File

#!/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()