Files
uwca-semantic-mac/experiments/e7_meta.py
T

94 lines
3.4 KiB
Python
Executable File

"""E7 — Meta-training over a multi-dimensional task family: held-out
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="E7-meta",
log_state=eta_log)
torch.save(m_meta.state_dict(), lib.DATA / "e7_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"E7-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"[E7] {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"[E7] eta final={etas[-1]:.3f} max={max(etas):.3f} "
f"gnorm max={max(gns):.3f}", flush=True)
save_json("e7_meta.json", out)