feat: 新增外部数据源拉取与 DPO 输出格式支持
- 支持从 PostgreSQL 数据库拉取结构化数据作为训练来源 - 新增 DPO (Direct Preference Optimization) 输出类型 - 支持 chosen/rejected 字段的编辑、校验和发布 - 完善数据预处理切分逻辑和元数据管理 - 移除 OCR 扫描 PDF 功能,保持基础文本解析能力 Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -1840,6 +1840,43 @@ def test_external_source_never_returns_fake_success(tmp_path: Path) -> None:
|
||||
assert response.json()["detail"]["code"] == 501
|
||||
|
||||
|
||||
def test_external_source_mode_belongs_to_step_three_structured_task(tmp_path: Path) -> None:
|
||||
client, _, _ = make_client(tmp_path)
|
||||
local_task_id = client.post(
|
||||
"/modelTF/data-process",
|
||||
json={
|
||||
"name": "本地结构化任务",
|
||||
"process_type": "structured",
|
||||
"config": {"source_mode": "local"},
|
||||
},
|
||||
).json()["data"]["id"]
|
||||
rejected = client.post(
|
||||
f"/modelTF/data-process/{local_task_id}/external/test",
|
||||
json={"type": "mysql", "url": "mysql://db.example/test"},
|
||||
)
|
||||
assert rejected.status_code == 409
|
||||
|
||||
external_task_id = client.post(
|
||||
"/modelTF/data-process",
|
||||
json={
|
||||
"name": "外部结构化任务",
|
||||
"process_type": "structured",
|
||||
"config": {"source_mode": "external"},
|
||||
},
|
||||
).json()["data"]["id"]
|
||||
accepted_as_external = client.post(
|
||||
f"/modelTF/data-process/{external_task_id}/external/test",
|
||||
json={"type": "mysql", "url": "mysql://db.example/test"},
|
||||
)
|
||||
assert accepted_as_external.status_code == 501
|
||||
|
||||
local_upload = client.post(
|
||||
f"/modelTF/data-process/{external_task_id}/source-files",
|
||||
files={"files": ("records.jsonl", b'{"id":1}\n', "application/jsonl")},
|
||||
)
|
||||
assert local_upload.status_code == 409
|
||||
|
||||
|
||||
def test_regenerate_endpoint_prepares_an_existing_published_task(tmp_path: Path) -> None:
|
||||
client, store, _ = make_client(tmp_path)
|
||||
task_id = client.post(
|
||||
|
||||
@@ -88,6 +88,74 @@ def test_generate_model_records_uses_prompt_auth_and_stable_split() -> None:
|
||||
assert progress_updates == [(1, 1)]
|
||||
|
||||
|
||||
def test_generate_model_records_builds_native_dpo_pair() -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
payload = json.loads(request.content)
|
||||
system_prompt = payload["messages"][0]["content"]
|
||||
assert '"chosen"' in system_prompt
|
||||
assert '"rejected"' in system_prompt
|
||||
assert "直接偏好优化" in system_prompt
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": json.dumps({
|
||||
"items": [{
|
||||
"instruction": "系统如何处理扫描 PDF?",
|
||||
"input": "",
|
||||
"chosen": "仅在没有文本层时调用 OCR,并保留页码。",
|
||||
"rejected": "所有 PDF 都重复执行 OCR。",
|
||||
}],
|
||||
}, ensure_ascii=False),
|
||||
},
|
||||
}],
|
||||
},
|
||||
)
|
||||
|
||||
records = generate_model_records(
|
||||
[{"id": "preview-dpo", "edited_content": "扫描 PDF 缺少文本层时执行 OCR。"}],
|
||||
model={"name": "model", "api_url": "https://model.example/v1"},
|
||||
config={"output_type": "dpo", "generation_retries": 0},
|
||||
task_id="task-dpo",
|
||||
split={"train": 100, "validation": 0, "test": 0},
|
||||
qa_pairs_per_item=1,
|
||||
client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
|
||||
assert records[0]["status"] == "valid"
|
||||
assert records[0]["chosen"] == "仅在没有文本层时调用 OCR,并保留页码。"
|
||||
assert records[0]["rejected"] == "所有 PDF 都重复执行 OCR。"
|
||||
assert records[0]["output"] == records[0]["chosen"]
|
||||
|
||||
|
||||
def test_generate_model_records_rejects_equal_dpo_pair() -> None:
|
||||
def handler(_request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"choices": [{"message": {"content": json.dumps({
|
||||
"items": [{
|
||||
"instruction": "问题",
|
||||
"chosen": "相同回答",
|
||||
"rejected": "相同回答",
|
||||
}],
|
||||
}, ensure_ascii=False)}}]},
|
||||
)
|
||||
|
||||
records = generate_model_records(
|
||||
[{"id": "preview-dpo-invalid", "edited_content": "来源"}],
|
||||
model={"name": "model", "api_url": "https://model.example/v1"},
|
||||
config={"output_type": "dpo", "generation_retries": 0},
|
||||
task_id="task-dpo-invalid",
|
||||
split={"train": 100, "validation": 0, "test": 0},
|
||||
qa_pairs_per_item=1,
|
||||
client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
|
||||
assert records[0]["status"] == "invalid"
|
||||
assert "chosen equals rejected" in records[0]["error"]
|
||||
|
||||
|
||||
def test_minimax_m3_uses_split_reasoning_and_completion_token_budget() -> None:
|
||||
requests: list[dict[str, object]] = []
|
||||
|
||||
|
||||
@@ -37,6 +37,10 @@ def test_runtime_migration_fails_fast_on_incompatible_schema() -> None:
|
||||
for value in ("idle", "queued", "running", "completed", "failed", "cancelled"):
|
||||
assert f"'{value}'" in sql
|
||||
assert "CREATE TABLE IF NOT EXISTS data_process_results" in sql
|
||||
assert "ADD COLUMN IF NOT EXISTS chosen TEXT NOT NULL DEFAULT ''" in sql
|
||||
assert "ADD COLUMN IF NOT EXISTS rejected TEXT NOT NULL DEFAULT ''" in sql
|
||||
assert "ADD COLUMN IF NOT EXISTS original_chosen TEXT" in sql
|
||||
assert "ADD COLUMN IF NOT EXISTS original_rejected TEXT" in sql
|
||||
assert sql.count("BEGIN;") == 1
|
||||
assert sql.rstrip().endswith("COMMIT;")
|
||||
|
||||
|
||||
@@ -1335,6 +1335,39 @@ def test_publish_rejects_invalid_reasoning_output_format() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_publish_dpo_writes_chosen_and_rejected_jsonl() -> None:
|
||||
conn = _PublishConnection(
|
||||
[
|
||||
{
|
||||
"id": "result-dpo",
|
||||
"status": "valid",
|
||||
"instruction": "如何处理扫描 PDF?",
|
||||
"input": "",
|
||||
"output": "仅在无文本层时执行 OCR。",
|
||||
"chosen": "仅在无文本层时执行 OCR。",
|
||||
"rejected": "所有 PDF 都执行 OCR。",
|
||||
"preview_item_id": "preview-dpo",
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
published = _PublishStore(conn, {"output_type": "dpo"}).publish(
|
||||
"task-dpo",
|
||||
{
|
||||
"dataset_name": "偏好数据",
|
||||
"storage_type": "local",
|
||||
"format": "dpo",
|
||||
"split": {"train": 100, "validation": 0, "test": 0},
|
||||
},
|
||||
)
|
||||
|
||||
record = conn.records[0]["raw"]
|
||||
assert record["chosen"] == "仅在无文本层时执行 OCR。"
|
||||
assert record["rejected"] == "所有 PDF 都执行 OCR。"
|
||||
assert "output" not in record
|
||||
assert published["datasets"][0]["metadata"]["format"] == "dpo"
|
||||
|
||||
|
||||
def test_source_storage_descriptor_accepts_owned_local_and_legacy_db_references() -> None:
|
||||
task_id = "dpt_task"
|
||||
source_file_id = "dpsf_source"
|
||||
|
||||
Reference in New Issue
Block a user