#!/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()