feat: 新增外部数据源拉取与 DPO 输出格式支持

- 支持从 PostgreSQL 数据库拉取结构化数据作为训练来源
- 新增 DPO (Direct Preference Optimization) 输出类型
- 支持 chosen/rejected 字段的编辑、校验和发布
- 完善数据预处理切分逻辑和元数据管理
- 移除 OCR 扫描 PDF 功能,保持基础文本解析能力

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
caoxiaozhu
2026-08-11 14:17:45 +08:00
parent f809825a7d
commit 5f6e7523cf
26 changed files with 927 additions and 144 deletions

View File

@@ -57,7 +57,62 @@ def _sentence_chunks(text: str) -> list[str]:
@lru_cache(maxsize=1)
def _tokenizer() -> tiktoken.Encoding:
return tiktoken.get_encoding("cl100k_base")
"""加载 cl100k_base 编码器,优先在线下载,失败时使用本地缓存以支持离线环境。"""
import os
import base64
# 先设置缓存目录环境变量
offline_cache = os.path.expanduser("~/.cache/tiktoken")
os.environ.setdefault("TIKTOKEN_CACHE_DIR", offline_cache)
try:
# 尝试标准方式加载
return tiktoken.get_encoding("cl100k_base")
except Exception:
# 如果失败,尝试手动从本地文件构造
try:
from pathlib import Path
local_file = Path(offline_cache) / "9b5ad71b2ce5302211f9c61530b329a4922fc6a4"
if not local_file.exists():
# 尝试另一个可能的文件名
local_file = Path(offline_cache) / "cl100k_base.tiktoken"
if local_file.exists():
# 读取 BPE 文件内容
with open(local_file, "rb") as f:
contents = f.read()
# 解析 BPE 文件
mergeable_ranks = {}
for line in contents.splitlines():
if line:
token, rank = line.split()
mergeable_ranks[base64.b64decode(token)] = int(rank)
# 构造 Encoding 对象
import tiktoken.core
return tiktoken.core.Encoding(
name="cl100k_base",
pat_str=r"""'(?i:[sdmt]|ll|ve|re)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]++[\r\n]*|\s*[\r\n]|\s+(?!\S)|\s+""",
mergeable_ranks=mergeable_ranks,
special_tokens={
"<|endoftext|>": 100257,
"<|fim_prefix|>": 100258,
"<|fim_middle|>": 100259,
"<|fim_suffix|>": 100260,
"<|endofprompt|>": 100276,
},
)
except Exception:
pass
raise RuntimeError(
f"无法加载 cl100k_base 编码器\n"
f"请确保以下任一条件满足:\n"
f"1. 服务器可以访问网络\n"
f"2. 本地存在缓存文件: {offline_cache}/9b5ad71b2ce5302211f9c61530b329a4922fc6a4"
)
def _text_chunks(

Binary file not shown.

View File

@@ -29,7 +29,12 @@ class _TerminalModelGenerationError(ModelGenerationError):
OUTPUT_TYPE_STANDARD = "standard"
OUTPUT_TYPE_REASONING = "reasoning"
SUPPORTED_OUTPUT_TYPES = {OUTPUT_TYPE_STANDARD, OUTPUT_TYPE_REASONING}
OUTPUT_TYPE_DPO = "dpo"
SUPPORTED_OUTPUT_TYPES = {
OUTPUT_TYPE_STANDARD,
OUTPUT_TYPE_REASONING,
OUTPUT_TYPE_DPO,
}
REASONING_DETAIL_NORMAL = "normal"
REASONING_DETAIL_DETAILED = "detailed"
SUPPORTED_REASONING_DETAILS = {
@@ -285,6 +290,18 @@ def _prompt_messages(
"这是思维链输出模式,即使其他提示语要求省略分析,也不得省略 reasoning。"
"不要自行添加 <think> 标签,系统会在保存时统一组装。"
)
elif output_type == OUTPUT_TYPE_DPO:
schema = (
'{"items":[{"instruction":"...","input":"...",'
'"chosen":"...","rejected":"..."}]}'
)
output_rule = (
"你正在生成用于直接偏好优化DPO的成对偏好数据。"
"instruction、chosen 和 rejected 均不得为空chosen 必须是忠于来源、"
"准确完整的优选回答rejected 必须是表面合理但存在明确质量缺陷的拒选回答。"
"两者不得相同rejected 不得包含违法危险内容,也不得用空白、乱码或无关文本凑数。"
"不要输出分析过程或 <think> 标签。"
)
else:
schema = '{"items":[{"instruction":"...","input":"...","output":"..."}]}'
output_rule = (
@@ -461,9 +478,13 @@ def generate_model_records(
"instruction": failure_instruction,
"input": content,
"output": "",
"chosen": "",
"rejected": "",
"original_instruction": failure_instruction,
"original_input": content,
"original_output": "",
"original_chosen": "",
"original_rejected": "",
"status": "invalid",
"error": error_message,
"split": "train",
@@ -479,6 +500,8 @@ def generate_model_records(
input_text = normalize_text(
str(value.get("input") or value.get("context") or "")
)
chosen = ""
rejected = ""
if output_type == OUTPUT_TYPE_REASONING:
reasoning = normalize_text(
re.sub(
@@ -508,6 +531,36 @@ def generate_model_records(
)
valid = bool(instruction and reasoning and answer)
missing_error = "model result is missing instruction, reasoning or answer"
elif output_type == OUTPUT_TYPE_DPO:
chosen = normalize_text(str(value.get("chosen") or ""))
rejected = normalize_text(str(value.get("rejected") or ""))
chosen = normalize_text(
re.sub(
r"<think>[\s\S]*?(?:</think>|$)",
"",
chosen,
flags=re.IGNORECASE,
)
)
rejected = normalize_text(
re.sub(
r"<think>[\s\S]*?(?:</think>|$)",
"",
rejected,
flags=re.IGNORECASE,
)
)
output = chosen
valid = bool(
instruction
and chosen
and rejected
and chosen.strip() != rejected.strip()
)
missing_error = (
"model result is missing instruction, chosen or rejected, "
"or chosen equals rejected"
)
else:
output = normalize_text(
str(
@@ -527,7 +580,10 @@ def generate_model_records(
)
valid = bool(instruction and output)
missing_error = "model result is missing instruction or output"
raw_id = f"{preview_id}:{variant_index + 1}:{instruction}:{output}"
raw_id = (
f"{preview_id}:{variant_index + 1}:{instruction}:"
f"{output}:{rejected}"
)
result_id = f"result_{hashlib.sha256(raw_id.encode()).hexdigest()[:16]}"
results.append(
{
@@ -536,9 +592,13 @@ def generate_model_records(
"instruction": instruction,
"input": input_text,
"output": output,
"chosen": chosen,
"rejected": rejected,
"original_instruction": instruction,
"original_input": input_text,
"original_output": output,
"original_chosen": chosen,
"original_rejected": rejected,
"status": "valid" if valid else "invalid",
"error": (None if valid else missing_error),
"split": "train",

View File

@@ -140,6 +140,12 @@ def _reasoning_output_is_valid(value: Any) -> bool:
)
def _dpo_fields_are_valid(row: dict[str, Any]) -> bool:
chosen = str(row.get("chosen") or "").strip()
rejected = str(row.get("rejected") or "").strip()
return bool(chosen and rejected and chosen != rejected)
def _preview_config_value(config: dict[str, Any], key: str, default: Any) -> Any:
if key in config:
return config[key]
@@ -846,7 +852,11 @@ class DataProcessStore:
)
instruction = str(record.get("instruction") or raw.get("instruction") or "")
input_text = str(record.get("input") or raw.get("input") or "")
output = str(record.get("output") or raw.get("output") or "")
chosen = str(raw.get("chosen") or "")
rejected = str(raw.get("rejected") or "")
output = str(
record.get("output") or raw.get("output") or chosen or ""
)
split = str(record.get("split") or raw.get("split") or "") or None
status = str(record.get("status") or "valid")
if status not in {"valid", "modified", "invalid"}:
@@ -856,10 +866,11 @@ class DataProcessStore:
"""
INSERT INTO data_process_results
(id, task_id, preview_item_id, instruction, input, output,
original_instruction, original_input, original_output, status,
chosen, rejected, original_instruction, original_input,
original_output, original_chosen, original_rejected, status,
error, split, quality_score, created_at, updated_at)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
NULL, %s, '{}', %s, %s)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
%s, %s, NULL, %s, '{}', %s, %s)
""",
(
result_id,
@@ -868,9 +879,13 @@ class DataProcessStore:
instruction,
input_text,
output,
chosen,
rejected,
instruction,
input_text,
output,
chosen,
rejected,
status,
split,
created_at,
@@ -2024,9 +2039,11 @@ class DataProcessStore:
"""
INSERT INTO data_process_results
(id, task_id, preview_item_id, instruction, input, output,
original_instruction, original_input, original_output, status, error,
chosen, rejected, original_instruction, original_input,
original_output, original_chosen, original_rejected, status, error,
split, quality_score, created_at, updated_at)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
%s, %s, %s, %s, %s, %s, %s)
""",
(
result.get("id") or new_id("dpr"),
@@ -2035,9 +2052,13 @@ class DataProcessStore:
result.get("instruction") or "",
result.get("input") or "",
result.get("output") or "",
result.get("chosen") or "",
result.get("rejected") or "",
result.get("original_instruction", result.get("instruction") or ""),
result.get("original_input", result.get("input") or ""),
result.get("original_output", result.get("output") or ""),
result.get("original_chosen", result.get("chosen") or ""),
result.get("original_rejected", result.get("rejected") or ""),
result.get("status") or "valid",
result.get("error"),
result.get("split"),
@@ -2121,7 +2142,7 @@ class DataProcessStore:
rows = conn.execute(
"""
SELECT status, instruction, output
SELECT status, instruction, output, chosen, rejected
FROM data_process_results
WHERE task_id=%s
""",
@@ -2139,6 +2160,10 @@ class DataProcessStore:
_task_output_type(task) == "reasoning"
and not _reasoning_output_is_valid(row.get("output"))
)
or (
_task_output_type(task) == "dpo"
and not _dpo_fields_are_valid(row)
)
)
if invalid_count:
raise InvalidStateError(
@@ -2210,7 +2235,9 @@ class DataProcessStore:
def update_result(
self, task_id: str, result_id: str, payload: dict[str, Any]
) -> dict[str, Any]:
allowed = {"instruction", "input", "output", "quality_score"}
allowed = {
"instruction", "input", "output", "chosen", "rejected", "quality_score"
}
values = {key: value for key, value in payload.items() if key in allowed}
if "quality_score" in values:
values["quality_score"] = json_dumps(values["quality_score"])
@@ -2232,20 +2259,28 @@ class DataProcessStore:
current_updated_at = _serialize_value(current.get("updated_at"))
if expected_updated_at and expected_updated_at != current_updated_at:
raise ConflictError("data process result was modified by another request")
output_type = _task_output_type(task)
if output_type == "dpo" and "chosen" in values:
values["output"] = values["chosen"]
merged = {**current, **values}
quality = payload.get("quality_score") or {}
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"
output_type != "reasoning"
or _reasoning_output_is_valid(merged.get("output"))
)
hard_valid = instruction_valid and output_valid and reasoning_valid
dpo_valid = output_type != "dpo" or _dpo_fields_are_valid(merged)
hard_valid = instruction_valid and output_valid and reasoning_valid and dpo_valid
quality_valid = bool(quality.get("is_valid", hard_valid))
changed = any(
str(merged.get(field) or "")
!= str(merged.get(f"original_{field}") or "")
for field in ("instruction", "input", "output")
for field in (
("instruction", "input", "chosen", "rejected")
if output_type == "dpo"
else ("instruction", "input", "output")
)
)
status = "invalid" if not hard_valid or not quality_valid else (
"modified" if changed else "valid"
@@ -2255,6 +2290,8 @@ class DataProcessStore:
format_error = (
"思维链输出必须包含非空的 <think>...</think> 推理过程和最终答案"
if instruction_valid and output_valid and not reasoning_valid
else "DPO 输出必须包含不同的非空 Chosen 和 Rejected 回答"
if instruction_valid and not dpo_valid
else "Instruction 和 Output 不能为空"
if not instruction_valid or not output_valid
else None
@@ -2318,18 +2355,23 @@ class DataProcessStore:
instruction = str(replacement.get("instruction") or "").strip()
input_text = str(replacement.get("input") or "").strip()
output = str(replacement.get("output") or "").strip()
chosen = str(replacement.get("chosen") or "").strip()
rejected = str(replacement.get("rejected") or "").strip()
quality_score = replacement.get("quality_score") or {}
if not instruction or not output or not bool(quality_score.get("is_valid")):
raise InvalidStateError("regenerated result did not pass quality validation")
if _task_output_type(task) == "reasoning" and not _reasoning_output_is_valid(output):
raise InvalidStateError("regenerated reasoning result has an invalid output format")
if _task_output_type(task) == "dpo" and not _dpo_fields_are_valid(replacement):
raise InvalidStateError("regenerated DPO result has invalid preference fields")
now = utcnow()
row = conn.execute(
"""
UPDATE data_process_results
SET instruction=%s, input=%s, output=%s,
SET instruction=%s, input=%s, output=%s, chosen=%s, rejected=%s,
original_instruction=%s, original_input=%s, original_output=%s,
original_chosen=%s, original_rejected=%s,
status='valid', error=NULL, quality_score=%s, updated_at=%s
WHERE id=%s AND task_id=%s
RETURNING *
@@ -2338,9 +2380,13 @@ class DataProcessStore:
instruction,
input_text,
output,
chosen,
rejected,
instruction,
input_text,
output,
chosen,
rejected,
json_dumps(quality_score),
now,
result_id,
@@ -2434,6 +2480,10 @@ class DataProcessStore:
_task_output_type(task) == "reasoning"
and not _reasoning_output_is_valid(row.get("output"))
)
or (
_task_output_type(task) == "dpo"
and not _dpo_fields_are_valid(row)
)
)
if invalid_count:
raise InvalidStateError(f"task contains {invalid_count} invalid results")
@@ -2449,15 +2499,27 @@ class DataProcessStore:
requested_split,
seed=task_id,
)
records = [
{
"instruction": row["instruction"],
"input": row["input"],
"output": row["output"],
"split": assignment,
}
for row, assignment in zip(rows, assignments, strict=True)
]
if _task_output_type(task) == "dpo":
records = [
{
"instruction": row["instruction"],
"input": row["input"],
"chosen": row["chosen"],
"rejected": row["rejected"],
"split": assignment,
}
for row, assignment in zip(rows, assignments, strict=True)
]
else:
records = [
{
"instruction": row["instruction"],
"input": row["input"],
"output": row["output"],
"split": assignment,
}
for row, assignment in zip(rows, assignments, strict=True)
]
split_order = ("train", "validation", "test")
split_counts = {
split_name: assignments.count(split_name) for split_name in split_order
@@ -2498,7 +2560,11 @@ class DataProcessStore:
"reasoning_detail": _task_reasoning_detail(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",
"format": (
"dpo"
if _task_output_type(task) == "dpo"
else payload.get("format") or "alpaca_jsonl"
),
"split": requested_split,
}
@@ -2741,7 +2807,7 @@ class DataProcessStore:
record["split"],
record["instruction"],
record["input"],
record["output"],
record.get("output") or record.get("chosen") or "",
json_dumps(
{
**record,