EchoLoc / data_construction /generate_response_controls.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
22.5 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Generate response-side TTS controls from accepted query candidates.
This is the second semantic planning stage:
query candidate -> assistant/TTS response Global+Control.
The output intentionally keeps both `response_generated_control` and
`final_generated_control` so existing Qwen3TTS generation scripts can consume it
with their normal control key.
"""
from __future__ import annotations
import argparse
import json
import re
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Sequence, Set, Tuple
from common_llm import DEFAULT_BASE_URL, chat_completion, extract_json_object
from prompt_loader import get_prompt, render_prompt
DEFAULT_INPUT = "/workspace/echoloc/Dataset/Novel/query_data/v2_2000/filtered/query_candidates_keep.jsonl"
DEFAULT_OUTPUT = "/workspace/echoloc/Dataset/Novel/query_data/v2_2000/thinker_targets/response_controls.jsonl"
DEFAULT_FAILED_OUTPUT = "/workspace/echoloc/Dataset/Novel/query_data/v2_2000/thinker_targets/response_controls_failed.jsonl"
CONTROL_KEY = "response_generated_control"
ALLOWED_GENDERS = {"Male", "Female"}
ALLOWED_AGES = {"Child", "Teen", "Young_Adult", "Middle_Aged", "Senior"}
ALLOWED_LEVELS = {"Level_1_Subtle", "Level_2_Mild", "Level_3_Strong", "Level_4_Extreme"}
ACTION_OR_MARKUP_RE = re.compile(r"[\(\(\[\【].*?[\)\)\]\】]|旁白|动作|叹气|沉默|停顿|笑声:|哭声:")
INSTRUCTION_PREFIX_RE = re.compile(r"^\s*(请|帮我|麻烦|用|以).{0,40}(说|读|念|朗读|播报)[::,,]")
def iter_jsonl(path: Path) -> Iterable[Dict[str, Any]]:
with path.open("r", encoding="utf-8") as f:
for line in f:
if line.strip():
yield json.loads(line)
def load_existing_qids(path: Path) -> Set[str]:
qids: Set[str] = set()
if not path.exists():
return qids
with path.open("r", encoding="utf-8") as f:
for line in f:
try:
qid = json.loads(line).get("qid")
except json.JSONDecodeError:
continue
if qid:
qids.add(str(qid))
return qids
def append_jsonl(path: Path, rows: Sequence[Dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
def norm_text(text: Any) -> str:
return re.sub(r"\s+", "", str(text or "")).strip()
def clean_text(text: Any) -> str:
text = str(text or "")
text = re.sub(r"\s+", " ", text).strip()
text = re.sub(r"^[\"'“”‘’]+|[\"'“”‘’]+$", "", text).strip()
return text
def extract_quoted_payload(text: str) -> str:
"""Best-effort extraction of target utterance from an instruction query."""
if not text:
return ""
patterns = [
r"(?:说|读|念|朗读|播报)[::]\s*[‘'“\"](.+?)[’'”\"]\s*$",
r"[‘'“\"](.+?)[’'”\"]\s*$",
]
for pattern in patterns:
match = re.search(pattern, text, flags=re.S)
if match:
return clean_text(match.group(1))
return ""
def item_language(item: Dict[str, Any]) -> str:
language = str(item.get("language") or (item.get("source") or {}).get("language") or "").strip().lower()
if language:
return language
source_candidate = item.get("source_query_candidate") or {}
language = str(source_candidate.get("language") or ((source_candidate.get("source") or {}).get("language")) or "").strip().lower()
if language:
return language
qid = str(item.get("qid") or "")
if qid.startswith("vstyle_en_") or "_en_" in qid:
return "en"
return "zh"
def get_language_prompt(base_key: str, language: str, default: str) -> str:
if language and language != "zh":
prompt = get_prompt(f"{base_key}.{language}", "")
if prompt:
return prompt
return get_prompt(base_key, default)
def render_language_prompt(base_key: str, language: str, default: str, replacements: Dict[str, Any]) -> str:
if language and language != "zh":
rendered = render_prompt(f"{base_key}.{language}", "", replacements)
if rendered:
return rendered
return render_prompt(base_key, default, replacements)
def build_generation_messages(
item: Dict[str, Any],
previous_control: Optional[Dict[str, Any]] = None,
feedback: str = "",
) -> List[Dict[str, str]]:
language = item_language(item)
system = """
你是 Omni 情感语音训练数据的“回复规划 + TTS 声学控制”Agent。输入是一条已经质检通过的用户 query candidate。你要生成后续 talker 需要合成的高表现力回复语音控制,也就是 `response_generated_control`。
两类 query 的任务完全不同:
1. dialogue:
- 用户是在向 AI 倾诉、抱怨、请求回应或表达情绪。
- 你需要生成 AI 助手应该说出的回复台词,并给出该回复的 Global+Control。
- 回复要接住 visible_query.text 中可见/可听的情绪,必要时参考 target_contract。
- hidden_context/source_event 只能帮助理解抽象处境,不能把小说专名、未说出口的人物关系或具体剧情写进回复。
2. instruction:
- 用户是在要求 TTS 系统“用某种风格说某段话”。
- 你需要解析这条指令,生成真正应该被朗读/表演的目标文本,以及目标风格对应的 Global+Control。
- 绝对不要把整句用户指令当作 sample_text。通常 sample_text 应是引号中的目标台词。
Global/Control 规范:
- Global 只描述目标说话人的稳定音色/年龄/性别/人设/风格,必须包含明确年龄和性别。
- Control 描述每个台词片段的局部静态情绪、语速、音量、音高、气息、韵律。
- Control 可以 1-4 段。只有存在自然情绪、语义或强度变化时才分段。
- `sample_text` 必须是纯净 TTS 文本:不要 Markdown、不要括号动作、不要“AI:”、不要解释。
- VAD 坐标范围是 [0,1]:Valence 愉悦度,Arousal 唤醒度,Dominance 支配感。
- `instruct_en` 是对应中文声学描述的英文翻译。
输出只允许是合法 JSON,不要 Markdown:
{
"response_generated_control": {
"Global": {
"Gender": "Male | Female",
"Age": "Child | Teen | Young_Adult | Middle_Aged | Senior",
"Persona": "String",
"Style": "String",
"instruct_zh": "一句话全局音色/人设描述,必须包含明确年龄和性别",
"instruct_en": "English translation"
},
"Control": [
{
"emotion": "String",
"level": "Level_1_Subtle | Level_2_Mild | Level_3_Strong | Level_4_Extreme",
"instruct_zh": "该段静态声学控制",
"instruct_en": "English translation",
"vad_coordinates": [0.5, 0.5, 0.5],
"vad_explanation": "一句话说明 VAD 依据",
"sample_text": "该段要合成的纯净回复/目标台词"
}
]
},
"planning_reasoning": "简述如何从 query 推断回复目标、哪些隐藏信息没有使用、为何这样分段"
}
"""
system = get_language_prompt("response_controls.generation.system", language, system)
payload = {
"qid": item.get("qid"),
"query_type": item.get("query_type"),
"visible_query": item.get("visible_query"),
"query_voice": item.get("query_voice"),
"source_event": item.get("source_event"),
"hidden_context": item.get("hidden_context"),
"target_contract": item.get("target_contract"),
"quality_hints": item.get("quality_hints"),
}
user = get_language_prompt(
"response_controls.generation.user_prefix",
language,
"请为下面 query candidate 生成 response_generated_control:\n",
)
user += json.dumps(payload, ensure_ascii=False, indent=2)
if previous_control is not None or feedback:
user += render_language_prompt(
"response_controls.generation.retry",
language,
"\n\n上一次生成未通过,请根据反馈重写。\n上一次 control:\n__PREVIOUS_CONTROL_JSON__\n反馈:\n__FEEDBACK__",
{
"PREVIOUS_CONTROL_JSON": json.dumps(previous_control, ensure_ascii=False, indent=2) if previous_control is not None else "",
"FEEDBACK": feedback,
},
)
return [
{"role": "system", "content": system.strip()},
{"role": "user", "content": user},
]
def build_eval_messages(item: Dict[str, Any], control: Dict[str, Any], deterministic_issues: Sequence[str]) -> List[Dict[str, str]]:
language = item_language(item)
system = """
你是 query->response TTS control 的质检专家。请判断 generated control 是否适合合成后续 talker 的回复/目标语音。
硬失败项:
1. dialogue 类型没有生成 AI 助手回复,而是复读用户原话。
2. dialogue 类型引入小说专名、隐藏人物关系或用户未说出口的具体事实。
3. instruction 类型把整条用户指令当作 sample_text,而不是提取目标台词。
4. instruction 类型丢失用户明确要求的声学风格。
5. sample_text 包含括号动作、旁白、Markdown、解释性文字或角色前缀。
6. Global 缺少明确年龄/性别,或枚举字段/VAD 不合法。
评分:
- 5:完全可用,回复/目标文本自然,声学控制细致,分段合理。
- 4:可用,有轻微措辞瑕疵。
- 3:可修复,但存在明显偏差。
- 1/2:不可用。
只输出 JSON:
{"score": 5, "reason": "简短理由", "repair_suggestions": []}
"""
system = get_language_prompt("response_controls.eval.system", language, system)
payload = {
"qid": item.get("qid"),
"query_type": item.get("query_type"),
"visible_query": item.get("visible_query"),
"target_contract": item.get("target_contract"),
"deterministic_issues": list(deterministic_issues),
"generated_control": control,
}
return [
{"role": "system", "content": system.strip()},
{"role": "user", "content": json.dumps(payload, ensure_ascii=False, indent=2)},
]
def parse_control(raw: str) -> Tuple[Dict[str, Any], str]:
parsed = extract_json_object(raw)
reasoning = str(parsed.get("planning_reasoning", ""))
control = parsed.get(CONTROL_KEY) or parsed.get("final_generated_control") or parsed.get("generated_control") or parsed
if isinstance(control, dict) and CONTROL_KEY in control:
control = control[CONTROL_KEY]
return control, reasoning
def validate_control(item: Dict[str, Any], control: Dict[str, Any]) -> List[str]:
issues: List[str] = []
if not isinstance(control, dict):
return ["control_not_object"]
global_ctrl = control.get("Global")
segments = control.get("Control")
if not isinstance(global_ctrl, dict):
issues.append("missing_Global")
global_ctrl = {}
if not isinstance(segments, list) or not segments:
issues.append("missing_Control")
segments = []
gender = global_ctrl.get("Gender")
age = global_ctrl.get("Age")
global_zh = str(global_ctrl.get("instruct_zh", ""))
if gender not in ALLOWED_GENDERS:
issues.append("Gender_must_be_Male_or_Female")
if age not in ALLOWED_AGES:
issues.append("Age_enum_invalid")
if not re.search(r"(男|女|男性|女性|男孩|女孩|少年|少女|老人|老年|中年|年轻)", global_zh):
issues.append("Global_instruct_zh_missing_age_gender")
sample_texts: List[str] = []
for idx, seg in enumerate(segments):
if not isinstance(seg, dict):
issues.append(f"Control_{idx}_not_object")
continue
for key in ["emotion", "level", "instruct_zh", "instruct_en", "vad_coordinates", "vad_explanation", "sample_text"]:
if key not in seg:
issues.append(f"Control_{idx}_missing_{key}")
if seg.get("level") not in ALLOWED_LEVELS:
issues.append(f"Control_{idx}_level_invalid")
vad = seg.get("vad_coordinates")
if (
not isinstance(vad, list)
or len(vad) != 3
or any(not isinstance(x, (int, float)) or x < 0 or x > 1 for x in vad)
):
issues.append(f"Control_{idx}_vad_invalid")
sample_text = clean_text(seg.get("sample_text", ""))
if not sample_text:
issues.append(f"Control_{idx}_sample_text_empty")
if ACTION_OR_MARKUP_RE.search(sample_text):
issues.append(f"Control_{idx}_sample_text_contains_action_or_markup")
sample_texts.append(sample_text)
joined = "".join(sample_texts)
visible = str((item.get("visible_query") or {}).get("text", ""))
qtype = item.get("query_type")
if qtype == "dialogue":
if norm_text(joined) == norm_text(visible):
issues.append("dialogue_response_must_not_repeat_user_query")
if len(joined) < 8:
issues.append("dialogue_response_too_short")
if len(joined) > 260:
issues.append("dialogue_response_too_long")
elif qtype == "instruction":
quoted = extract_quoted_payload(visible)
if quoted and norm_text(joined) != norm_text(quoted):
issues.append("instruction_sample_text_should_equal_quoted_target_text")
if not quoted and INSTRUCTION_PREFIX_RE.search(joined):
issues.append("instruction_sample_text_looks_like_full_user_instruction")
if len(joined) > 220:
issues.append("instruction_target_text_too_long")
return sorted(set(issues))
def score_with_llm(
item: Dict[str, Any],
control: Dict[str, Any],
issues: Sequence[str],
args: argparse.Namespace,
) -> Tuple[int, str, List[str]]:
if issues:
return 0, "deterministic validation failed: " + "; ".join(issues), list(issues)
if args.no_eval:
return 5, "deterministic validation passed; LLM eval disabled", []
raw = chat_completion(
build_eval_messages(item, control, issues),
model=args.model,
base_url=args.base_url,
api_key_env=args.api_key_env,
temperature=0.1,
enable_thinking=args.enable_thinking,
stream=args.stream,
timeout=args.http_timeout,
)
parsed = extract_json_object(raw)
score = int(parsed.get("score", 0))
reason = str(parsed.get("reason", ""))
suggestions = parsed.get("repair_suggestions") or []
if not isinstance(suggestions, list):
suggestions = [str(suggestions)]
return score, reason, [str(x) for x in suggestions]
def normalize_row(
item: Dict[str, Any],
control: Dict[str, Any],
input_idx: int,
model: str,
score: int,
reason: str,
planning_reasoning: str,
) -> Dict[str, Any]:
qid = item.get("qid") or f"query_{input_idx:06d}"
sample_text = "".join(str(seg.get("sample_text", "")) for seg in control.get("Control", []))
language = item_language(item)
row = {
"qid": qid,
"query_type": item.get("query_type"),
"audio_content": sample_text,
"language": language,
"instruct_id": qid,
"file_name": f"response_{input_idx:06d}",
"ability": "response_tts_from_query",
"visible_query": item.get("visible_query"),
"source_event": item.get("source_event", ""),
"query_voice": item.get("query_voice", {}),
"target_contract": item.get("target_contract", {}),
"source_query_candidate": {
"language": language,
"source": item.get("source"),
"visible_query": item.get("visible_query"),
"hidden_context": item.get("hidden_context"),
"source_event": item.get("source_event"),
"target_contract": item.get("target_contract"),
"quality_hints": item.get("quality_hints"),
},
CONTROL_KEY: control,
"final_generated_control": control,
"response_control_generation": {
"model": model,
"prompt_version": "query_response_control_v1_grounded",
"score": score,
"reason": reason,
"planning_reasoning": planning_reasoning,
},
}
return row
def generate_one(task: Tuple[int, Dict[str, Any]], args: argparse.Namespace) -> Dict[str, Any]:
input_idx, item = task
attempts: List[Dict[str, Any]] = []
previous_control: Optional[Dict[str, Any]] = None
feedback = ""
best: Optional[Tuple[Dict[str, Any], int, str, str]] = None
for attempt in range(1, args.max_attempts + 1):
try:
raw = chat_completion(
build_generation_messages(item, previous_control=previous_control, feedback=feedback),
model=args.model,
base_url=args.base_url,
api_key_env=args.api_key_env,
temperature=args.temperature,
enable_thinking=args.enable_thinking,
stream=args.stream,
timeout=args.http_timeout,
)
control, planning_reasoning = parse_control(raw)
issues = validate_control(item, control)
score, reason, suggestions = score_with_llm(item, control, issues, args)
attempts.append(
{
"attempt": attempt,
"score": score,
"reason": reason,
"issues": issues,
"repair_suggestions": suggestions,
"control_preview": control,
}
)
if best is None or score > best[1]:
best = (control, score, reason, planning_reasoning)
if score >= args.accept_score:
return {
"status": "ok",
"row": normalize_row(item, control, input_idx, args.model, score, reason, planning_reasoning),
"failed": None,
}
previous_control = control
feedback = reason
if suggestions:
feedback += "\n修复建议:" + json.dumps(suggestions, ensure_ascii=False)
except Exception as exc:
attempts.append({"attempt": attempt, "error": repr(exc)})
feedback = repr(exc)
failed = {
"qid": item.get("qid"),
"query_type": item.get("query_type"),
"status": "failed",
"model": args.model,
"attempts": attempts,
"visible_query": item.get("visible_query"),
}
if best is not None and args.keep_best:
control, score, reason, planning_reasoning = best
return {
"status": "ok",
"row": normalize_row(item, control, input_idx, args.model, score, reason, planning_reasoning),
"failed": failed,
}
return {"status": "failed", "row": None, "failed": failed}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--input", default=DEFAULT_INPUT)
parser.add_argument("--output", default=DEFAULT_OUTPUT)
parser.add_argument("--failed_output", default=DEFAULT_FAILED_OUTPUT)
parser.add_argument("--model", default="deepseek-v4-pro")
parser.add_argument("--base_url", default=DEFAULT_BASE_URL)
parser.add_argument("--api_key_env", default="DASHSCOPE_API_KEY")
parser.add_argument("--temperature", type=float, default=0.35)
parser.add_argument("--http_timeout", type=int, default=180)
parser.add_argument("--num_workers", type=int, default=4)
parser.add_argument("--limit", type=int, default=0)
parser.add_argument("--max_attempts", type=int, default=3)
parser.add_argument("--accept_score", type=int, default=4)
parser.add_argument("--resume", action="store_true")
parser.add_argument("--enable_thinking", action="store_true")
parser.add_argument("--stream", action="store_true")
parser.add_argument("--no_eval", action="store_true")
parser.add_argument("--keep_best", action="store_true")
return parser.parse_args()
def main() -> int:
args = parse_args()
input_path = Path(args.input)
output_path = Path(args.output)
failed_path = Path(args.failed_output)
output_path.parent.mkdir(parents=True, exist_ok=True)
failed_path.parent.mkdir(parents=True, exist_ok=True)
done = load_existing_qids(output_path) if args.resume else set()
indexed_items = [(idx, item) for idx, item in enumerate(iter_jsonl(input_path), start=1) if str(item.get("qid")) not in done]
if args.limit > 0:
indexed_items = indexed_items[: args.limit]
print(
f"[INFO] items={len(indexed_items)} resume_done={len(done)} output={output_path} failed_output={failed_path}",
flush=True,
)
ok_count = 0
fail_count = 0
with ThreadPoolExecutor(max_workers=max(1, args.num_workers)) as executor:
pending = {}
iterator = iter(indexed_items)
def submit_next() -> bool:
try:
task = next(iterator)
except StopIteration:
return False
pending[executor.submit(generate_one, task, args)] = task[1].get("qid")
return True
while len(pending) < args.num_workers and submit_next():
pass
processed = 0
while pending:
finished, _ = wait(pending, return_when=FIRST_COMPLETED)
ok_rows: List[Dict[str, Any]] = []
failed_rows: List[Dict[str, Any]] = []
for fut in finished:
pending.pop(fut, None)
result = fut.result()
processed += 1
if result["row"] is not None:
ok_rows.append(result["row"])
ok_count += 1
if result["failed"] is not None:
failed_rows.append(result["failed"])
if result["status"] == "failed":
fail_count += 1
submit_next()
append_jsonl(output_path, ok_rows)
append_jsonl(failed_path, failed_rows)
print(f"[PROGRESS] processed={processed}/{len(indexed_items)} ok={ok_count} failed={fail_count}", flush=True)
return 0
if __name__ == "__main__":
raise SystemExit(main())