File size: 2,305 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
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Create jsonl shards for unfinished e2e Omni inference."""

from __future__ import annotations

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


BASE_DIR = Path("/workspace/echoloc/Dataset/Novel/query_data/v2_2000")
DEFAULT_INPUT = BASE_DIR / "eval_inputs/e2e_input_aligned.jsonl"
DEFAULT_OUT_DIR = BASE_DIR / "eval_inputs/e2e_remaining_shards"
DEFAULT_WAV_DIR = BASE_DIR / "model_outputs/qwen3omni_e2e_from_query_audio/wavs"


def iter_jsonl(path: Path) -> Iterable[Dict[str, Any]]:
    with path.open("r", encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if line:
                yield json.loads(line)


def write_jsonl(path: Path, rows: List[Dict[str, Any]]) -> None:
    with path.open("w", encoding="utf-8") as f:
        for row in rows:
            f.write(json.dumps(row, ensure_ascii=False) + "\n")


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--input_jsonl", type=Path, default=DEFAULT_INPUT)
    ap.add_argument("--out_dir", type=Path, default=DEFAULT_OUT_DIR)
    ap.add_argument("--wav_dir", type=Path, default=DEFAULT_WAV_DIR)
    ap.add_argument("--num_shards", type=int, default=4)
    args = ap.parse_args()

    args.out_dir.mkdir(parents=True, exist_ok=True)
    rows = list(iter_jsonl(args.input_jsonl))
    remaining = [
        row for row in rows
        if not (args.wav_dir / f"{row['qid']}.wav").exists()
    ]

    shards: List[List[Dict[str, Any]]] = [[] for _ in range(args.num_shards)]
    for i, row in enumerate(remaining):
        shards[i % args.num_shards].append(row)

    for i, shard in enumerate(shards):
        write_jsonl(args.out_dir / f"e2e_remaining_shard_{i:02d}.jsonl", shard)

    summary = {
        "input_total": len(rows),
        "done_wavs": len(rows) - len(remaining),
        "remaining": len(remaining),
        "num_shards": args.num_shards,
        "shard_sizes": [len(x) for x in shards],
        "out_dir": str(args.out_dir),
    }
    with (args.out_dir / "summary.json").open("w", encoding="utf-8") as f:
        json.dump(summary, f, ensure_ascii=False, indent=2)
    print(json.dumps(summary, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()