#!/usr/bin/env python3 """ analyze_drl_improvements.py Reads results_improve//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()