"""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)