Initial release: code for WCL2026-1544 (context-aware embedding masking via DRL)
This commit is contained in:
@@ -0,0 +1,129 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user