EchoLoc / data_construction /build_thinker_sft_data.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
7.9 kB
#!/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())