805 lines
34 KiB
Python
Executable File
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()
|