File size: 3,114 Bytes
d8bfe4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Sample aligned reference-inference inputs by qid."""

from __future__ import annotations

import argparse
import json
import random
from pathlib import Path
from typing import Any, Dict, Iterable, List, Set


def read_jsonl(path: Path) -> List[Dict[str, Any]]:
    rows: List[Dict[str, Any]] = []
    if not path.exists():
        return rows
    with path.open("r", encoding="utf-8") as handle:
        for line in handle:
            line = line.strip()
            if line:
                rows.append(json.loads(line))
    return rows


def write_jsonl(path: Path, rows: Iterable[Dict[str, Any]]) -> int:
    path.parent.mkdir(parents=True, exist_ok=True)
    n = 0
    with path.open("w", encoding="utf-8") as handle:
        for row in rows:
            handle.write(json.dumps(row, ensure_ascii=False) + "\n")
            n += 1
    return n


def row_id(row: Dict[str, Any]) -> str:
    return str(row.get("qid") or row.get("id") or "")


def filter_by_qids(rows: Iterable[Dict[str, Any]], qids: Set[str]) -> List[Dict[str, Any]]:
    return [row for row in rows if row_id(row) in qids]


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--e2e_input", type=Path, required=True)
    parser.add_argument("--talker_input", type=Path, required=True)
    parser.add_argument("--response_input", type=Path, required=True)
    parser.add_argument("--out_dir", type=Path, required=True)
    parser.add_argument("--sample_size", type=int, default=50)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    e2e_rows = read_jsonl(args.e2e_input)
    if args.sample_size > 0 and len(e2e_rows) > args.sample_size:
        e2e_rows = random.Random(args.seed).sample(e2e_rows, args.sample_size)
    e2e_rows = sorted(e2e_rows, key=row_id)
    qids = {row_id(row) for row in e2e_rows if row_id(row)}

    talker_rows = sorted(filter_by_qids(read_jsonl(args.talker_input), qids), key=row_id)
    response_rows = sorted(filter_by_qids(read_jsonl(args.response_input), qids), key=row_id)

    written = {
        "e2e_input_aligned": write_jsonl(args.out_dir / "e2e_input_aligned.jsonl", e2e_rows),
        "talker_input_all": write_jsonl(args.out_dir / "talker_input_all.jsonl", talker_rows),
        "response_tts_input_aligned": write_jsonl(args.out_dir / "response_tts_input_aligned.jsonl", response_rows),
    }
    summary = {
        "source": {
            "e2e_input": str(args.e2e_input),
            "talker_input": str(args.talker_input),
            "response_input": str(args.response_input),
        },
        "out_dir": str(args.out_dir),
        "seed": args.seed,
        "sample_size": args.sample_size,
        "selected_qids": sorted(qids),
        "written": written,
    }
    args.out_dir.mkdir(parents=True, exist_ok=True)
    (args.out_dir / "sample_summary.json").write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
    print(json.dumps(summary, ensure_ascii=False, indent=2))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())