EchoLoc / experiments /reviewer_analyses /interface_ablation.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
2.21 kB
from __future__ import annotations
import argparse
from pathlib import Path
from common import group_by, markdown_table, read_csv, safe_mean, safe_std, write_csv
VSTYLE_DIMS = ["overall", "ESMOS", "IEMOS", "TTMOS", "PVMOS", "NCMOS"]
RENDER_DIMS = ["autopcp_ref", "autopcp_ref_var", "emo2vec_ref", "audonnx_ref", "emo2vec_ref_unsmooth", "audonnx_ref_unsmooth", "wer"]
THINKER_DIMS = ["emotion_f1", "intensity_acc", "transition_f1", "vad_bin_f1", "global_style_f1"]
def main() -> None:
parser = argparse.ArgumentParser(description="Build the explicit interface ablation table.")
parser.add_argument("--vstyle", required=True, help="CSV with system, seed, and VStyle/MOS dimensions.")
parser.add_argument("--render", required=True, help="CSV with system, seed, and rendering metrics.")
parser.add_argument("--thinker", default="", help="Optional CSV with system, seed, and per-field Thinker metrics.")
parser.add_argument("--out-dir", required=True)
args = parser.parse_args()
vstyle = read_csv(Path(args.vstyle))
render = read_csv(Path(args.render))
thinker = read_csv(Path(args.thinker)) if args.thinker else []
systems = sorted(set(r["system"] for r in vstyle) | set(r["system"] for r in render) | set(r["system"] for r in thinker))
v_by_system = group_by(vstyle, "system")
r_by_system = group_by(render, "system")
t_by_system = group_by(thinker, "system")
rows = []
for system in systems:
row = {"system": system}
for dim in VSTYLE_DIMS:
row[dim] = _mean_pm(v_by_system.get(system, []), dim)
for dim in RENDER_DIMS:
row[dim] = _mean_pm(r_by_system.get(system, []), dim)
for dim in THINKER_DIMS:
row[dim] = _mean_pm(t_by_system.get(system, []), dim)
rows.append(row)
out_dir = Path(args.out_dir)
write_csv(out_dir / "tab_interface.csv", rows)
(out_dir / "tab_interface.md").write_text(markdown_table(rows), encoding="utf-8")
def _mean_pm(rows, key):
vals = [r.get(key, "") for r in rows if r.get(key, "") != ""]
if not vals:
return "--"
return f"{safe_mean(vals):.3f} +/- {safe_std(vals):.3f}"
if __name__ == "__main__":
main()