130 lines
5.9 KiB
Python
130 lines
5.9 KiB
Python
#!/usr/bin/env python3
|
|
# ---------------------------------------------------------------
|
|
# Multi-seed re-evaluation of the Table II retrieval metric.
|
|
# For each seed s in --seeds:
|
|
# - Proposed DRL : load checkpoint drl_U4_100ep_s{s}
|
|
# - Fixed-Orth : fixed orthogonal masks + transceiver (seed s)
|
|
# - Static (Sem) : free masks + transceiver, MSE+CosSim (seed s)
|
|
# - Static (CE) : free masks + transceiver, symbol-CE (seed s)
|
|
# Evaluation uses the SAME paired-channel protocol as
|
|
# eval_task_oriented.py (eval_seed=12345, 200 trials per SNR).
|
|
# Writes per-(method,seed,snr) rows and prints mean +/- std across
|
|
# seeds for each (method, snr).
|
|
# ---------------------------------------------------------------
|
|
import os, csv, argparse
|
|
from types import SimpleNamespace
|
|
import numpy as np
|
|
import torch
|
|
|
|
from eval_task_oriented import (load_emb, sweep_metrics, train_transceiver,
|
|
build_fixed_orthogonal_masks, load_drl)
|
|
|
|
SNRS = list(range(0, 31, 5))
|
|
|
|
|
|
def main():
|
|
p = argparse.ArgumentParser()
|
|
p.add_argument("--seeds", type=int, nargs="+",
|
|
default=[0, 7, 42, 123, 2025, 2026])
|
|
p.add_argument("--users", type=int, default=4)
|
|
p.add_argument("--users-max", type=int, default=8)
|
|
p.add_argument("--mux-factor", type=int, default=4)
|
|
p.add_argument("--d-bert", type=int, default=768)
|
|
p.add_argument("--hidden", type=int, default=256)
|
|
p.add_argument("--rank", type=int, default=64)
|
|
p.add_argument("--embed-file", type=str, default="../bert_agnews_8000.pt")
|
|
p.add_argument("--pool-size", type=int, default=8000)
|
|
p.add_argument("--channel", default="rayleigh")
|
|
p.add_argument("--train-snr", type=float, nargs="+",
|
|
default=[0, 5, 10, 15, 20, 25])
|
|
p.add_argument("--lr", type=float, default=1e-3)
|
|
p.add_argument("--ce-tau", type=float, default=16.0)
|
|
p.add_argument("--epochs", type=int, default=100)
|
|
p.add_argument("--steps-per-epoch", type=int, default=200)
|
|
p.add_argument("--eval-trials", type=int, default=200)
|
|
p.add_argument("--ckpt-tmpl", type=str,
|
|
default="../results_sweeps/drl_U4_100ep_s{seed}/drl_mask_policy.pt")
|
|
p.add_argument("--out-dir", type=str,
|
|
default="../results_sweeps/task_oriented")
|
|
a = p.parse_args()
|
|
|
|
device = torch.device("mps" if torch.backends.mps.is_available()
|
|
else ("cuda" if torch.cuda.is_available() else "cpu"))
|
|
print(f"[INFO] device={device} seeds={a.seeds}")
|
|
U, d_s = a.users, a.d_bert * a.mux_factor
|
|
emb = load_emb(a.embed_file, a.d_bert, a.pool_size, device)
|
|
|
|
# method -> snr -> list of acc across seeds
|
|
acc = {m: {s: [] for s in SNRS}
|
|
for m in ["proposed_drl", "fixed_orth", "static_sem", "static_ce"]}
|
|
raw_rows = []
|
|
|
|
for seed in a.seeds:
|
|
print(f"\n########## SEED {seed} ##########")
|
|
args = SimpleNamespace(users=U, users_max=a.users_max,
|
|
mux_factor=a.mux_factor, d_bert=a.d_bert,
|
|
hidden=a.hidden, rank=a.rank, lr=a.lr,
|
|
ce_tau=a.ce_tau, epochs=a.epochs,
|
|
steps_per_epoch=a.steps_per_epoch,
|
|
channel=a.channel, train_snr=a.train_snr,
|
|
seed=seed)
|
|
|
|
print(f"--- Proposed DRL (seed {seed}) ---")
|
|
ckpt = a.ckpt_tmpl.format(seed=seed)
|
|
trx, mfn = load_drl(ckpt, args, device)
|
|
rows = sweep_metrics(trx, emb, mfn, device, U, a.eval_trials)
|
|
for r in rows:
|
|
acc["proposed_drl"][r[0]].append(r[3]); raw_rows.append(["proposed_drl", seed, *r])
|
|
|
|
print(f"--- Fixed-Orth (seed {seed}) ---")
|
|
Mfix = build_fixed_orthogonal_masks(U, d_s, seed=seed)
|
|
trx, mfn = train_transceiver(emb, args, device, "fixed", fixed_masks=Mfix)
|
|
rows = sweep_metrics(trx, emb, mfn, device, U, a.eval_trials)
|
|
for r in rows:
|
|
acc["fixed_orth"][r[0]].append(r[3]); raw_rows.append(["fixed_orth", seed, *r])
|
|
|
|
print(f"--- Static Sem (seed {seed}) ---")
|
|
trx, mfn = train_transceiver(emb, args, device, "sem")
|
|
rows = sweep_metrics(trx, emb, mfn, device, U, a.eval_trials)
|
|
for r in rows:
|
|
acc["static_sem"][r[0]].append(r[3]); raw_rows.append(["static_sem", seed, *r])
|
|
|
|
print(f"--- Static CE (seed {seed}) ---")
|
|
trx, mfn = train_transceiver(emb, args, device, "ce")
|
|
rows = sweep_metrics(trx, emb, mfn, device, U, a.eval_trials)
|
|
for r in rows:
|
|
acc["static_ce"][r[0]].append(r[3]); raw_rows.append(["static_ce", seed, *r])
|
|
|
|
os.makedirs(a.out_dir, exist_ok=True)
|
|
raw_out = os.path.join(a.out_dir, "taskmetric_multiseed_raw.csv")
|
|
with open(raw_out, "w", newline="") as f:
|
|
w = csv.writer(f)
|
|
w.writerow(["method", "seed", "snr_db", "cos_sim", "orthogonality", "top1_acc"])
|
|
w.writerows(raw_rows)
|
|
|
|
agg_out = os.path.join(a.out_dir, "taskmetric_multiseed_agg.csv")
|
|
with open(agg_out, "w", newline="") as f:
|
|
w = csv.writer(f)
|
|
w.writerow(["method", "snr_db", "n_seeds", "acc_mean_pct", "acc_std_pct"])
|
|
for m in acc:
|
|
for s in SNRS:
|
|
vals = np.array(acc[m][s]) * 100.0
|
|
w.writerow([m, s, len(vals), f"{vals.mean():.2f}", f"{vals.std():.2f}"])
|
|
|
|
print(f"\n[DONE] raw -> {raw_out}\n agg -> {agg_out}")
|
|
print("\n===== mean +/- std (%, top-1 retrieval Acc) =====")
|
|
names = {"fixed_orth": "Fixed-Orth", "static_ce": "Static, CE",
|
|
"static_sem": "Static, Sem", "proposed_drl": "Proposed"}
|
|
hdr = "Method " + "".join(f"{s:>13}" for s in SNRS)
|
|
print(hdr)
|
|
for m in ["fixed_orth", "static_ce", "static_sem", "proposed_drl"]:
|
|
line = f"{names[m]:<12}"
|
|
for s in SNRS:
|
|
v = np.array(acc[m][s]) * 100.0
|
|
line += f"{v.mean():>6.1f}±{v.std():>4.1f}"
|
|
print(line)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|