| 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() |
|
|