check_family_enum.py measures the attack the manuscript now states in Section III-A: the winning correlation is an index-free verifier, so ranking the 63 non-constant Walsh rows by mean winning correlation recovers the user set from one frame in 0.905 of 200 trials at 10 dB and from four frames in 0.990, using nothing outside the stated threat model. Under the invariance refresh it recovers it in none, because the entry permutation relabels the codebook the adversary must align against. V8 and V9 read the trained codebook through main_model(), which retrains on every call, and a codebook trained on CUDA is not the one trained on CPU. The shipped verify_math.csv therefore read PASS here and FAIL for anyone running this package without a GPU. model_main.pt is 7 KB and fixes the codebook, which is what both checks are about; delete it to retrain. V1-V11 now pass on both. New checks: V10, the format-matched OMA reference Section VI-B quotes, and V11, the closed-form against Monte Carlo comparison the manuscript claimed and never stored. V3a's bias-linearity result was computed and printed but never written to the CSV, so the one linearity claim the paper quotes was the one this package could not show. check_consistency.py gains 21 assertions, covering five data files that no assertion read (users, csi, semantic, cov_attack, sec_jam) and the trend claims it structurally could not see, since it compared values and not shapes. README: the figure map named stages that do not write the artifacts they list, so following it did not reproduce Figs. 4 and 6; the reproduction block was five scripts short; and the refresh numbers were from a superseded run (nearly three, 15.0 to 64.8 bits) against the manuscript's 2.3 and 23.8 to 364.6.
427 lines
18 KiB
Python
427 lines
18 KiB
Python
"""Shared library for the scalable shared embedding (SSE) letter.
|
|
|
|
System model (real-vector convention, declared in the paper):
|
|
d-dimensional real embedding frame, U users, single transmitter.
|
|
Proposed SSE: the frame is split into P periods of length L = d/P and
|
|
one unit codebook B in R^{Vu x L} is reused in every period, so the
|
|
vocabulary size is V = Vu^P while the codebook stores only Vu*L
|
|
numbers. Index v maps to base-Vu digits (i_1,...,i_P).
|
|
Per-user periodic masks mu_u in R^L are repeated over the P periods.
|
|
Tx: y = (1/c) * sum_u e(s_u) .* m_u, c fixes unit average frame power.
|
|
Channel: user u sees h_u * y + n, h_u^2 ~ Exp(1) (Rayleigh magnitude,
|
|
known at the receiver), n ~ N(0, sigma^2 I), sigma^2 = 1/(d*snr).
|
|
Rx u: equalize by h_u, per period correlate with the masked unit
|
|
codewords b_i .* mu_u and take the argmax digit.
|
|
|
|
Device: CUDA when available (run under WSL), CPU fallback.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import math
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import torch
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
DATA = ROOT / "data"
|
|
FIG = ROOT / "fig"
|
|
DATA.mkdir(exist_ok=True)
|
|
FIG.mkdir(exist_ok=True)
|
|
|
|
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
|
|
# ----------------------------------------------------------------------
|
|
# global configuration
|
|
# ----------------------------------------------------------------------
|
|
D = 256 # embedding dimension (real)
|
|
U = 4 # users
|
|
VU = 16 # unit codebook size
|
|
P_MAX = 4 # periods for the main configuration, V = 16^4 = 65536
|
|
SEED = 1
|
|
|
|
TRAIN_ITERS = 4000
|
|
TRAIN_BATCH = 256
|
|
TRAIN_SNR_DB = (0.0, 20.0)
|
|
LR = 3e-3
|
|
|
|
|
|
def set_seed(seed: int = SEED) -> None:
|
|
torch.manual_seed(seed)
|
|
np.random.seed(seed)
|
|
if torch.cuda.is_available():
|
|
torch.cuda.manual_seed_all(seed)
|
|
|
|
|
|
def snr_to_sigma2(snr_db: torch.Tensor | float,
|
|
d: int = D) -> torch.Tensor | float:
|
|
"""Per-dimension noise variance for unit frame power and E[h^2]=1.
|
|
Pass the model's actual frame dimension d when it differs from the
|
|
module default, otherwise the effective SNR shifts with d."""
|
|
snr = 10.0 ** (torch.as_tensor(snr_db, dtype=torch.float64) / 10.0)
|
|
return (1.0 / (d * snr)).float()
|
|
|
|
|
|
def rayleigh_gain(shape, device=DEVICE) -> torch.Tensor:
|
|
"""|g| with g ~ CN(0,1): h^2 ~ Exp(1), E[h^2] = 1."""
|
|
u = torch.rand(shape, device=device).clamp_min(1e-12)
|
|
return torch.sqrt(-torch.log(u))
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# proposed scalable shared embedding
|
|
# ----------------------------------------------------------------------
|
|
class SSE(torch.nn.Module):
|
|
"""Periodic unit-codebook shared embedding with per-user periodic masks."""
|
|
|
|
def __init__(self, P: int = P_MAX, vu: int = VU, d: int = D, users: int = U):
|
|
super().__init__()
|
|
assert d % P == 0, "embedding dimension must split into P periods"
|
|
self.P, self.vu, self.d, self.users = P, vu, d, users
|
|
self.L = d // P
|
|
self.B = torch.nn.Parameter(torch.randn(vu, self.L) / math.sqrt(self.L))
|
|
self.W = torch.nn.Parameter(torch.randn(users, self.L) / math.sqrt(self.L))
|
|
self.logit_scale = torch.nn.Parameter(torch.tensor(2.0))
|
|
# transmit power normalizer, calibrated after training (buffer)
|
|
self.register_buffer("c", torch.tensor(1.0))
|
|
|
|
@property
|
|
def V(self) -> int:
|
|
return self.vu ** self.P
|
|
|
|
def unit_codebook(self) -> torch.Tensor:
|
|
return self.B / self.B.norm(dim=1, keepdim=True).clamp_min(1e-8)
|
|
|
|
def masks(self) -> torch.Tensor:
|
|
return self.W / self.W.norm(dim=1, keepdim=True).clamp_min(1e-8) * math.sqrt(self.L)
|
|
|
|
def tx_frame(self, digits: torch.Tensor) -> torch.Tensor:
|
|
"""digits: (N, U, P) ints -> unnormalized tx frame (N, P, L)."""
|
|
Bn = self.unit_codebook()
|
|
e = Bn[digits] / math.sqrt(self.P) # (N,U,P,L)
|
|
m = self.masks() # (U,L)
|
|
x = e * m[None, :, None, :]
|
|
return x.sum(dim=1) # (N,P,L)
|
|
|
|
def calibrate_power(self, n: int = 65536) -> None:
|
|
with torch.no_grad():
|
|
digits = torch.randint(self.vu, (n, self.users, self.P), device=self.B.device)
|
|
y = self.tx_frame(digits)
|
|
self.c.fill_(float(y.pow(2).sum(dim=(1, 2)).mean().sqrt()))
|
|
|
|
def scores(self, r: torch.Tensor) -> torch.Tensor:
|
|
"""r: equalized rx frame (N,P,L) -> scores (N,U,P,Vu)."""
|
|
Bn = self.unit_codebook() # (Vu,L)
|
|
m = self.masks() # (U,L)
|
|
cand = Bn[None, :, :] * m[:, None, :] # (U,Vu,L)
|
|
return torch.einsum("npl,uvl->nupv", r, cand)
|
|
|
|
def forward(self, digits: torch.Tensor, snr_db: torch.Tensor,
|
|
h: torch.Tensor | None = None):
|
|
"""digits (N,U,P), snr_db (N,) -> per-user scores and rx frames."""
|
|
N = digits.shape[0]
|
|
y = self.tx_frame(digits) / self.c # (N,P,L)
|
|
if h is None:
|
|
h = rayleigh_gain((N, self.users), device=y.device)
|
|
sigma = snr_to_sigma2(snr_db, self.d).to(y.device).sqrt() # (N,)
|
|
n = torch.randn(N, self.users, self.P, self.L, device=y.device)
|
|
y_rx = h[:, :, None, None] * y[:, None] + sigma[:, None, None, None] * n
|
|
r = y_rx / h[:, :, None, None].clamp_min(1e-6) # equalized (N,U,P,L)
|
|
Bn = self.unit_codebook()
|
|
m = self.masks()
|
|
cand = Bn[None, :, :] * m[:, None, :] # (U,Vu,L)
|
|
scores = torch.einsum("nupl,uvl->nupv", r, cand)
|
|
return scores
|
|
|
|
def n_params(self) -> int:
|
|
return self.B.numel() + self.W.numel()
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# conventional scheme: full unstructured codebook (prior shared embedding)
|
|
# ----------------------------------------------------------------------
|
|
class FullCodebook(torch.nn.Module):
|
|
"""Directly trained V x d codebook with full-d per-user masks."""
|
|
|
|
def __init__(self, V: int, d: int = D, users: int = U):
|
|
super().__init__()
|
|
self.Vc, self.d, self.users = V, d, users
|
|
self.E = torch.nn.Parameter(torch.randn(V, d) / math.sqrt(d))
|
|
self.W = torch.nn.Parameter(torch.randn(users, d) / math.sqrt(d))
|
|
self.logit_scale = torch.nn.Parameter(torch.tensor(2.0))
|
|
self.register_buffer("c", torch.tensor(1.0))
|
|
|
|
@property
|
|
def V(self) -> int:
|
|
return self.Vc
|
|
|
|
def codebook(self) -> torch.Tensor:
|
|
return self.E / self.E.norm(dim=1, keepdim=True).clamp_min(1e-8)
|
|
|
|
def masks(self) -> torch.Tensor:
|
|
return self.W / self.W.norm(dim=1, keepdim=True).clamp_min(1e-8) * math.sqrt(self.d)
|
|
|
|
def tx_frame(self, idx: torch.Tensor) -> torch.Tensor:
|
|
En = self.codebook()
|
|
e = En[idx] # (N,U,d)
|
|
m = self.masks()
|
|
return (e * m[None]).sum(dim=1) # (N,d)
|
|
|
|
def calibrate_power(self, n: int = 65536) -> None:
|
|
with torch.no_grad():
|
|
idx = torch.randint(self.Vc, (n, self.users), device=self.E.device)
|
|
y = self.tx_frame(idx)
|
|
self.c.fill_(float(y.pow(2).sum(dim=1).mean().sqrt()))
|
|
|
|
def forward(self, idx: torch.Tensor, snr_db: torch.Tensor,
|
|
h: torch.Tensor | None = None):
|
|
N = idx.shape[0]
|
|
y = self.tx_frame(idx) / self.c # (N,d)
|
|
if h is None:
|
|
h = rayleigh_gain((N, self.users), device=y.device)
|
|
sigma = snr_to_sigma2(snr_db).to(y.device).sqrt()
|
|
n = torch.randn(N, self.users, self.d, device=y.device)
|
|
y_rx = h[:, :, None] * y[:, None] + sigma[:, None, None] * n
|
|
r = y_rx / h[:, :, None].clamp_min(1e-6) # (N,U,d)
|
|
En = self.codebook()
|
|
m = self.masks()
|
|
cand = En[None, :, :] * m[:, None, :] # (U,V,d)
|
|
scores = torch.einsum("nud,uvd->nuv", r, cand)
|
|
return scores
|
|
|
|
def n_params(self) -> int:
|
|
return self.E.numel() + self.W.numel()
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# training
|
|
# ----------------------------------------------------------------------
|
|
def train_sse(model: SSE, iters: int = TRAIN_ITERS, batch: int = TRAIN_BATCH,
|
|
lr: float = LR, log_every: int = 0, eval_snr: float = 10.0,
|
|
eval_frames: int = 20000, seed: int = SEED):
|
|
"""Digit-wise cross entropy under superposition; cost independent of V."""
|
|
set_seed(seed)
|
|
model.to(DEVICE)
|
|
opt = torch.optim.Adam(model.parameters(), lr=lr)
|
|
ce = torch.nn.CrossEntropyLoss()
|
|
curve = []
|
|
t0 = time.time()
|
|
for it in range(1, iters + 1):
|
|
digits = torch.randint(model.vu, (batch, model.users, model.P), device=DEVICE)
|
|
snr = torch.empty(batch).uniform_(*TRAIN_SNR_DB)
|
|
model.calibrate_power(8192)
|
|
scores = model(digits, snr) * model.logit_scale.exp()
|
|
loss = ce(scores.reshape(-1, model.vu), digits.reshape(-1))
|
|
opt.zero_grad(); loss.backward(); opt.step()
|
|
if log_every and (it % log_every == 0 or it == 1):
|
|
ser = eval_ser_sse(model, [eval_snr], frames=eval_frames)[0]
|
|
curve.append((it, time.time() - t0, float(loss.detach()), ser))
|
|
model.calibrate_power()
|
|
return curve
|
|
|
|
|
|
def train_full(model: FullCodebook, iters: int = TRAIN_ITERS,
|
|
batch: int = TRAIN_BATCH, lr: float = LR, log_every: int = 0,
|
|
eval_snr: float = 10.0, eval_frames: int = 20000,
|
|
seed: int = SEED):
|
|
set_seed(seed)
|
|
model.to(DEVICE)
|
|
opt = torch.optim.Adam(model.parameters(), lr=lr)
|
|
ce = torch.nn.CrossEntropyLoss()
|
|
curve = []
|
|
t0 = time.time()
|
|
for it in range(1, iters + 1):
|
|
idx = torch.randint(model.Vc, (batch, model.users), device=DEVICE)
|
|
snr = torch.empty(batch).uniform_(*TRAIN_SNR_DB)
|
|
model.calibrate_power(8192)
|
|
scores = model(idx, snr) * model.logit_scale.exp()
|
|
loss = ce(scores.reshape(-1, model.Vc), idx.reshape(-1))
|
|
opt.zero_grad(); loss.backward(); opt.step()
|
|
if log_every and (it % log_every == 0 or it == 1):
|
|
ser = eval_ser_full(model, [eval_snr], frames=eval_frames)[0]
|
|
curve.append((it, time.time() - t0, float(loss.detach()), ser))
|
|
model.calibrate_power()
|
|
return curve
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# evaluation (Monte Carlo)
|
|
# ----------------------------------------------------------------------
|
|
@torch.no_grad()
|
|
def eval_ser_sse(model: SSE, snr_list, frames: int = 2_000_000,
|
|
chunk: int = 100_000, seed: int = 777):
|
|
"""Frame (vocabulary-symbol) error rate: any wrong digit is an error."""
|
|
model.eval().to(DEVICE)
|
|
out = []
|
|
for snr_db in snr_list:
|
|
g = torch.Generator(device="cpu").manual_seed(seed + int(10 * snr_db))
|
|
err = tot = 0
|
|
for n0 in range(0, frames, chunk):
|
|
n = min(chunk, frames - n0)
|
|
digits = torch.randint(model.vu, (n, model.users, model.P),
|
|
generator=g).to(DEVICE)
|
|
snr = torch.full((n,), float(snr_db))
|
|
scores = model(digits, snr)
|
|
wrong = (scores.argmax(-1) != digits).any(dim=2) # (n,U)
|
|
err += int(wrong.sum()); tot += n * model.users
|
|
out.append(err / tot)
|
|
return out
|
|
|
|
|
|
@torch.no_grad()
|
|
def eval_ser_full(model: FullCodebook, snr_list, frames: int = 2_000_000,
|
|
chunk: int = 50_000, seed: int = 777):
|
|
model.eval().to(DEVICE)
|
|
if model.Vc >= 4096:
|
|
chunk = max(512, (1 << 23) // model.Vc)
|
|
out = []
|
|
for snr_db in snr_list:
|
|
g = torch.Generator(device="cpu").manual_seed(seed + int(10 * snr_db))
|
|
err = tot = 0
|
|
for n0 in range(0, frames, chunk):
|
|
n = min(chunk, frames - n0)
|
|
idx = torch.randint(model.Vc, (n, model.users), generator=g).to(DEVICE)
|
|
snr = torch.full((n,), float(snr_db))
|
|
scores = model(idx, snr)
|
|
err += int((scores.argmax(-1) != idx).sum()); tot += n * model.users
|
|
out.append(err / tot)
|
|
return out
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# conventional digital scheme: OMA with BPSK per real dimension (QPSK
|
|
# per complex subcarrier under the declared real-imaginary stacking)
|
|
# ----------------------------------------------------------------------
|
|
def oma_ser(snr_db_list, bits: int = 16, n_grid: int = 200_000):
|
|
"""Semi-analytic expression under the real-dimension convention.
|
|
All bits of a vocabulary symbol share the same flat-fading gain, so
|
|
SER = E_h[1 - (1 - Q(h sqrt(snr)))^bits] with h^2 ~ Exp(1); the
|
|
expectation is evaluated by numerical integration on a dense grid."""
|
|
x = (np.arange(n_grid) + 0.5) / n_grid # uniform quantiles
|
|
h = np.sqrt(-np.log(1.0 - x)) # inverse-CDF transform
|
|
out = []
|
|
for s in snr_db_list:
|
|
g = 10.0 ** (s / 10.0)
|
|
q = 0.5 * np.array([math.erfc(v / math.sqrt(2.0)) for v in
|
|
np.clip(h * math.sqrt(g), 0, 38)])
|
|
out.append(float(np.mean(1.0 - (1.0 - q) ** bits)))
|
|
return out
|
|
|
|
|
|
def oma_ser_orth(snr_db_list, P: int = 4, vu: int = 16, L: int = 64,
|
|
n_h: int = 20_000, n_z: int = 2001):
|
|
"""Format-matched OMA reference.
|
|
|
|
The binary reference of oma_ser_keylen spends 16 of its L exclusive
|
|
dimensions on antipodal bits, a one-bit-per-dimension format inside
|
|
a log2(V)/L = 0.25 bit-per-dimension budget. The better uncoded use
|
|
of the same allocation is the format the proposed scheme itself
|
|
uses: P orthogonal decisions among vu candidates, each over L/P
|
|
exclusive dimensions, which needs exactly vu = L/P of them and so
|
|
fits the allocation with nothing to spare.
|
|
|
|
Energy accounting matches oma_ser, where one unit of energy on a
|
|
dimension gives 2Es/N0 = snr, so an L-dimension user spending its L
|
|
units on P symbols puts L/P units in each. Given the fading gain h
|
|
the correct matched-filter output is N(h sqrt(Es), N0/2) against
|
|
vu-1 outputs N(0, N0/2), so a digit is right with probability
|
|
E_z[Phi(z + h sqrt((L/P) snr))^(vu-1)] and the index is right when
|
|
all P digits are.
|
|
"""
|
|
from scipy.special import log_ndtr
|
|
x = (np.arange(n_h) + 0.5) / n_h
|
|
h = np.sqrt(-np.log(1.0 - x)) # h^2 ~ Exp(1)
|
|
z = np.linspace(-8.0, 8.0, n_z)
|
|
phi = np.exp(-0.5 * z * z) / math.sqrt(2.0 * math.pi)
|
|
out = []
|
|
for s in snr_db_list:
|
|
a = h * math.sqrt((L / P) * 10.0 ** (s / 10.0))
|
|
pc = np.empty_like(a)
|
|
for i in range(0, a.size, 2048): # bound the working set
|
|
blk = a[i:i + 2048][:, None]
|
|
pc[i:i + 2048] = np.trapezoid(
|
|
phi * np.exp((vu - 1) * log_ndtr(z[None, :] + blk)),
|
|
z, axis=1)
|
|
out.append(float(np.mean(1.0 - pc ** P)))
|
|
return out
|
|
|
|
|
|
@torch.no_grad()
|
|
def oma_ser_mc(snr_db_list, bits: int = 16, frames: int = 2_000_000,
|
|
chunk: int = 200_000, seed: int = 777):
|
|
"""Monte Carlo check of the closed form (same channel conventions)."""
|
|
out = []
|
|
for s in snr_db_list:
|
|
sigma = math.sqrt(1.0 / (10.0 ** (s / 10.0))) # per-dim, unit Es
|
|
err = tot = 0
|
|
torch.manual_seed(seed + int(10 * s))
|
|
for n0 in range(0, frames, chunk):
|
|
n = min(chunk, frames - n0)
|
|
h = rayleigh_gain((n, 1))
|
|
b = torch.randint(0, 2, (n, bits), device=DEVICE) * 2.0 - 1.0
|
|
y = h * b + sigma * torch.randn(n, bits, device=DEVICE)
|
|
wrong = ((y * b) < 0).any(dim=1)
|
|
err += int(wrong.sum()); tot += n
|
|
out.append(err / tot)
|
|
return out
|
|
|
|
|
|
# ----------------------------------------------------------------------
|
|
# analysis: union bound with residual-interference Gaussian approximation
|
|
# ----------------------------------------------------------------------
|
|
@torch.no_grad()
|
|
def sse_union_bound(model: SSE, snr_db_list, n_mc_h: int = 200_000):
|
|
"""P_e <= 1 - (1 - min(1, sum_pairs)) ... digit union bound averaged
|
|
over the Rayleigh gain, with measured residual interference treated
|
|
as additional Gaussian noise (approximation stated in the paper)."""
|
|
model.eval().to(DEVICE)
|
|
Bn = model.unit_codebook()
|
|
m = model.masks()
|
|
c = float(model.c)
|
|
# residual interference power per dimension at user u (measured)
|
|
digits = torch.randint(model.vu, (65536, model.users, model.P), device=DEVICE)
|
|
e = Bn[digits] / math.sqrt(model.P)
|
|
x = e * m[None, :, None, :] # (N,U,P,L)
|
|
y = x.sum(dim=1) / c # (N,P,L)
|
|
# per-user: signal = x_u/c, interference = (y - x_u/c)
|
|
interf_pw = []
|
|
dmin2 = []
|
|
for u in range(model.users):
|
|
su = x[:, u] / c # (N,P,L)
|
|
iu = y - su
|
|
# project interference onto the normalized candidate directions
|
|
cand = Bn * m[u][None, :] # (Vu,L)
|
|
cn = cand / cand.norm(dim=1, keepdim=True).clamp_min(1e-8)
|
|
proj = torch.einsum("npl,vl->npv", iu, cn)
|
|
interf_pw.append(float(proj.pow(2).mean()))
|
|
# pairwise distances of the scaled candidate set (tx side scaling)
|
|
cs = cand / (c * math.sqrt(model.P))
|
|
dd = torch.cdist(cs, cs)
|
|
dmin2.append(float((dd + torch.eye(model.vu, device=DEVICE) * 1e9).min() ** 2))
|
|
res = []
|
|
g = rayleigh_gain(n_mc_h)
|
|
for s in snr_db_list:
|
|
sig2 = float(snr_to_sigma2(torch.tensor(s)))
|
|
pe_users = []
|
|
for u in range(model.users):
|
|
# effective noise per dim after equalization: sig2/h^2 + interf
|
|
sig_eff2 = sig2 / g.pow(2) + interf_pw[u]
|
|
arg = torch.sqrt(torch.clamp(torch.tensor(dmin2[u], device=DEVICE)
|
|
/ (4.0 * sig_eff2), min=0.0))
|
|
q = 0.5 * torch.erfc(arg / math.sqrt(2.0))
|
|
p_digit = torch.clamp((model.vu - 1) * q, max=1.0)
|
|
p_frame = 1.0 - (1.0 - p_digit) ** model.P
|
|
pe_users.append(float(p_frame.mean()))
|
|
res.append(sum(pe_users) / len(pe_users))
|
|
return res
|
|
|
|
|
|
def write_csv(path: Path, header: list[str], rows) -> None:
|
|
with open(path, "w") as f:
|
|
f.write(",".join(header) + "\n")
|
|
for row in rows:
|
|
f.write(",".join(f"{v:.10g}" if isinstance(v, float) else str(v)
|
|
for v in row) + "\n")
|
|
print("[csv]", path)
|