105 lines
3.3 KiB
Python
Executable File
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()
|