94 lines
3.4 KiB
Python
Executable File
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)
|