EchoLoc / inference /benchmarks /export_benchmark_outputs.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
8.65 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Export thinker+talker chain outputs into benchmark-specific layouts."""
from __future__ import annotations
import argparse
import json
import os
import shutil
from pathlib import Path
from typing import Any
def read_json(path: Path) -> Any:
with path.open("r", encoding="utf-8") as f:
return json.load(f)
def write_json(path: Path, obj: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
json.dump(obj, f, ensure_ascii=False, indent=2)
def read_jsonl(path: Path) -> list[dict[str, Any]]:
rows = []
with path.open("r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
def id_map(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
return {str(r["id"]): r for r in rows if "id" in r}
def manifest_map(path: Path) -> dict[str, dict[str, Any]]:
if not path.is_file():
return {}
return id_map(read_jsonl(path))
def link_or_copy(src: str | Path, dst: str | Path) -> None:
src_p = Path(src)
dst_p = Path(dst)
dst_p.parent.mkdir(parents=True, exist_ok=True)
if dst_p.exists() or dst_p.is_symlink():
dst_p.unlink()
try:
os.symlink(src_p, dst_p)
except OSError:
shutil.copy2(src_p, dst_p)
def chain_records(chain_input: Path, thinker_output: Path, talker_manifest: Path) -> list[dict[str, Any]]:
inputs = id_map(read_jsonl(chain_input))
thinkers = id_map(read_jsonl(thinker_output))
manifests = manifest_map(talker_manifest)
ids = list(inputs.keys())
out = []
for sid in ids:
inp = inputs.get(sid, {})
thk = thinkers.get(sid, {})
man = manifests.get(sid, {})
out.append(
{
**inp,
"thinker_text": thk.get("thinker_text", ""),
"thinker_style": thk.get("thinker_style", ""),
"struct_pass": thk.get("struct_pass", False),
"wav": man.get("wav", ""),
"talker_text": man.get("text", thk.get("thinker_text", "")),
}
)
return out
def export_uro(args: argparse.Namespace) -> None:
rows = chain_records(Path(args.chain_input), Path(args.thinker_output), Path(args.talker_manifest))
out_dir = Path(args.out_dir)
audio_dir = out_dir / "audio"
pred_rows = []
question_rows = []
gt_rows = []
for idx, row in enumerate(rows):
sid = f"{idx:04d}"
pred = row.get("thinker_text") or row.get("talker_text") or ""
pred_rows.append({sid: pred})
question_rows.append({sid: row.get("source_text", "")})
gt_rows.append({sid: row.get("target_text", row.get("source_text", ""))})
wav = row.get("wav")
if wav and Path(wav).is_file():
link_or_copy(wav, audio_dir / f"{sid}.wav")
write_jsonl(out_dir / "pred_text.jsonl", pred_rows)
write_jsonl(out_dir / "question_text.jsonl", question_rows)
write_jsonl(out_dir / "gt_text.jsonl", gt_rows)
write_jsonl(out_dir / "chain_records.jsonl", rows)
print(f"[write] URO output dir={out_dir} rows={len(rows)}")
def export_generic(args: argparse.Namespace) -> None:
rows = chain_records(Path(args.chain_input), Path(args.thinker_output), Path(args.talker_manifest))
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
write_jsonl(out_dir / "chain_records.jsonl", rows)
print(f"[write] generic output dir={out_dir} rows={len(rows)}")
def export_echomind_asr(args: argparse.Namespace) -> None:
rows = chain_records(Path(args.chain_input), Path(args.thinker_output), Path(args.talker_manifest))
by_case = {str(r["case_id"]): r for r in rows}
root = Path(args.root_dir)
input_dir = root / "dataset" / f"data_{args.data_type}"
data = read_json(input_dir / f"script_info_{args.data_type}.json")
for d in data:
r = by_case.get(str(d.get("case_id")))
if r:
d["model"] = args.model_name
d["predicted_answer"] = r.get("thinker_text", "")
out = root / "output" / f"output_{args.data_type}" / "asr_output" / args.model_name / f"{args.model_name}_asr.json"
write_json(out, data)
print(f"[write] {out} rows={len(data)}")
def export_echomind_mcq(args: argparse.Namespace) -> None:
rows = chain_records(Path(args.chain_input), Path(args.thinker_output), Path(args.talker_manifest))
by_qid = {str(r["question_id"]): r for r in rows}
root = Path(args.root_dir)
input_dir = root / "dataset" / f"data_{args.data_type}"
data = read_json(input_dir / "MCQ" / args.mcq_file)
for d in data:
r = by_qid.get(str(d.get("question_id")))
if r:
d["predicted_audio_answer"] = r.get("thinker_text", "")
d["predicted_text_answer"] = r.get("thinker_text", "")
d["prediected_audio_file"] = r.get("wav", "")
stem = Path(args.mcq_file).stem
out = (
root
/ "output"
/ f"output_{args.data_type}"
/ "mcq_output"
/ args.model_name
/ "audio_output_true"
/ f"{args.model_name}_mcq_{stem}.json"
)
write_json(out, data)
print(f"[write] {out} rows={len(data)}")
def export_echomind_response(args: argparse.Namespace) -> None:
rows = chain_records(Path(args.chain_input), Path(args.thinker_output), Path(args.talker_manifest))
by_pair = {(str(r.get("case_id")), str(r.get("voice_type"))): r for r in rows}
root = Path(args.root_dir)
input_dir = root / "dataset" / f"data_{args.data_type}"
data = read_json(input_dir / f"script_info_{args.data_type}.json")
for d in data:
d["model"] = args.model_name
d["ouput_prompt"] = args.system_prompt
d.setdefault("output_content", {"target": {}, "neutral": {}, "alternative": {}})
for voice_type in ("target", "neutral", "alternative"):
r = by_pair.get((str(d.get("case_id")), voice_type))
if not r:
continue
d["output_content"].setdefault(voice_type, {})
d["output_content"][voice_type]["file"] = Path(str(r.get("wav", ""))).name
d["output_content"][voice_type]["response_audio_transcript"] = r.get("thinker_text", "")
d["output_content"][voice_type]["response_text"] = r.get("thinker_text", "")
out = (
root
/ "output"
/ f"output_{args.data_type}"
/ "response_output"
/ args.model_name
/ args.system_prompt
/ f"{args.model_name}_{args.system_prompt}_output_responses.json"
)
write_json(out, data)
print(f"[write] {out} rows={len(data)}")
def main() -> None:
ap = argparse.ArgumentParser()
sub = ap.add_subparsers(dest="cmd", required=True)
def add_chain_args(p: argparse.ArgumentParser) -> None:
p.add_argument("--chain-input", required=True)
p.add_argument("--thinker-output", required=True)
p.add_argument("--talker-manifest", required=True)
p = sub.add_parser("uro")
add_chain_args(p)
p.add_argument("--out-dir", required=True)
p.set_defaults(func=export_uro)
p = sub.add_parser("generic")
add_chain_args(p)
p.add_argument("--out-dir", required=True)
p.set_defaults(func=export_generic)
p = sub.add_parser("echomind-asr")
add_chain_args(p)
p.add_argument("--root-dir", required=True)
p.add_argument("--data-type", default="synthesis")
p.add_argument("--model-name", default="echoloc")
p.set_defaults(func=export_echomind_asr)
p = sub.add_parser("echomind-mcq")
add_chain_args(p)
p.add_argument("--root-dir", required=True)
p.add_argument("--data-type", default="synthesis")
p.add_argument("--model-name", default="echoloc")
p.add_argument("--mcq-file", required=True)
p.set_defaults(func=export_echomind_mcq)
p = sub.add_parser("echomind-response")
add_chain_args(p)
p.add_argument("--root-dir", required=True)
p.add_argument("--data-type", default="synthesis")
p.add_argument("--model-name", default="echoloc")
p.add_argument("--system-prompt", default="enhance")
p.set_defaults(func=export_echomind_response)
args = ap.parse_args()
args.func(args)
if __name__ == "__main__":
main()