Initial release: code for WCL2026-1544 (context-aware embedding masking via DRL)
This commit is contained in:
Executable
+104
@@ -0,0 +1,104 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
analyze_drl_improvements.py
|
||||
Reads results_improve/<tag>/drl_snr_sweep.csv for each experiment
|
||||
variant and compares against the existing MAML/Joint/baseline runs.
|
||||
Prints a per-SNR and averaged CosSim table and a "gap to MAML" column.
|
||||
"""
|
||||
import csv, os, sys
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parent.parent
|
||||
IMP_DIR = ROOT / "results_improve"
|
||||
SWEEP = ROOT / "results_sweeps"
|
||||
RES = ROOT / "results_drl"
|
||||
LONG = ROOT / "results_drl_long"
|
||||
|
||||
TRAIN_SNRS = {0, 5, 10, 15, 20, 25}
|
||||
|
||||
|
||||
def load(path):
|
||||
if not os.path.exists(path):
|
||||
return None
|
||||
with open(path) as f:
|
||||
rows = list(csv.DictReader(f))
|
||||
return rows if rows else None
|
||||
|
||||
|
||||
def stats(rows):
|
||||
per = {int(float(r["snr_db"])): float(r["cos_sim"]) for r in rows}
|
||||
snrs = sorted(per.keys())
|
||||
vals = [per[s] for s in snrs]
|
||||
tr_vals = [per[s] for s in snrs if s in TRAIN_SNRS]
|
||||
orth = [float(r["orthogonality"]) for r in rows]
|
||||
return {
|
||||
"per": per,
|
||||
"avg_all": sum(vals) / len(vals),
|
||||
"avg_train": sum(tr_vals) / len(tr_vals) if tr_vals else None,
|
||||
"orth_mean": sum(orth) / len(orth),
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
# Load baselines
|
||||
bench = {}
|
||||
for name, path in [
|
||||
("Joint 100ep", RES / "joint_snr_sweep.csv"),
|
||||
("MAML 200ep", SWEEP / "maml_U4_200ep" / "maml_snr_sweep.csv"),
|
||||
("DRL 200ep (paper, results_drl_long)",
|
||||
LONG / "drl_snr_sweep.csv"),
|
||||
]:
|
||||
rows = load(path)
|
||||
if rows is not None:
|
||||
bench[name] = stats(rows)
|
||||
|
||||
# Load improvement variants
|
||||
improv = {}
|
||||
if IMP_DIR.exists():
|
||||
for d in sorted(IMP_DIR.iterdir()):
|
||||
path = d / "drl_snr_sweep.csv"
|
||||
rows = load(path)
|
||||
if rows is not None:
|
||||
improv[f"Improvement: {d.name}"] = stats(rows)
|
||||
|
||||
all_rows = {**bench, **improv}
|
||||
|
||||
if not all_rows:
|
||||
print("No results found. Run run_drl_improvements.sh first.")
|
||||
return
|
||||
|
||||
# Print header
|
||||
snrs = sorted(next(iter(all_rows.values()))["per"].keys())
|
||||
header = f"{'method':42s} | " + " | ".join(
|
||||
f"{s:>5}dB" for s in snrs) + " | avg_all | avg_tr | orth"
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
|
||||
maml_ref = bench.get("MAML 200ep", {}).get("avg_all")
|
||||
|
||||
for name, st in all_rows.items():
|
||||
per = st["per"]
|
||||
row_str = f"{name:42s} | " + " | ".join(
|
||||
f"{per.get(s, float('nan')):7.4f}" for s in snrs)
|
||||
gap = f" (vs MAML {st['avg_all']-maml_ref:+.4f})" \
|
||||
if maml_ref is not None else ""
|
||||
row_str += f" | {st['avg_all']:7.4f} | "
|
||||
row_str += f"{st['avg_train']:6.4f}" if st["avg_train"] is not None else " N/A"
|
||||
row_str += f" | {st['orth_mean']:.4f}{gap}"
|
||||
print(row_str)
|
||||
|
||||
print()
|
||||
if maml_ref is not None:
|
||||
print(f"MAML 200ep avg_all CosSim = {maml_ref:.4f}")
|
||||
joint_ref = bench.get("Joint 100ep", {}).get("avg_all")
|
||||
if joint_ref is not None:
|
||||
mm_gap = maml_ref - joint_ref
|
||||
print(f"MAML - Joint gap = {mm_gap:+.4f}")
|
||||
for name in [k for k in all_rows if "DRL" in k or "Improvement" in k]:
|
||||
dg = all_rows[name]["avg_all"] - joint_ref
|
||||
rec = dg / mm_gap * 100 if mm_gap != 0 else float("nan")
|
||||
print(f" {name:42s} -> recovered {rec:5.1f}% of MAML-over-Joint gap")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user