Reproducibility package: UWCA semantic multiple access (TWC submission)
This commit is contained in:
Executable
+144
@@ -0,0 +1,144 @@
|
||||
"""E9 — Online relevance acquisition for sparse top-k attention (R1.2, R3.6)
|
||||
and measured selection/sorting overhead (R2.4).
|
||||
|
||||
Protocol (U=32, 8 clusters of 4, k=4):
|
||||
frames 1..3 : full attention; the BS estimates beta_hat from the decoded
|
||||
embeddings by an EWMA of pairwise cosines (no oracle knowledge)
|
||||
frames >=4 : top-k attention using beta_hat (self + k-1 best peers)
|
||||
Reports the per-frame fidelity trajectory against the oracle top-k (true
|
||||
clusters) and full attention, plus wall-clock timing of the full pipeline
|
||||
including estimation and argpartition selection for U in {8..128}.
|
||||
"""
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import lib
|
||||
from lib import (UWCA, DEVICE, block_masks, channel, gen_embeddings,
|
||||
mean_cos, save_json, semantic_loss, ser, set_seed)
|
||||
|
||||
rng = set_seed(42)
|
||||
d, H, k = 64, 4, 4
|
||||
U = 32
|
||||
G = U // k # 8 clusters of 4
|
||||
scen = {"beta_u": [0.65] * U, "scenes": [i // k for i in range(U)]}
|
||||
masks = block_masks(U, d)
|
||||
snrs = [0.0, 10.0, 20.0]
|
||||
|
||||
|
||||
def gen(n):
|
||||
return gen_embeddings(n, d, U, rng, scen).to(DEVICE)
|
||||
|
||||
|
||||
def cluster_mask():
|
||||
m = torch.zeros(U, U, dtype=torch.bool, device=DEVICE)
|
||||
for u in range(U):
|
||||
c = u // k
|
||||
m[u, c * k:(c + 1) * k] = True
|
||||
return m
|
||||
|
||||
|
||||
def topk_from_beta(bhat):
|
||||
m = torch.zeros(U, U, dtype=torch.bool, device=DEVICE)
|
||||
b = bhat.clone()
|
||||
b.fill_diagonal_(2.0) # always keep self
|
||||
idx = torch.topk(b, k, dim=1).indices
|
||||
m.scatter_(1, idx, True)
|
||||
return m
|
||||
|
||||
|
||||
print("[E9] training U=32 model (full attention)", flush=True)
|
||||
model = UWCA(d, U, H).to(DEVICE)
|
||||
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]
|
||||
opt = torch.optim.Adam([{"params": other, "lr": 1e-3},
|
||||
{"params": mask_p, "lr": 0.1}])
|
||||
for ep in range(1, 181):
|
||||
opt.zero_grad()
|
||||
loss = 0.0
|
||||
for s in snrs:
|
||||
E = gen(64)
|
||||
ch = channel(E, snr_db=s)
|
||||
Eh = model(ch["yI"], ch["yQ"])
|
||||
loss = loss + semantic_loss(Eh, E, 0.1)
|
||||
(loss / len(snrs)).backward()
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
|
||||
opt.step()
|
||||
if ep % 45 == 0:
|
||||
print(f" [E9] ep {ep}/180 loss={float(loss)/len(snrs):.4f}", flush=True)
|
||||
|
||||
# ---- online protocol trajectory at 10 dB
|
||||
T, warm, gamma_ewma = 16, 3, 0.5
|
||||
oracle = cluster_mask()
|
||||
traj = {"frame": list(range(1, T + 1)), "online_cos": [], "oracle_cos": [],
|
||||
"full_cos": [], "beta_err": []}
|
||||
bhat = torch.zeros(U, U, device=DEVICE)
|
||||
model.eval()
|
||||
Btrue = torch.zeros(U, U, device=DEVICE)
|
||||
for u in range(U):
|
||||
for v in range(U):
|
||||
if u != v and scen["scenes"][u] == scen["scenes"][v]:
|
||||
Btrue[u, v] = scen["beta_u"][u] * scen["beta_u"][v]
|
||||
|
||||
with torch.no_grad():
|
||||
for t in range(1, T + 1):
|
||||
E = gen(256)
|
||||
ch = channel(E, snr_db=10.0)
|
||||
tk = None if t <= warm else topk_from_beta(bhat)
|
||||
Eh = model(ch["yI"], ch["yQ"], topk_mask=tk)
|
||||
# BS-side estimate from decoded embeddings only
|
||||
Cb = torch.einsum("nud,nvd->uv", Eh, Eh) / Eh.shape[0]
|
||||
Cb.fill_diagonal_(0.0)
|
||||
bhat = gamma_ewma * bhat + (1 - gamma_ewma) * Cb
|
||||
Ehf = model(ch["yI"], ch["yQ"])
|
||||
Eho = model(ch["yI"], ch["yQ"], topk_mask=oracle)
|
||||
traj["online_cos"].append(mean_cos(Eh, E))
|
||||
traj["full_cos"].append(mean_cos(Ehf, E))
|
||||
traj["oracle_cos"].append(mean_cos(Eho, E))
|
||||
traj["beta_err"].append(float((bhat - Btrue).abs().mean()))
|
||||
print(f"[E9] frame {t}: online={traj['online_cos'][-1]:.4f} "
|
||||
f"oracle={traj['oracle_cos'][-1]:.4f} "
|
||||
f"full={traj['full_cos'][-1]:.4f}", flush=True)
|
||||
|
||||
# ---- wall-clock overhead incl. estimation + argpartition selection
|
||||
timing = {"U": [8, 16, 32, 64, 128], "full_ms": [], "topk_ms": [],
|
||||
"select_ms": []}
|
||||
for Ut in timing["U"]:
|
||||
dt = max(d, 2 * Ut) # keep at least 2 dims per user slot
|
||||
mt = UWCA(dt, Ut, H).to(DEVICE).eval()
|
||||
ch = {"yI": torch.randn(256, dt, device=DEVICE),
|
||||
"yQ": torch.randn(256, dt, device=DEVICE)}
|
||||
bh = torch.rand(Ut, Ut, device=DEVICE)
|
||||
with torch.no_grad():
|
||||
for _ in range(3):
|
||||
mt(ch["yI"], ch["yQ"]) # warm-up
|
||||
torch.cuda.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
for _ in range(20):
|
||||
mt(ch["yI"], ch["yQ"])
|
||||
torch.cuda.synchronize()
|
||||
t_full = (time.perf_counter() - t0) / 20 * 1e3
|
||||
t0 = time.perf_counter()
|
||||
for _ in range(20):
|
||||
b = bh.clone(); b.fill_diagonal_(2.0)
|
||||
idx = torch.topk(b, k, dim=1).indices
|
||||
tkm = torch.zeros(Ut, Ut, dtype=torch.bool, device=DEVICE)
|
||||
tkm.scatter_(1, idx, True)
|
||||
mt(ch["yI"], ch["yQ"], topk_mask=tkm)
|
||||
torch.cuda.synchronize()
|
||||
t_topk = (time.perf_counter() - t0) / 20 * 1e3
|
||||
t0 = time.perf_counter()
|
||||
for _ in range(100):
|
||||
b = bh.clone(); b.fill_diagonal_(2.0)
|
||||
idx = torch.topk(b, k, dim=1).indices
|
||||
torch.cuda.synchronize()
|
||||
t_sel = (time.perf_counter() - t0) / 100 * 1e3
|
||||
timing["full_ms"].append(t_full)
|
||||
timing["topk_ms"].append(t_topk)
|
||||
timing["select_ms"].append(t_sel)
|
||||
print(f"[E9] U={Ut}: full={t_full:.2f}ms topk={t_topk:.2f}ms "
|
||||
f"select={t_sel:.3f}ms", flush=True)
|
||||
|
||||
save_json("e9_topk_online.json", {"trajectory": traj, "timing": timing,
|
||||
"k": k, "warm_frames": warm})
|
||||
Reference in New Issue
Block a user