Reproducibility package: UWCA semantic multiple access (TWC submission)

This commit is contained in:
Ki-Ho Lee
2026-08-25 17:55:00 +09:00
commit 6a65d931f7
47 changed files with 10004 additions and 0 deletions
Executable
+490
View File
@@ -0,0 +1,490 @@
"""
Shared library for TWC revision-2 experiments (new submission).
Single-signal uplink model matching the manuscript:
y = sum_v g_v (e_v ⊙ m_v) + n, g_v = |h_v| e^{jΔφ_v}
All experiments import from here. Seed fixed = 42.
"""
import copy
import json
import math
import time
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
DATA = Path(__file__).resolve().parent / "data"
DATA.mkdir(exist_ok=True)
SCENARIOS = {
"HIGH": {"beta_u": [0.65, 0.65, 0.60, 0.60], "scenes": [0, 0, 0, 0]},
"LOW": {"beta_u": [0.65, 0.05, 0.05, 0.05], "scenes": [0, 1, 2, 3]},
"MIX": {"beta_u": [0.65, 0.65, 0.05, 0.05], "scenes": [0, 0, 1, 2]},
# 8-slot high-correlation scenario for dynamic-user experiment
"HIGH8": {"beta_u": [0.60] * 8, "scenes": [0] * 8},
}
SNR_GRID = np.arange(0.0, 20.0 + 1e-6, 2.0)
def beta_matrix(scen):
"""True relevance matrix beta_{u,v} = beta_u beta_v [same scene], else 0."""
b = np.asarray(scen["beta_u"], dtype=np.float64)
sc = np.asarray(scen["scenes"])
U = len(b)
B = np.zeros((U, U))
for u in range(U):
for v in range(U):
B[u, v] = 1.0 if u == v else (b[u] * b[v] if sc[u] == sc[v] else 0.0)
return B
def gen_embeddings(n, d, U, rng, scen):
"""Unit-norm ground-truth embeddings (n, U, d): e_u = sqrt(1-b^2) p + b s."""
blend, scenes = scen["beta_u"], scen["scenes"]
svecs = {}
for sc in sorted(set(scenes)):
v = rng.standard_normal(d)
svecs[sc] = v / (np.linalg.norm(v) + 1e-8)
embs = []
for u in range(U):
b = blend[u]
p = rng.standard_normal((n, d))
p /= np.linalg.norm(p, axis=-1, keepdims=True) + 1e-8
e = np.sqrt(max(1 - b * b, 0.0)) * p + b * svecs[scenes[u]][None, :]
e /= np.linalg.norm(e, axis=-1, keepdims=True) + 1e-8
embs.append(e)
return torch.from_numpy(np.stack(embs, 1)).float()
class ViewNets(nn.Module):
"""Fixed random per-user nonlinear view functions g_u([kappa*s; p_u])."""
def __init__(self, d, U, seed=7):
super().__init__()
g = torch.Generator().manual_seed(seed)
self.nets = nn.ModuleList()
for _ in range(U):
l1 = nn.Linear(2 * d, 2 * d)
l2 = nn.Linear(2 * d, d)
for l in (l1, l2):
nn.init.normal_(l.weight, std=(2.0 / l.in_features) ** 0.5, generator=g)
nn.init.zeros_(l.bias)
self.nets.append(nn.Sequential(l1, nn.Tanh(), l2))
for p in self.parameters():
p.requires_grad_(False)
@torch.no_grad()
def gen(self, n, d, U, rng, kappa, shared_scene=True):
if shared_scene:
s = rng.standard_normal((n, d)) / math.sqrt(d)
s = np.repeat(s[:, None, :], U, axis=1)
else:
s = rng.standard_normal((n, U, d)) / math.sqrt(d)
p = rng.standard_normal((n, U, d)) / math.sqrt(d)
s = torch.from_numpy(s).float().to(DEVICE)
p = torch.from_numpy(p).float().to(DEVICE)
outs = []
for u in range(U):
x = torch.cat([kappa * s[:, u], p[:, u]], -1)
outs.append(F.normalize(self.nets[u](x), dim=-1))
return torch.stack(outs, 1) # (n, U, d)
def block_masks(U, d, device=DEVICE):
dpu = d // U
m = torch.zeros(U, d, device=device)
for u in range(U):
m[u, u * dpu:(u + 1) * dpu] = 1.0
return m
def channel(E, snr_db, *, phase_sigma_deg=0.0, fading="rayleigh", rician_K_dB=None,
nakagami_m=None, offsets=None, h_err_sigma=0.0, masks=None, use_masks=True):
"""Single-signal uplink. Returns dict with yI, yQ, h (true magnitude), h_hat.
E : (n, U, d) ground-truth embeddings on DEVICE.
offsets : (n, U) integer per-user timing offsets (symbols), or None.
Noise convention: per-rail noise std = sqrt(mean|y_tx|^2 / snr_lin); with
phase_sigma = 0 the model reduces exactly to the real-valued model.
"""
n, U, d = E.shape
if masks is None:
masks = block_masks(U, d, E.device)
X = E * masks[None] if use_masks else E.clone()
if offsets is not None:
Xs = torch.zeros_like(X)
offs = offsets
for u in range(U):
for o in offs[:, u].unique():
o = int(o.item())
idx = offs[:, u] == o
if o == 0:
Xs[idx, u] = X[idx, u]
else:
Xs[idx, u, o:] = X[idx, u, :d - o]
X = Xs
if fading == "rayleigh":
hI = torch.randn(n, U, device=E.device) * (0.5 ** 0.5)
hQ = torch.randn(n, U, device=E.device) * (0.5 ** 0.5)
hmag = (hI ** 2 + hQ ** 2).sqrt()
elif fading == "rician":
K = 10 ** (rician_K_dB / 10.0)
mu = math.sqrt(K / (K + 1))
sig = math.sqrt(1.0 / (2 * (K + 1)))
hI = mu + torch.randn(n, U, device=E.device) * sig
hQ = torch.randn(n, U, device=E.device) * sig
hmag = (hI ** 2 + hQ ** 2).sqrt()
elif fading == "nakagami":
m = nakagami_m
gam = torch.distributions.Gamma(m, m).sample((n, U)).to(E.device)
hmag = gam.sqrt()
else:
raise ValueError(fading)
dphi = torch.randn(n, U, device=E.device) * math.radians(phase_sigma_deg)
gI = hmag * torch.cos(dphi)
gQ = hmag * torch.sin(dphi)
yI = (gI[:, :, None] * X).sum(1)
yQ = (gQ[:, :, None] * X).sum(1)
P = (yI ** 2 + yQ ** 2).mean()
nstd = (P / (10 ** (snr_db / 10.0))).sqrt()
yI = yI + torch.randn_like(yI) * nstd
yQ = yQ + torch.randn_like(yQ) * nstd
h_hat = hmag * (1 + torch.randn_like(hmag) * h_err_sigma) if h_err_sigma > 0 else hmag
return {"yI": yI, "yQ": yQ, "h": hmag, "h_hat": h_hat, "nvar": float(nstd ** 2)}
# ---------------------------------------------------------------- decoders --
class UWCA(nn.Module):
"""User-wise cross-attention decoder on the single superimposed signal.
iq=True : keys/values read the stacked [I; Q] rails (2d input).
forward(yI, yQ, active) with active (n, U) bool or None.
"""
def __init__(self, d, U, H=4, iq=False):
super().__init__()
assert d % H == 0
self.d, self.U, self.H, self.dk = d, U, H, d // H
self.iq = iq
din = 2 * d if iq else d
self.q_vectors = nn.Parameter(torch.randn(U, d) * d ** -0.5)
self.eta = nn.Parameter(torch.ones(1))
self.W_K = nn.Linear(din, d, bias=False)
self.W_V = nn.Linear(din, d, bias=False)
self.W_O = nn.Linear(d, d, bias=False)
self.norm = nn.LayerNorm(d)
init_logits = torch.full((U, d), -3.0)
dpu = d // U
for u in range(U):
init_logits[u, u * dpu:(u + 1) * dpu] = 3.0
self.mask_logits = nn.Parameter(init_logits)
def soft_masks(self):
return torch.sigmoid(self.mask_logits)
def forward(self, yI, yQ=None, active=None, topk_mask=None, return_alpha=False):
# yI: (n, d) shared signal, or (n, U, d) per-candidate aligned copies
U, H, dk = self.U, self.H, self.dk
m = self.soft_masks() # (U, d)
if yI.dim() == 2:
n, d = yI.shape
yIc = yI[:, None, :].expand(-1, U, -1)
else:
n, _, d = yI.shape
yIc = yI
R = yIc * m[None] # (n, U, d)
if self.iq:
yQc = yQ[:, None, :].expand(-1, U, -1) if yQ.dim() == 2 else yQ
RQ = yQc * m[None]
Rin = torch.cat([R, RQ], -1) # (n, U, 2d)
else:
Rin = R
K = self.W_K(Rin).view(n, U, H, dk) # (n, Uk, H, dk)
V = self.W_V(Rin).view(n, U, H, dk)
Q = self.q_vectors.view(U, H, dk) # (Uq, H, dk)
scores = torch.einsum("qhk,nihk->nqhi", Q, K) * self.eta / dk ** 0.5
if active is not None: # (n, U) bool
scores = scores.masked_fill(~active[:, None, None, :], -1e9)
if topk_mask is not None: # (U, U) bool keep
scores = scores.masked_fill(~topk_mask[None, :, None, :], -1e9)
alpha = F.softmax(scores, dim=-1) # (n, Uq, H, Uk)
ctx = torch.einsum("nqhi,nihk->nqhk", alpha, V).reshape(n, U, d)
own = yIc * m[None]
out = F.normalize(self.norm(self.W_O(ctx) + own), dim=-1)
if return_alpha:
return out, alpha.mean(dim=(0, 2))
return out
class Encoder(nn.Module):
"""Trainable semantic encoder for the end-to-end experiment."""
def __init__(self, d):
super().__init__()
self.net = nn.Sequential(nn.Linear(d, 2 * d), nn.LayerNorm(2 * d),
nn.GELU(), nn.Linear(2 * d, d))
def forward(self, x):
return F.normalize(self.net(x), dim=-1)
def ofdma_decode(yI, masks):
return F.normalize(yI[:, None, :] * masks[None], dim=-1)
def tdma_proj_decode(E, snr_db, rng_t):
"""Orthogonal scheme with an arbitrary (random orthonormal) d/U-dim
projection per user instead of coordinate masks: z_u = h_u P_u e_u + n."""
n, U, d = E.shape
dpu = d // U
Q, _ = torch.linalg.qr(torch.randn(d, d, generator=rng_t).to(E.device))
outs = []
hmag = (torch.randn(n, U, device=E.device) ** 2 +
torch.randn(n, U, device=E.device) ** 2).sqrt() * 0.5 ** 0.5
snr_lin = 10 ** (snr_db / 10.0)
for u in range(U):
P = Q[u * dpu:(u + 1) * dpu] # (dpu, d)
z = hmag[:, u:u + 1] * (E[:, u] @ P.T) # (n, dpu)
nstd = (z.pow(2).mean() / snr_lin).sqrt()
z = z + torch.randn_like(z) * nstd
outs.append(F.normalize(z @ P, dim=-1))
return torch.stack(outs, 1)
def lmmse_decode(yI, h, nvar, B, masks, genie=True):
"""Closed-form linear MMSE on the block model. C_vv = I/d, C_uv = B_uv I/d.
Block-diagonal C_y => per-block Wiener weights w_uv = h_v B_uv/d / (h_v^2/d + nvar).
genie=False zeroes the cross terms (correlation-blind)."""
n, d = yI.shape
U = masks.shape[0]
Bm = torch.as_tensor(B, dtype=torch.float32, device=yI.device)
if not genie:
Bm = torch.eye(U, device=yI.device)
yb = yI[:, None, :] * masks[None] # (n, Uv, d) block pieces
w = (h[:, None, :] * Bm[None] / d) / (h[:, None, :] ** 2 / d + nvar) # (n,Uu,Uv)
est = torch.einsum("nuv,nvd->nud", w, yb)
return F.normalize(est, dim=-1)
def noma_sic_decode(E, snr_db):
"""Full-band power-domain NOMA with SIC (no masks)."""
n, U, d = E.shape
pa = torch.tensor([0.40, 0.30, 0.20, 0.10], device=E.device)[:U]
pa = pa / pa.sum()
h = (torch.randn(n, U, 1, device=E.device) ** 2 +
torch.randn(n, U, 1, device=E.device) ** 2).sqrt() * 0.5 ** 0.5
y = (E * pa.sqrt()[None, :, None] * h).sum(1)
nstd = (y.pow(2).mean() / 10 ** (snr_db / 10.0)).sqrt()
y = y + torch.randn(n, d, device=E.device) * nstd
order = torch.argsort(pa, descending=True)
res = y.clone()
out = torch.zeros_like(E)
for ui in order:
u = int(ui.item())
eh = F.normalize(res / (h[:, u] + 1e-8), dim=-1)
out[:, u] = eh
res = res - h[:, u] * pa[u].sqrt() * eh
return out
# ------------------------------------------------------------------ losses --
def semantic_loss(Ehat, E, lam=0.1, active=None):
cos = (Ehat * E).sum(-1)
if active is not None:
distortion = ((1 - cos) * active).sum() / active.sum()
else:
distortion = (1 - cos).mean()
U = Ehat.shape[1]
emb = Ehat.mean(0)
ec = emb - emb.mean(1, keepdim=True)
en = F.normalize(ec, dim=1)
C = en @ en.T
off = C[~torch.eye(U, dtype=torch.bool, device=Ehat.device)].abs().mean()
return distortion + lam * off
def ser(Ehat, E, tau=0.45, active=None):
bad = ((Ehat * E).sum(-1) < tau).float()
if active is not None:
return float((bad * active).sum() / active.sum())
return float(bad.mean())
def mean_cos(Ehat, E, active=None):
c = (Ehat * E).sum(-1)
if active is not None:
return float((c * active).sum() / active.sum())
return float(c.mean())
def sample_corr(A, Bt):
"""Mean per-sample Pearson correlation across the d dims of two (n,d) tensors."""
Ac = A - A.mean(-1, keepdim=True)
Bc = Bt - Bt.mean(-1, keepdim=True)
num = (Ac * Bc).sum(-1)
den = Ac.norm(dim=-1) * Bc.norm(dim=-1) + 1e-9
return float((num / den).mean())
# ----------------------------------------------------------------- training --
def train_multitask(model, gen_fn, tasks, epochs=300, batch=64, lam=0.1,
outer_lr=1e-3, mask_lr_mult=100.0, log_every=50, tag="",
encoder=None, extra_loss=None, log_state=None):
"""Multi-SNR/-condition aggregated training (the manuscript's outer objective
without inner adaptation). gen_fn(batch) -> E ground truth on DEVICE.
tasks: list of dicts of channel kwargs incl. 'snr_db'."""
params = []
mask_p = [p for nm, p in model.named_parameters() if "mask_logits" in nm]
other = [p for nm, p in model.named_parameters() if "mask_logits" not in nm]
params = [{"params": other, "lr": outer_lr},
{"params": mask_p, "lr": outer_lr * mask_lr_mult}]
if encoder is not None:
params.append({"params": encoder.parameters(), "lr": outer_lr})
opt = torch.optim.Adam(params)
hist = []
for ep in range(1, epochs + 1):
loss_acc = 0.0
opt.zero_grad()
for t in tasks:
E = gen_fn(batch)
if encoder is not None:
n, U, d = E.shape
E = encoder(E.reshape(-1, d)).reshape(n, U, d)
ch = channel(E, **t)
Eh = model(ch["yI"], ch["yQ"])
L = semantic_loss(Eh, E, lam)
if extra_loss is not None:
L = L + extra_loss(E)
loss_acc += L
(loss_acc / len(tasks)).backward()
nn.utils.clip_grad_norm_(model.parameters(), 5.0)
opt.step()
hist.append(float(loss_acc) / len(tasks))
if log_state is not None:
gn = sum(float(p.grad.norm()) ** 2 for p in model.parameters()
if p.grad is not None) ** 0.5
log_state.append({"ep": ep, "eta": float(model.eta), "gnorm": gn})
if ep % log_every == 0:
print(f" [{tag}] ep {ep}/{epochs} loss={hist[-1]:.4f}", flush=True)
return hist
def fomaml_train(model, gen_fn, tasks, epochs=200, batch=64, lam=0.1,
inner_lr=0.01, inner_steps=5, outer_lr=1e-3,
tasks_per_step=8, log_every=25, tag="fomaml", log_state=None,
rng=None):
"""Proper first-order MAML: inner SGD on support, outer update from query
gradients evaluated at the adapted parameters."""
opt = torch.optim.Adam(model.parameters(), lr=outer_lr)
rng = rng or np.random.default_rng(0)
hist = []
names = [nm for nm, _ in model.named_parameters()]
for ep in range(1, epochs + 1):
idx = rng.choice(len(tasks), size=min(tasks_per_step, len(tasks)),
replace=False)
grads = {nm: torch.zeros_like(p) for nm, p in model.named_parameters()}
qloss_acc = 0.0
for ti in idx:
t = tasks[ti]
adapted = copy.deepcopy(model)
iopt = torch.optim.SGD(adapted.parameters(), lr=inner_lr)
for _ in range(inner_steps):
E = gen_fn(batch)
ch = channel(E, **t)
L = semantic_loss(adapted(ch["yI"], ch["yQ"]), E, lam)
iopt.zero_grad(); L.backward(); iopt.step()
E = gen_fn(batch)
ch = channel(E, **t)
qL = semantic_loss(adapted(ch["yI"], ch["yQ"]), E, lam)
adapted.zero_grad(); qL.backward()
for nm, p in adapted.named_parameters():
if p.grad is not None:
grads[nm] += p.grad
qloss_acc += float(qL)
opt.zero_grad()
for nm, p in model.named_parameters():
p.grad = grads[nm] / len(idx)
nn.utils.clip_grad_norm_(model.parameters(), 5.0)
opt.step()
hist.append(qloss_acc / len(idx))
if log_state is not None:
gn = sum(float(g.norm()) ** 2 for g in grads.values()) ** 0.5 / len(idx)
log_state.append({"ep": ep, "eta": float(model.eta), "gnorm": gn})
if ep % log_every == 0:
print(f" [{tag}] ep {ep}/{epochs} qloss={hist[-1]:.4f} "
f"eta={float(model.eta):.3f}", flush=True)
return hist
def adapt(model, gen_fn, task, steps=5, inner_lr=0.01, batch=64, lam=0.1):
adapted = copy.deepcopy(model)
iopt = torch.optim.SGD(adapted.parameters(), lr=inner_lr)
for _ in range(steps):
E = gen_fn(batch)
ch = channel(E, **task)
L = semantic_loss(adapted(ch["yI"], ch["yQ"]), E, lam)
iopt.zero_grad(); L.backward(); iopt.step()
return adapted
# --------------------------------------------------------------- evaluation --
@torch.no_grad()
def eval_scheme(scheme, gen_fn, task, n_mc=200, batch=64, tau=0.45, model=None,
B=None, masks=None, encoder=None, rng_t=None, active_fn=None,
topk_mask=None):
"""Returns (ser, cos) for one task/channel config."""
s_acc = c_acc = 0.0
for _ in range(n_mc):
E = gen_fn(batch)
if encoder is not None:
n, U, d = E.shape
E = encoder(E.reshape(-1, d)).reshape(n, U, d)
act = active_fn(E.shape[0]) if active_fn is not None else None
if scheme == "uwca":
ch = channel(E, **task)
Eh = model(ch["yI"], ch["yQ"], active=act, topk_mask=topk_mask)
elif scheme == "ofdma":
ch = channel(E, **task)
m = masks if masks is not None else block_masks(E.shape[1], E.shape[2],
E.device)
Eh = ofdma_decode(ch["yI"], m)
elif scheme == "sfdma":
ch = channel(E, **task)
m = masks
Eh = torch.stack([F.normalize((ch["yI"] * m[u]) /
(ch["h_hat"][:, u:u + 1] + 1e-8), dim=-1)
for u in range(m.shape[0])], 1)
elif scheme == "noma":
Eh = noma_sic_decode(E, task["snr_db"])
elif scheme in ("lmmse_genie", "lmmse_blind"):
ch = channel(E, **task)
Eh = lmmse_decode(ch["yI"], ch["h_hat"], ch["nvar"], B, masks,
genie=(scheme == "lmmse_genie"))
elif scheme == "tdma_proj":
Eh = tdma_proj_decode(E, task["snr_db"], rng_t)
else:
raise ValueError(scheme)
s_acc += ser(Eh, E, tau, act)
c_acc += mean_cos(Eh, E, act)
return s_acc / n_mc, c_acc / n_mc
def save_json(name, obj):
p = DATA / name
with open(p, "w") as f:
json.dump(obj, f, indent=1)
print(f"saved -> {p}", flush=True)
def set_seed(seed=42):
torch.manual_seed(seed)
np.random.seed(seed)
return np.random.default_rng(seed)