Files
WCL/analyze_drl_improvements.py

105 lines
3.3 KiB
Python
Executable File

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