File size: 7,898 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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Build Qwen3-Omni thinker SFT jsonl from aligned Novel query data."""

from __future__ import annotations

import argparse
import collections
import datetime as dt
import json
import random
import re
import shutil
from pathlib import Path
from typing import Any, Dict, Iterable, List, Tuple

from prompt_loader import render_prompt


DEFAULT_INPUT = Path("/workspace/echoloc/Dataset/Novel/query_data/v2_2000/eval_inputs/e2e_input_aligned.jsonl")
DEFAULT_OUT_DIR = Path("/workspace/echoloc/Dataset/Novel/query_data/v2_2000/sft/thinker_emochange")
EMO_TOKEN = "<|EMO_CHANGE|>"


def normalize_emo_change(text: str, language: str) -> str:
    text = (text or "").replace("<EMO_CHANGE>", EMO_TOKEN)
    if language == "en":
        text = re.sub(r"\s*<\|EMO_CHANGE\|>\s*", f" {EMO_TOKEN} ", text)
        text = re.sub(r" {2,}", " ", text)
        return text.strip()
    return re.sub(r"\s*<\|EMO_CHANGE\|>\s*", EMO_TOKEN, text).strip()


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


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


def build_sample(row: Dict[str, Any]) -> Dict[str, Any]:
    language = row.get("language", "zh") or "zh"
    query_type = row.get("query_type", "unknown") or "unknown"
    style = (row.get("oracle_thinker_style") or "").strip()
    text = normalize_emo_change(row.get("oracle_thinker_text") or "", language)
    answer = render_prompt("thinker_sft.answer_template", "Style: __STYLE__\n\nText: __TEXT__", {"STYLE": style, "TEXT": text})
    return {
        "id": row["id"],
        "task": "novel_query_emochange",
        "audio_url": row["query_audio_path"],
        "language": language,
        "ability": f"novel/{query_type}",
        "query_type": query_type,
        "answer": answer,
        "thinking": "",
        "qid": row.get("qid", row.get("id")),
        "source_query": row.get("query_text", ""),
        "query_audio_control_path": row.get("query_audio_control_path", ""),
        "oracle_response_text": row.get("oracle_response_text", ""),
        "source_event": row.get("source_event", ""),
        "emo_change_spacing_rule": "en:space_around; zh:no_space_around",
    }


def split_by_query_type(samples: List[Dict[str, Any]], val_ratio: float, seed: int) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
    by_type: Dict[str, List[Dict[str, Any]]] = collections.defaultdict(list)
    for sample in samples:
        by_type[sample.get("query_type") or "unknown"].append(sample)

    rng = random.Random(seed)
    train: List[Dict[str, Any]] = []
    val: List[Dict[str, Any]] = []
    for key in sorted(by_type):
        group = by_type[key]
        rng.shuffle(group)
        n_val = max(1, round(len(group) * val_ratio))
        val.extend(group[:n_val])
        train.extend(group[n_val:])
    rng.shuffle(train)
    rng.shuffle(val)
    return train, val


def spacing_stats(rows: Iterable[Dict[str, Any]]) -> Dict[str, int]:
    counts: collections.Counter[Tuple[str, bool, bool]] = collections.Counter()
    for row in rows:
        answer = row.get("answer", "")
        start = 0
        while True:
            idx = answer.find(EMO_TOKEN, start)
            if idx < 0:
                break
            left = answer[idx - 1] if idx > 0 else ""
            right_pos = idx + len(EMO_TOKEN)
            right = answer[right_pos] if right_pos < len(answer) else ""
            counts[(row.get("language", ""), left == " ", right == " ")] += 1
            start = right_pos
    return {str(key): value for key, value in sorted(counts.items())}


def validate_spacing(rows: Iterable[Dict[str, Any]]) -> List[str]:
    errors: List[str] = []
    for row in rows:
        answer = row.get("answer", "")
        start = 0
        while True:
            idx = answer.find(EMO_TOKEN, start)
            if idx < 0:
                break
            left = answer[idx - 1] if idx > 0 else ""
            right_pos = idx + len(EMO_TOKEN)
            right = answer[right_pos] if right_pos < len(answer) else ""
            language = row.get("language", "")
            if language == "zh" and (left == " " or right == " "):
                errors.append(f"{row.get('id')}: zh token has surrounding space")
            if language == "en" and not (left == " " and right == " "):
                errors.append(f"{row.get('id')}: en token lacks surrounding spaces")
            start = right_pos
    return errors


def backup_existing(out_dir: Path) -> str:
    existing = [out_dir / name for name in ("train.jsonl", "val.jsonl", "manifest.json") if (out_dir / name).exists()]
    if not existing:
        return ""
    stamp = dt.datetime.now().strftime("%Y%m%d_%H%M%S")
    backup_dir = out_dir / f"backup_before_sft_rebuild_{stamp}"
    backup_dir.mkdir(parents=True, exist_ok=True)
    for path in existing:
        shutil.copy2(path, backup_dir / path.name)
    return str(backup_dir)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", type=Path, default=DEFAULT_INPUT)
    parser.add_argument("--out_dir", type=Path, default=DEFAULT_OUT_DIR)
    parser.add_argument("--val_ratio", type=float, default=0.05)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--no_backup", action="store_true")
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    samples = [build_sample(row) for row in iter_jsonl(args.input)]
    train, val = split_by_query_type(samples, args.val_ratio, args.seed)

    errors = validate_spacing(train) + validate_spacing(val)
    if errors:
        for error in errors[:20]:
            print(f"[ERROR] {error}")
        raise SystemExit(f"invalid EMO_CHANGE spacing: {len(errors)} errors")

    args.out_dir.mkdir(parents=True, exist_ok=True)
    (args.out_dir / "logs").mkdir(exist_ok=True)
    (args.out_dir / "output").mkdir(exist_ok=True)
    backup_dir = "" if args.no_backup else backup_existing(args.out_dir)

    write_jsonl(args.out_dir / "train.jsonl", train)
    write_jsonl(args.out_dir / "val.jsonl", val)

    summary = {
        "created_at": dt.datetime.now().isoformat(timespec="seconds"),
        "source": str(args.input),
        "output_dir": str(args.out_dir),
        "backup_before_rebuild": backup_dir,
        "split_seed": args.seed,
        "split_method": "stratified_by_query_type",
        "val_ratio": args.val_ratio,
        "emo_change_spacing_rule": "en: add one space before and after token; zh: remove spaces around token",
        "total": len(samples),
        "train": len(train),
        "val": len(val),
        "query_type_total": dict(collections.Counter(s.get("query_type") for s in samples)),
        "query_type_train": dict(collections.Counter(s.get("query_type") for s in train)),
        "query_type_val": dict(collections.Counter(s.get("query_type") for s in val)),
        "emo_change_total": sum(EMO_TOKEN in s["answer"] for s in samples),
        "emo_change_train": sum(EMO_TOKEN in s["answer"] for s in train),
        "emo_change_val": sum(EMO_TOKEN in s["answer"] for s in val),
        "spacing_train": spacing_stats(train),
        "spacing_val": spacing_stats(val),
    }
    (args.out_dir / "manifest.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())