Files
uwca-semantic-mac/rev2/e2_phase_iui.py
T

130 lines
4.7 KiB
Python
Executable File

"""E2 — Full complex-baseband phase-error model with inter-user leakage (R1.1)
plus CSI amplitude-error robustness (R2.3).
Three evaluation models on HIGH:
scalar : real channel, per-user cos(dphi) attenuation only (old model)
complex-I : full complex superposition; decoder reads the in-phase rail only
complex-IQ: full complex superposition; decoder reads both rails (proposed)
Two trained decoders (phase-augmented training, sigma_phi ~ U[0,20] deg):
m_real (iq=False) and m_iq (iq=True).
Also evaluates a decoder trained at sigma_phi=0 to expose training mismatch,
and a CSI amplitude error sweep for SFDMA (divides by h) vs UWCA (no CSI).
"""
import numpy as np
import torch
import lib
from lib import (SCENARIOS, UWCA, DEVICE, block_masks, channel, eval_scheme,
gen_embeddings, mean_cos, save_json, ser, set_seed,
train_multitask)
rng = set_seed(42)
d, U, H = 64, 4, 4
masks = block_masks(U, d)
scen = SCENARIOS["HIGH"]
def gen(n):
return gen_embeddings(n, d, U, rng, scen).to(DEVICE)
snrs = [0.0, 4.0, 8.0, 12.0, 16.0, 20.0]
rng_ph = np.random.default_rng(3)
def aug_tasks():
return [{"snr_db": s, "phase_sigma_deg": float(rng_ph.uniform(0, 20))}
for s in snrs]
class AugTaskList:
"""List-like view that resamples phase residuals each epoch."""
def __init__(self):
self._t = aug_tasks()
self._n = 0
def __len__(self):
return len(self._t)
def __iter__(self):
self._n += 1
self._t = aug_tasks()
return iter(self._t)
m_real = UWCA(d, U, H, iq=False).to(DEVICE)
train_multitask(m_real, gen, AugTaskList(), epochs=300, tag="E2-real-aug")
m_iq = UWCA(d, U, H, iq=True).to(DEVICE)
train_multitask(m_iq, gen, AugTaskList(), epochs=300, tag="E2-iq-aug")
m_zero = UWCA(d, U, H, iq=False).to(DEVICE)
train_multitask(m_zero, gen, [{"snr_db": s} for s in snrs], epochs=300,
tag="E2-zerophase")
torch.save(m_iq.state_dict(), lib.DATA / "e2_uwca_iq.pt")
sig_grid = [0, 5, 10, 15, 20, 30]
out = {"sigma_phi_deg": sig_grid, "snr_eval": [10.0, 20.0], "curves": {}}
@torch.no_grad()
def run(model, sig, snr, mode):
s_acc = c_acc = 0.0
n_mc = 200
for _ in range(n_mc):
E = gen(64)
if mode == "scalar":
# magnitude attenuation only: fold cos(dphi) into the gain, Q rail ignored
ch = channel(E, snr_db=snr, phase_sigma_deg=sig)
# scalar model == complex-I when masks are disjoint; emulate the
# old analytic model by discarding the Q rail entirely
Eh = model(ch["yI"], torch.zeros_like(ch["yQ"]))
elif mode == "cI":
ch = channel(E, snr_db=snr, phase_sigma_deg=sig)
Eh = model(ch["yI"], torch.zeros_like(ch["yQ"])) \
if not model.iq else model(ch["yI"], ch["yQ"])
elif mode == "cIQ":
ch = channel(E, snr_db=snr, phase_sigma_deg=sig)
Eh = model(ch["yI"], ch["yQ"])
s_acc += ser(Eh, E)
c_acc += mean_cos(Eh, E)
return s_acc / n_mc, c_acc / n_mc
for label, model, mode in [("scalar_augtrain", m_real, "scalar"),
("complexI_augtrain", m_real, "cI"),
("complexIQ_iqtrain", m_iq, "cIQ"),
("complexI_zerotrain", m_zero, "cI")]:
cur = {}
for snr in out["snr_eval"]:
cur[str(snr)] = {"ser": [], "cos": []}
for sig in sig_grid:
s, c = run(model, sig, snr, mode)
cur[str(snr)]["ser"].append(s)
cur[str(snr)]["cos"].append(c)
print(f"[E2] {label} snr={snr}: SER={cur[str(snr)]['ser']}", flush=True)
out["curves"][label] = cur
# soft-mask overlap of the trained decoders (quantifies the IUI channel)
with torch.no_grad():
for label, model in [("m_real", m_real), ("m_iq", m_iq)]:
m = model.soft_masks()
ov = (m @ m.T) / (m.norm(dim=1, keepdim=True) * m.norm(dim=1) + 1e-9)
off = ov[~torch.eye(U, dtype=torch.bool, device=ov.device)]
out[f"mask_overlap_{label}"] = {"mean": float(off.mean()),
"max": float(off.max())}
# CSI amplitude error: SFDMA (uses h) vs UWCA (no explicit CSI), sigma_phi=10
h_grid = [0.0, 0.05, 0.1, 0.2]
csi = {"h_err": h_grid, "uwca_ser": [], "sfdma_ser": []}
for he in h_grid:
t = {"snr_db": 10.0, "phase_sigma_deg": 10.0, "h_err_sigma": he}
s_u, _ = eval_scheme("uwca", gen, t, n_mc=200, model=m_iq, masks=masks)
s_f, _ = eval_scheme("sfdma", gen, t, n_mc=200, masks=masks)
csi["uwca_ser"].append(s_u)
csi["sfdma_ser"].append(s_f)
print(f"[E2-CSI] h_err={he}: UWCA {s_u:.3f} SFDMA {s_f:.3f}", flush=True)
out["csi_error"] = csi
save_json("e2_phase_iui.json", out)