feat(data-process): 支持思维链输出类型

This commit is contained in:
caoxiaozhu
2026-07-27 13:08:44 +08:00
parent ecafb7eb13
commit de2e8952b5
14 changed files with 443 additions and 46 deletions

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
import hashlib
import json
import re
import uuid
from collections.abc import Iterator, Sequence
from contextlib import contextmanager
@@ -87,6 +88,31 @@ def _json_value(value: Any, default: Any) -> Any:
return default
def _task_output_type(task: dict[str, Any]) -> str:
config = _json_value(task.get("config"), {})
if not isinstance(config, dict):
return "standard"
return str(config.get("output_type") or config.get("outputType") or "standard")
def _reasoning_output_is_valid(value: Any) -> bool:
match = re.fullmatch(
r"\s*<think>\s*(?P<reasoning>[\s\S]*?)\s*</think>\s*(?P<answer>[\s\S]+?)\s*",
str(value or ""),
flags=re.IGNORECASE,
)
return bool(
match
and match.group("reasoning").strip()
and match.group("answer").strip()
and all(
tag not in part.lower()
for tag in ("<think", "</think")
for part in (match.group("reasoning"), match.group("answer"))
)
)
def _preview_config_value(config: dict[str, Any], key: str, default: Any) -> Any:
if key in config:
return config[key]
@@ -1565,10 +1591,13 @@ class DataProcessStore:
raise ConflictError("data process result was modified by another request")
merged = {**current, **values}
quality = payload.get("quality_score") or {}
hard_valid = bool(
str(merged.get("instruction") or "").strip()
and str(merged.get("output") or "").strip()
instruction_valid = bool(str(merged.get("instruction") or "").strip())
output_valid = bool(str(merged.get("output") or "").strip())
reasoning_valid = (
_task_output_type(task) != "reasoning"
or _reasoning_output_is_valid(merged.get("output"))
)
hard_valid = instruction_valid and output_valid and reasoning_valid
quality_valid = bool(quality.get("is_valid", hard_valid))
changed = any(
str(merged.get(field) or "")
@@ -1580,8 +1609,16 @@ class DataProcessStore:
)
values["status"] = status
flags = quality.get("flags") if isinstance(quality, dict) else None
format_error = (
"思维链输出必须包含非空的 <think>...</think> 推理过程和最终答案"
if instruction_valid and output_valid and not reasoning_valid
else "Instruction 和 Output 不能为空"
if not instruction_valid or not output_valid
else None
)
values["error"] = ", ".join(str(flag) for flag in flags or []) or (
"quality validation failed" if status == "invalid" else None
format_error
or ("quality validation failed" if status == "invalid" else None)
)
values["updated_at"] = utcnow()
assignments = ", ".join(f"{key}=%s" for key in values)
@@ -1671,6 +1708,10 @@ class DataProcessStore:
if row["status"] == "invalid"
or not str(row.get("instruction") or "").strip()
or not str(row.get("output") or "").strip()
or (
_task_output_type(task) == "reasoning"
and not _reasoning_output_is_valid(row.get("output"))
)
)
if invalid_count:
raise InvalidStateError(f"task contains {invalid_count} invalid results")
@@ -1731,6 +1772,7 @@ class DataProcessStore:
"source": "data_process",
"storage_backend": "database",
"source_task_id": task_id,
"output_type": _task_output_type(task),
"source_file_ids": [item["id"] for item in self._source_ids(conn, task_id)],
"source_result_ids": source_result_ids,
"format": payload.get("format") or "alpaca_jsonl",