Reproducibility package: UWCA semantic multiple access (TWC submission)
This commit is contained in:
Executable
+93
@@ -0,0 +1,93 @@
|
||||
"""E7 v2 — Meta-training over the multi-dimensional task family, OOD transfer,
|
||||
adaptation sweep, and eta/gradient logging.
|
||||
|
||||
meta : the paper's first-order meta-training aggregated over the FULL
|
||||
36-task family (SNR x {Rayleigh, Rician K=5,10 dB} x phase {0,10 deg})
|
||||
lookup : per-SNR specialists trained on Rayleigh / no phase error, indexed by
|
||||
nearest SNR (the 1-D lookup table)
|
||||
Test on held-out task combinations (incl. unseen Nakagami fading), zero-shot
|
||||
and with S in {1,5,10,20} inner adaptation steps from each initialization.
|
||||
"""
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import lib
|
||||
from lib import (SCENARIOS, UWCA, DEVICE, adapt, block_masks, eval_scheme,
|
||||
gen_embeddings, save_json, 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]
|
||||
fads = [{"fading": "rayleigh"},
|
||||
{"fading": "rician", "rician_K_dB": 5.0},
|
||||
{"fading": "rician", "rician_K_dB": 10.0}]
|
||||
phis = [0.0, 10.0]
|
||||
family = [{"snr_db": s, "phase_sigma_deg": p, **f}
|
||||
for s in snrs for f in fads for p in phis]
|
||||
|
||||
eta_log = []
|
||||
m_meta = UWCA(d, U, H).to(DEVICE)
|
||||
train_multitask(m_meta, gen, family, epochs=300, tag="E7v2-meta",
|
||||
log_state=eta_log)
|
||||
torch.save(m_meta.state_dict(), lib.DATA / "e7v2_meta.pt")
|
||||
|
||||
specialists = {}
|
||||
for s in snrs:
|
||||
m = UWCA(d, U, H).to(DEVICE)
|
||||
train_multitask(m, gen, [{"snr_db": s}], epochs=150,
|
||||
tag=f"E7v2-spec{int(s)}")
|
||||
specialists[s] = m
|
||||
|
||||
|
||||
def lookup(snr):
|
||||
return specialists[min(snrs, key=lambda x: abs(x - snr))]
|
||||
|
||||
|
||||
test_tasks = {
|
||||
"ricianK20_phi15_snr10": {"snr_db": 10.0, "fading": "rician",
|
||||
"rician_K_dB": 20.0, "phase_sigma_deg": 15.0},
|
||||
"nakagami3_phi5_snr10": {"snr_db": 10.0, "fading": "nakagami",
|
||||
"nakagami_m": 3.0, "phase_sigma_deg": 5.0},
|
||||
"rayleigh_phi20_snr6": {"snr_db": 6.0, "fading": "rayleigh",
|
||||
"phase_sigma_deg": 20.0},
|
||||
"ricianK20_phi15_snr18": {"snr_db": 18.0, "fading": "rician",
|
||||
"rician_K_dB": 20.0, "phase_sigma_deg": 15.0},
|
||||
"indist_rayleigh_snr10": {"snr_db": 10.0, "fading": "rayleigh",
|
||||
"phase_sigma_deg": 0.0},
|
||||
}
|
||||
|
||||
S_grid = [0, 1, 5, 10, 20]
|
||||
out = {"S_grid": S_grid, "results": {}}
|
||||
for name, t in test_tasks.items():
|
||||
row = {}
|
||||
for label, base in [("meta", m_meta), ("lookup", lookup(t["snr_db"]))]:
|
||||
sers = []
|
||||
for S in S_grid:
|
||||
mdl = base if S == 0 else adapt(base, gen, t, steps=S,
|
||||
inner_lr=0.02)
|
||||
s, c = eval_scheme("uwca", gen, t, n_mc=150, model=mdl,
|
||||
masks=masks)
|
||||
sers.append({"ser": s, "cos": c})
|
||||
row[label] = sers
|
||||
out["results"][name] = row
|
||||
print(f"[E7v2] {name}: meta={[round(x['ser'],3) for x in row['meta']]} "
|
||||
f"lookup={[round(x['ser'],3) for x in row['lookup']]}", flush=True)
|
||||
|
||||
etas = [e["eta"] for e in eta_log]
|
||||
gns = [e["gnorm"] for e in eta_log]
|
||||
out["eta_traj"] = etas[::5]
|
||||
out["gnorm_traj"] = gns[::5]
|
||||
out["eta_final"], out["eta_max"] = etas[-1], max(etas)
|
||||
out["gnorm_max"] = max(gns)
|
||||
print(f"[E7v2] eta final={etas[-1]:.3f} max={max(etas):.3f} "
|
||||
f"gnorm max={max(gns):.3f}", flush=True)
|
||||
|
||||
save_json("e7_v2_meta.json", out)
|
||||
Reference in New Issue
Block a user