Initial release: code for WCL2026-1544 (context-aware embedding masking via DRL)
This commit is contained in:
@@ -0,0 +1,70 @@
|
||||
#!/usr/bin/env python3
|
||||
# ---------------------------------------------------------------
|
||||
# Fixed-orthogonal-mask scheme across user counts U = 1..6, for
|
||||
# overlaying the independent reference on Fig. 2 and Fig. 3 of the
|
||||
# WCL revision.
|
||||
#
|
||||
# Reuses the EXACT training/eval path of eval_task_oriented.py
|
||||
# (fixed QR masks, transceiver trained with MSE+CosSim, seed 0,
|
||||
# matched 100-epoch x 200-step budget) so the U=4 numbers reproduce
|
||||
# Table II. Evaluates the full SNR sweep per U so the same file
|
||||
# serves Fig 3(a) (CosSim vs SNR, U=4), Fig 2(b) (U=4, 20 dB) and
|
||||
# Fig 3(b) (throughput vs U at 10 dB).
|
||||
#
|
||||
# Output: ../results_sweeps/task_oriented/fixed_orth_byU.csv
|
||||
# ---------------------------------------------------------------
|
||||
import os, csv, argparse
|
||||
from types import SimpleNamespace
|
||||
import torch
|
||||
|
||||
from eval_task_oriented import (load_emb, train_transceiver, sweep_metrics,
|
||||
build_fixed_orthogonal_masks)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--users-list", type=int, nargs="+",
|
||||
default=[1, 2, 3, 4, 5, 6])
|
||||
ap.add_argument("--epochs", type=int, default=100)
|
||||
ap.add_argument("--steps-per-epoch", type=int, default=200)
|
||||
ap.add_argument("--eval-trials", type=int, default=200)
|
||||
ap.add_argument("--seed", type=int, default=0)
|
||||
ap.add_argument("--embed-file", type=str, default="../bert_agnews_8000.pt")
|
||||
ap.add_argument("--out", type=str,
|
||||
default="../results_sweeps/task_oriented/fixed_orth_byU.csv")
|
||||
a = ap.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}")
|
||||
d_bert, mux = 768, 4
|
||||
d_s = d_bert * mux
|
||||
emb = load_emb(a.embed_file, d_bert, 8000, device)
|
||||
|
||||
rows_all = []
|
||||
for U in a.users_list:
|
||||
print(f"\n=== Fixed-Orth scheme, U={U} ===", flush=True)
|
||||
args = SimpleNamespace(
|
||||
users=U, users_max=8, mux_factor=mux, d_bert=d_bert,
|
||||
hidden=256, rank=64, channel="rayleigh",
|
||||
train_snr=[0, 5, 10, 15, 20, 25], lr=1e-3, ce_tau=16.0,
|
||||
epochs=a.epochs, steps_per_epoch=a.steps_per_epoch,
|
||||
eval_trials=a.eval_trials, seed=a.seed)
|
||||
Mfix = build_fixed_orthogonal_masks(U, d_s, seed=0)
|
||||
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:
|
||||
rows_all.append((U, *r))
|
||||
|
||||
os.makedirs(os.path.dirname(a.out), exist_ok=True)
|
||||
with open(a.out, "w", newline="") as f:
|
||||
w = csv.writer(f)
|
||||
w.writerow(["users", "snr_db", "cos_sim", "orthogonality", "top1_acc"])
|
||||
for r in rows_all:
|
||||
w.writerow(r)
|
||||
print(f"\n[DONE] wrote {a.out}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user