Initial release: code for WCL2026-1544 (context-aware embedding masking via DRL)

This commit is contained in:
Ki-Ho Lee
2026-06-22 17:31:32 +09:00
commit 8b7f70d650
26 changed files with 3445 additions and 0 deletions
+129
View File
@@ -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()