diff --git a/README.md b/README.md index 0956bc2..5d0b07c 100644 --- a/README.md +++ b/README.md @@ -50,16 +50,15 @@ YG_FT/ ## 前后端一键启动 -首次使用前,请先按下方“后端启动”和“前端启动”说明安装依赖,并确保 -PostgreSQL 已可用。之后在项目根目录执行: +首次使用前请确保前端依赖已安装、PostgreSQL 已可用。之后在项目根目录执行: ```bash bash ./start.sh ``` -脚本会同时启动前端 `http://localhost:16801` 和后端 -`http://127.0.0.1:17861`,按 `Ctrl+C` 会同时停止两个服务。脚本只负责 -启动前后端,不会自动安装依赖,也不会启动 PostgreSQL、Redis 或算力服务。 +脚本会自动补装后端 `requirements.txt`,然后同时启动前端 `http://localhost:16801` 和后端 +`http://127.0.0.1:17861`,按 `Ctrl+C` 会同时停止两个服务。脚本不会自动安装 +前端依赖,也不会启动 PostgreSQL、Redis 或算力服务。 仅检查依赖和端口而不启动服务: diff --git a/backend/app/api/v1/endpoints/data_process.py b/backend/app/api/v1/endpoints/data_process.py index 3e8611f..b706779 100644 --- a/backend/app/api/v1/endpoints/data_process.py +++ b/backend/app/api/v1/endpoints/data_process.py @@ -16,7 +16,7 @@ from dataclasses import asdict from pathlib import Path from threading import BoundedSemaphore, Lock from typing import Any, Literal -from urllib.parse import quote, urlsplit +from urllib.parse import parse_qs, quote, urlsplit import httpx import psycopg @@ -676,8 +676,10 @@ def _run_generation( qa_pairs_per_item=int(pairs or 1), on_progress=report_progress, ) - elif output_type == "reasoning": - raise InvalidStateError("思维链输出必须配置可用的数据生成模型") + elif output_type in {"reasoning", "dpo"}: + if output_type == "reasoning": + raise InvalidStateError("思维链输出必须配置可用的数据生成模型") + raise InvalidStateError("DPO 输出必须配置可用的数据生成模型") else: generated = generate_standard_records( preview_items, @@ -1077,7 +1079,10 @@ async def upload_source_files( with api_errors(): task = store.get_task(task_id) process_type = str(task["process_type"]) - if process_type == "external": + source_mode = str( + _value(task.get("config") or {}, "source_mode", "sourceMode", "local") + ) + if process_type == "external" or source_mode == "external": raise InvalidStateError( "external tasks must import data through the external source endpoint" ) @@ -1375,6 +1380,10 @@ def _external_postgres_connection(payload: ExternalSourceRequest) -> psycopg.Con raise fail(501, f"external data source type is not supported: {payload.type}") if parsed_url.username or parsed_url.password: raise fail(400, "database credentials must use the account and password fields") + if set(parse_qs(parsed_url.query)) & { + "password", "secret", "token", "api_key", "user", "username" + }: + raise fail(400, "database URL query must not contain credentials") if payload.auth_mode not in {"none", "basic"}: raise fail(400, "PostgreSQL supports only none or basic authentication") if payload.auth_mode == "basic" and not payload.username: @@ -1413,10 +1422,14 @@ def _external_postgres_connection(payload: ExternalSourceRequest) -> psycopg.Con "set DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB=true only in a trusted deployment", ) kwargs: dict[str, Any] = { - "connect_timeout": 5, + "connect_timeout": payload.connect_timeout_seconds, "row_factory": dict_row, "application_name": "yg-ft-data-process-readonly", - "options": "-c default_transaction_read_only=on -c statement_timeout=30000", + "options": ( + "-c default_transaction_read_only=on " + f"-c statement_timeout={payload.statement_timeout_seconds * 1000}" + ), + "sslmode": payload.ssl_mode, } if payload.auth_mode == "basic" and payload.username: kwargs["user"] = payload.username @@ -1425,6 +1438,17 @@ def _external_postgres_connection(payload: ExternalSourceRequest) -> psycopg.Con return psycopg.connect(payload.url, **kwargs) +def _assert_external_source_task(task: dict[str, Any]) -> None: + if str(task.get("process_type")) == "external": + return + config = task.get("config") or {} + source_mode = str(_value(config, "source_mode", "sourceMode", "local")) + if str(task.get("process_type")) != "structured" or source_mode != "external": + raise InvalidStateError( + "external source access requires a structured task with source_mode=external" + ) + + @router.post("/{task_id}/external/test") def test_external_source( task_id: str, @@ -1433,10 +1457,7 @@ def test_external_source( ) -> dict[str, Any]: with api_errors(): task = store.get_task(task_id) - if str(task.get("process_type")) != "external": - raise InvalidStateError( - "external source access requires an external data processing task" - ) + _assert_external_source_task(task) try: with _external_postgres_connection(payload) as conn: conn.execute("SELECT 1 AS ok").fetchone() @@ -1462,12 +1483,14 @@ def pull_external_source( raise fail(400, "a read-only SELECT or WITH query is required for external pull") with api_errors(): task = store.get_task(task_id) - if str(task.get("process_type")) != "external": - raise InvalidStateError("external pull requires an external data processing task") + _assert_external_source_task(task) try: with _external_postgres_connection(payload) as conn: conn.execute("SET TRANSACTION READ ONLY") - conn.execute("SET LOCAL statement_timeout = '30s'") + conn.execute( + "SELECT set_config('statement_timeout', %s, true)", + (f"{payload.statement_timeout_seconds}s",), + ) cursor = conn.execute(query) rows: list[dict[str, Any]] = [] content_parts: list[str] = [] @@ -1518,6 +1541,9 @@ def pull_external_source( "external_type": payload.type, "external_host": urlsplit(payload.url).hostname, "external_limit": payload.limit, + "external_ssl_mode": payload.ssl_mode, + "external_connect_timeout_seconds": payload.connect_timeout_seconds, + "external_statement_timeout_seconds": payload.statement_timeout_seconds, }, } ], @@ -1566,7 +1592,10 @@ def _prepare_preview_items( for index, source in enumerate(sources): source_format = str(source.get("file_format") or "").lower() needs_structured_xlsx = not is_unstructured and source_format == "xlsx" - needs_layout_raw = is_unstructured and chunk_method == "layout_hybrid" + needs_layout_raw = ( + is_unstructured + and chunk_method == "layout_hybrid" + ) needs_pdf_noise = ( is_unstructured and not needs_layout_raw @@ -2038,6 +2067,8 @@ def restore_result( "instruction": current.get("original_instruction") or current.get("instruction") or "", "input": current.get("original_input") or current.get("input") or "", "output": current.get("original_output") or current.get("output") or "", + "chosen": current.get("original_chosen") or current.get("chosen") or "", + "rejected": current.get("original_rejected") or current.get("rejected") or "", } preview_id = current.get("preview_item_id") source_content = "" @@ -2070,6 +2101,8 @@ def restore_result( "instruction": restored["instruction"], "input": restored["input"], "output": restored["output"], + "chosen": restored["chosen"], + "rejected": restored["rejected"], "quality_score": asdict(quality), "expected_updated_at": current.get("updated_at"), }, @@ -2141,13 +2174,15 @@ def _generate_result_replacement( ).strip().lower() previous_instruction = str(current.get("instruction") or "")[:1000] previous_output = str(current.get("output") or "")[:1000] + previous_rejected = str(current.get("rejected") or "")[:1000] base_prompt = str( _value(config, "generation_prompt", "generationPrompt", "") or "" ) regeneration_instruction = ( "这是一次失败结果的重新生成。请使用新的提问角度和表达," "不要复述旧结果。旧问题:" - f"{previous_instruction or '无'};旧答案:{previous_output or '无'}。" + f"{previous_instruction or '无'};旧优选答案:{previous_output or '无'};" + f"旧拒选答案:{previous_rejected or '无'}。" ) runtime_config = { **config, diff --git a/backend/app/db/sql/002_data_process.sql b/backend/app/db/sql/002_data_process.sql index b05cb81..a3bc93a 100644 --- a/backend/app/db/sql/002_data_process.sql +++ b/backend/app/db/sql/002_data_process.sql @@ -268,9 +268,13 @@ CREATE TABLE IF NOT EXISTS data_process_results ( instruction TEXT NOT NULL, input TEXT NOT NULL DEFAULT '', output TEXT NOT NULL, + chosen TEXT NOT NULL DEFAULT '', + rejected TEXT NOT NULL DEFAULT '', original_instruction TEXT, original_input TEXT, original_output TEXT, + original_chosen TEXT, + original_rejected TEXT, status VARCHAR(20) NOT NULL DEFAULT 'valid' CHECK (status IN ('valid', 'modified', 'invalid')), error TEXT, @@ -280,6 +284,11 @@ CREATE TABLE IF NOT EXISTS data_process_results ( updated_at TIMESTAMPTZ NOT NULL DEFAULT now() ); +ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS chosen TEXT NOT NULL DEFAULT ''; +ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS rejected TEXT NOT NULL DEFAULT ''; +ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_chosen TEXT; +ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_rejected TEXT; + CREATE INDEX IF NOT EXISTS idx_data_process_results_task_status ON data_process_results(task_id, status, id); CREATE INDEX IF NOT EXISTS idx_data_process_results_task_split diff --git a/backend/app/modules/data_process/document_chunking.py b/backend/app/modules/data_process/document_chunking.py index 61d00b3..b418e79 100644 --- a/backend/app/modules/data_process/document_chunking.py +++ b/backend/app/modules/data_process/document_chunking.py @@ -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( diff --git a/backend/app/modules/data_process/document_chunking.txt b/backend/app/modules/data_process/document_chunking.txt new file mode 100644 index 0000000..65f913d Binary files /dev/null and b/backend/app/modules/data_process/document_chunking.txt differ diff --git a/backend/app/modules/data_process/generation.py b/backend/app/modules/data_process/generation.py index 8845505..0168e72 100644 --- a/backend/app/modules/data_process/generation.py +++ b/backend/app/modules/data_process/generation.py @@ -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。" "不要自行添加 标签,系统会在保存时统一组装。" ) + elif output_type == OUTPUT_TYPE_DPO: + schema = ( + '{"items":[{"instruction":"...","input":"...",' + '"chosen":"...","rejected":"..."}]}' + ) + output_rule = ( + "你正在生成用于直接偏好优化(DPO)的成对偏好数据。" + "instruction、chosen 和 rejected 均不得为空;chosen 必须是忠于来源、" + "准确完整的优选回答,rejected 必须是表面合理但存在明确质量缺陷的拒选回答。" + "两者不得相同;rejected 不得包含违法危险内容,也不得用空白、乱码或无关文本凑数。" + "不要输出分析过程或 标签。" + ) 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"[\s\S]*?(?:|$)", + "", + chosen, + flags=re.IGNORECASE, + ) + ) + rejected = normalize_text( + re.sub( + r"[\s\S]*?(?:|$)", + "", + 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", diff --git a/backend/app/modules/data_process/store.py b/backend/app/modules/data_process/store.py index 06925b6..3d29369 100644 --- a/backend/app/modules/data_process/store.py +++ b/backend/app/modules/data_process/store.py @@ -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 = ( "思维链输出必须包含非空的 ... 推理过程和最终答案" 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, diff --git a/backend/app/schemas/data_process.py b/backend/app/schemas/data_process.py index 1798121..aa3671f 100644 --- a/backend/app/schemas/data_process.py +++ b/backend/app/schemas/data_process.py @@ -2,6 +2,7 @@ from __future__ import annotations from enum import StrEnum from typing import Any, Literal +from urllib.parse import parse_qs, urlsplit from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator @@ -15,6 +16,31 @@ def _config_value(config: dict[str, Any], snake_name: str, camel_name: str, defa def _validate_process_config(config: dict[str, Any]) -> None: + output_type = _config_value(config, "output_type", "outputType", "standard") + if output_type not in {"standard", "reasoning", "dpo"}: + raise ValueError("output_type must be one of: standard, reasoning, dpo") + + source_mode = _config_value(config, "source_mode", "sourceMode", "local") + if source_mode not in {"local", "external"}: + raise ValueError("source_mode must be one of: local, external") + external_source = _config_value(config, "external_source", "externalSource", None) + if external_source is not None: + if not isinstance(external_source, dict): + raise ValueError("external_source must be an object") + if any( + key.lower() in {"password", "secret", "token", "api_key"} + for key in external_source + ): + raise ValueError("external_source must not persist credentials") + external_url = str(external_source.get("url") or "").strip() + if external_url: + parsed_external_url = urlsplit(external_url) + sensitive_query_keys = {"password", "secret", "token", "api_key", "user", "username"} + if parsed_external_url.username or parsed_external_url.password or ( + set(parse_qs(parsed_external_url.query)) & sensitive_query_keys + ): + raise ValueError("external_source URL must not contain credentials") + chunk_method = _config_value(config, "chunk_method", "chunkMethod", "layout_hybrid") if not isinstance(chunk_method, str) or chunk_method not in { "layout_hybrid", @@ -303,6 +329,9 @@ class ExternalSourceRequest(BaseModel): username: str | None = Field(default=None, max_length=150) password: str | None = Field(default=None, max_length=500) limit: int = Field(default=1000, ge=1, le=100_000) + connect_timeout_seconds: int = Field(default=5, ge=1, le=30) + statement_timeout_seconds: int = Field(default=30, ge=1, le=300) + ssl_mode: Literal["disable", "prefer", "require", "verify-ca", "verify-full"] = "prefer" class ExternalPullRequest(ExternalSourceRequest): @@ -324,6 +353,8 @@ class ResultUpdate(BaseModel): instruction: str | None = None input: str | None = None output: str | None = None + chosen: str | None = None + rejected: str | None = None expected_updated_at: str | None = None @@ -374,7 +405,7 @@ class PublishRequest(BaseModel): dataset_type: Literal["train", "test", "eval", "val", "other"] = "train" storage_type: Literal["local"] = "local" split: DatasetSplit = Field(default_factory=DatasetSplit) - format: Literal["alpaca_jsonl", "jsonl"] = "alpaca_jsonl" + format: Literal["alpaca_jsonl", "jsonl", "dpo"] = "alpaca_jsonl" description: str = "" @field_validator("dataset_name") diff --git a/backend/tests/test_data_process_api.py b/backend/tests/test_data_process_api.py index c93c4dc..2170935 100644 --- a/backend/tests/test_data_process_api.py +++ b/backend/tests/test_data_process_api.py @@ -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( diff --git a/backend/tests/test_data_process_generation.py b/backend/tests/test_data_process_generation.py index 968548d..a7eaeb0 100644 --- a/backend/tests/test_data_process_generation.py +++ b/backend/tests/test_data_process_generation.py @@ -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]] = [] diff --git a/backend/tests/test_data_process_migration.py b/backend/tests/test_data_process_migration.py index a5b533a..bcfaa6a 100644 --- a/backend/tests/test_data_process_migration.py +++ b/backend/tests/test_data_process_migration.py @@ -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;") diff --git a/backend/tests/test_data_process_store.py b/backend/tests/test_data_process_store.py index f397bfc..a27aeeb 100644 --- a/backend/tests/test_data_process_store.py +++ b/backend/tests/test_data_process_store.py @@ -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" diff --git a/docker/app/Dockerfile.backend b/docker/app/Dockerfile.backend index c5a7195..3d654f3 100644 --- a/docker/app/Dockerfile.backend +++ b/docker/app/Dockerfile.backend @@ -2,7 +2,8 @@ FROM python:3.12-slim ENV PYTHONDONTWRITEBYTECODE=1 \ PYTHONUNBUFFERED=1 \ - PIP_NO_CACHE_DIR=1 + PIP_NO_CACHE_DIR=1 \ + TIKTOKEN_CACHE_DIR=/opt/tiktoken_cache WORKDIR /app @@ -16,6 +17,9 @@ RUN python -c "import fastapi, uvicorn, psycopg, psycopg_pool, sqlalchemy, redis RUN mkdir -p /opt/yg-ft/logs/backend /data/yg-ft \ && chmod -R 0775 /opt/yg-ft /data/yg-ft +# 离线打包 tiktoken cl100k_base 词表,避免无网环境下运行时联网下载 +COPY docker/app/tiktoken /opt/tiktoken_cache + EXPOSE 8000 CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/frontend/scripts/regression-data-process-detail.mjs b/frontend/scripts/regression-data-process-detail.mjs index 865194b..3a3ff0c 100644 --- a/frontend/scripts/regression-data-process-detail.mjs +++ b/frontend/scripts/regression-data-process-detail.mjs @@ -101,7 +101,9 @@ assert.match(detailSource, /!\/\(\?:password\|secret\|token\|api_key\)\/i\.test\ assert.match(detailSource, /key !== 'generation_model_snapshot'/, '处理配置仍直接展示内部模型快照') assert.match(detailSource, /preprocessOptionLabelMap/, '处理配置没有把预处理内部枚举转换为中文') assert.match(detailSource, /output_type:\s*'输出类型'/, '处理配置没有显示输出类型名称') -assert.match(detailSource, /value === 'reasoning' \? '思维链回答' : '标准回答'/, '处理配置没有转换输出类型枚举') +assert.match(detailSource, /value === 'reasoning' \? '思维链回答' : value === 'dpo' \? 'DPO 偏好对' : '标准回答'/, '处理配置没有转换输出类型枚举') +assert.match(detailSource, /prop="chosen"[\s\S]*?prop="rejected"/, 'DPO 结果没有展示 Chosen 与 Rejected') +assert.match(detailSource, /publishForm\.format = isDpoOutput\.value \? 'dpo'/, 'DPO 发布没有锁定原生格式') assert.match(detailSource, /reasoning_detail:\s*'推理详细程度'/, '处理配置没有显示推理详细程度名称') assert.match(detailSource, /value === 'detailed' \? '详细推理' : '普通推理'/, '处理配置没有转换推理详细程度枚举') assert.match(detailSource, /key !== 'reasoning_detail' \|\| config\.output_type === 'reasoning'/, '标准回答任务不应展示无关的推理详细程度') diff --git a/frontend/scripts/regression-data-process-wizard.mjs b/frontend/scripts/regression-data-process-wizard.mjs index 0172d63..d733599 100644 --- a/frontend/scripts/regression-data-process-wizard.mjs +++ b/frontend/scripts/regression-data-process-wizard.mjs @@ -122,13 +122,13 @@ assert.match(regenerationSource, /route\.name === 'data-process-regenerate'[\s\S assert.match(viewSource, /const currentStep = ref\(0\)/, '重新生成必须从向导第一步开始') assert.match(viewSource, /const WIZARD_STEPS = \[/, '向导步骤尚未改为固定常量') -for (const title of ['创建任务', '大模型选择', '上传文件', '数据预览', '开始生成', '结果编辑与保存']) { +for (const title of ['创建任务', '大模型选择', '数据来源', '数据预览', '开始生成', '结果编辑与保存']) { assert.ok(viewSource.includes(`title: '${title}'`), `缺少固定步骤:${title}`) } assert.match( viewSource, /\{ id: 'create',[\s\S]*?\{ id: 'model',[\s\S]*?\{ id: 'upload',[\s\S]*?\{ id: 'preview',[\s\S]*?\{ id: 'generate',[\s\S]*?\{ id: 'results'/, - '六步向导顺序必须为创建任务、大模型选择、上传文件、数据预览、开始生成、结果编辑与保存', + '六步向导顺序必须为创建任务、大模型选择、数据来源、数据预览、开始生成、结果编辑与保存', ) assert.doesNotMatch(viewSource, /steps\s*=\s*computed|all\.filter/, '步骤仍根据处理类型动态增减') assert.match( @@ -380,7 +380,7 @@ for (const marker of [']*>/, + /]*>/, '无文件时未保留原有大拖拽上传区或上传配置', ) assert.match( @@ -1432,7 +1432,7 @@ assert.ok(sourceUploadSource.includes('旧版 DOC/PPT 请先转换'), '非结构 assert.ok(sourceUploadSource.includes('旧版 XLS 请先转换'), '结构化格式提示没有说明旧版 XLS 需转换') assert.match( sourceUploadSource, - /[\s\S]*?已选择 \{\{ uploadedFiles\.length \}\} 个文件[\s\S]*?\s* emit\('file-change', file\)"[^>]*>\s*继续上传<\/el-button>\s*<\/el-upload>\s*<\/div>\s*<\/div>/, + /[\s\S]*?已选择 \{\{ uploadedFiles\.length \}\} 个文件[\s\S]*?\s*]*>\s*继续上传<\/el-button>\s*<\/el-upload>\s*<\/div>\s*<\/div>/, '有文件时缺少标题右侧的继续上传触发器或上传配置', ) assert.match( @@ -1476,4 +1476,13 @@ assert.match( ) assert.match(previewSource, /height:\s*clamp\(560px,\s*calc\(100vh - 370px\),\s*720px\)/, '对照预览高度不足以展示切片正文') +assert.doesNotMatch(taskSetupSource, /外来数据源拉取/, '外部数据源仍错误地放在第一步处理类型中') +assert.match(sourceUploadSource, /

数据来源<\/h3>[\s\S]*?本地上传[\s\S]*?外部数据源/, '第三步缺少本地与外部数据来源选择') +for (const field of ['地址 / URL', '鉴权方式', 'SSL 模式', '连接超时', '查询超时', '只读查询语句', '拉取条数', '落地文件名']) { + assert.ok(sourceUploadSource.includes(field), `外部数据源标准配置缺少:${field}`) +} +assert.match(generationControlSource, /label="DPO 偏好对" value="dpo"/, '输出类型缺少 DPO 偏好对') +assert.match(resultEditorSource, /Chosen[\s\S]*?Rejected/, '结果编辑器缺少 DPO 成对字段') +assert.match(viewSource, /sourceConfigForBackend\(sourceMode\.value, externalSource\)/, '任务配置没有保存第三步数据来源模式') + console.log('数据处理六步向导回归检查通过') diff --git a/frontend/src/types/dataProcess.ts b/frontend/src/types/dataProcess.ts index 3581c06..921387f 100644 --- a/frontend/src/types/dataProcess.ts +++ b/frontend/src/types/dataProcess.ts @@ -4,7 +4,7 @@ export type DataProcessStatus = 'pending' | 'running' | 'completed' | 'failed' | export type DataProcessType = 'structured' | 'unstructured' | 'external' export type DataProcessResultStatus = 'valid' | 'modified' | 'invalid' export type DataProcessSplit = 'train' | 'validation' | 'test' -export type DataProcessOutputType = 'standard' | 'reasoning' +export type DataProcessOutputType = 'standard' | 'reasoning' | 'dpo' export type DataProcessReasoningDetail = 'normal' | 'detailed' export type DataProcessWorkflowStep = 'create' | 'model' | 'upload' | 'preview' | 'generate' | 'results' export type DataProcessPreviewStatus = 'idle' | 'queued' | 'running' | 'completed' | 'failed' | 'cancelled' @@ -221,6 +221,9 @@ export interface DataProcessExternalSourcePayload { username?: string password?: string limit: number + connect_timeout_seconds: number + statement_timeout_seconds: number + ssl_mode: 'disable' | 'prefer' | 'require' | 'verify-ca' | 'verify-full' query?: string file_name?: string } @@ -342,9 +345,13 @@ export interface DataProcessResult { instruction: string input: string output: string + chosen?: string + rejected?: string original_instruction?: string | null original_input?: string | null original_output?: string | null + original_chosen?: string | null + original_rejected?: string | null status: DataProcessResultStatus error?: string | null split?: DataProcessSplit | null @@ -356,6 +363,8 @@ export interface DataProcessResultUpdatePayload { instruction: string input: string output: string + chosen?: string + rejected?: string expected_updated_at?: string } @@ -410,7 +419,7 @@ export interface DataProcessPublishPayload { dataset_type: 'train' | 'test' | 'eval' | 'val' | 'other' storage_type: 'local' split: DataProcessDatasetSplit - format: 'alpaca_jsonl' | 'jsonl' + format: 'alpaca_jsonl' | 'jsonl' | 'dpo' } export interface DataProcessPublishResult { diff --git a/frontend/src/views/data-process/DataProcessCreateView.vue b/frontend/src/views/data-process/DataProcessCreateView.vue index 79cf870..21951fc 100644 --- a/frontend/src/views/data-process/DataProcessCreateView.vue +++ b/frontend/src/views/data-process/DataProcessCreateView.vue @@ -20,6 +20,7 @@ import { import { useDataProcessGeneration } from './create/useDataProcessGeneration' import { useDataProcessPreviewBuild } from './create/useDataProcessPreviewBuild' import { useDataProcessRegeneration } from './create/useDataProcessRegeneration' +import { createDefaultExternalSource, externalSourcePayload, restoreExternalSourceConfig, sourceConfigForBackend } from './create/externalSourceConfig' import { loadCanonicalSourceContent, mapDataProcessSourceFile, @@ -39,26 +40,20 @@ import { updateDataProcessPreview, updateDataProcessTask, updateDataProcessWorkflowStep, - type DataProcessExternalSourcePayload, type DataProcessPreviewItem, } from '@/api/modules/dataProcess' -import type { - DataProcessConfig, - DataProcessPreviewProgress, - DataProcessTask, - DataProcessWorkflowStep, -} from '@/types/dataProcess' +import type { DataProcessConfig, DataProcessPreviewProgress, DataProcessTask, DataProcessWorkflowStep } from '@/types/dataProcess' import type { ExternalDataSource, GenerationControlOptions, PreviewItem, ProcessType, + SourceMode, StepId, StructuredProcessOptions, UnstructuredProcessOptions, UploadedDataFile, } from './create/types' - const router = useRouter() const modelsStore = useModelsStore() const { list: modelList, loaded: modelsLoaded } = storeToRefs(modelsStore) @@ -71,7 +66,7 @@ const PREVIEW_MODEL_VERSION = 'backend-pipeline-v4' const WIZARD_STEPS = [ { id: 'create', title: '创建任务', desc: '填写任务信息与处理配置' }, { id: 'model', title: '大模型选择', desc: '选择生成模型并设置输出要求' }, - { id: 'upload', title: '上传文件', desc: '上传或接入待处理的源数据' }, + { id: 'upload', title: '数据来源', desc: '选择本地上传或外部数据源拉取' }, { id: 'preview', title: '数据预览', desc: '核对源文件与预览内容' }, { id: 'generate', title: '开始生成', desc: '确认摘要并启动处理' }, { id: 'results', title: '结果编辑与保存', desc: '检查、修改并保存结果' }, @@ -81,11 +76,13 @@ const currentStepId = computed(() => WIZARD_STEPS[currentStep.value]?.id const task = reactive({ name: '', description: '' }) const taskId = ref(null) const processType = ref('structured') +const sourceMode = ref('local') const structuredOptions = ref(createDefaultStructuredOptions()) const unstructuredOptions = ref(createDefaultUnstructuredOptions()) const modelSelectionOptions = computed(() => ( processType.value === 'unstructured' ? unstructuredOptions.value : structuredOptions.value )) +const activeOutputType = computed(() => modelSelectionOptions.value.outputType) const uploadedFiles = ref([]) const previewBuilding = ref(false) const { @@ -93,16 +90,7 @@ const { startPreviewBuild, stopPreviewPolling, } = useDataProcessPreviewBuild() -const externalSource = reactive({ - type: 'postgresql', - url: '', - authMode: 'none', - username: '', - password: '', - limit: 1000, - query: '', - fileName: 'external-data.jsonl', -}) +const externalSource = reactive(createDefaultExternalSource()) const externalPulling = ref(false) const externalConnected = ref(false) const fileName = computed(() => uploadedFiles.value.map(f => f.name).join(', ')) @@ -135,6 +123,7 @@ const { } = useDataProcessGeneration({ taskId, dirty, + outputType: activeOutputType, beforeGenerate: beforeStartGeneration, }) const { enqueueSourceUpload, sourceUploading } = useDataProcessSourceUpload({ @@ -174,7 +163,7 @@ const previewFiles = computed(() => uploadedFiles.value.map((file) => { })) const primaryActionLabel = computed(() => { if (currentStepId.value === 'create') return '继续:选择大模型' - if (currentStepId.value === 'model') return '继续:上传文件' + if (currentStepId.value === 'model') return '继续:选择数据来源' if (currentStepId.value === 'upload') { if (sourceUploading.value) return '正在上传' return previewBuilding.value ? '正在切分' : '继续:数据预览' @@ -198,7 +187,6 @@ function goToStep(stepId: StepId) { const nextStepIndex = WIZARD_STEPS.findIndex((step) => step.id === stepId) if (nextStepIndex >= 0) currentStep.value = nextStepIndex } - function updateModelSelectionOptions(value: GenerationControlOptions) { if (processType.value === 'unstructured') { unstructuredOptions.value = { ...unstructuredOptions.value, ...value } @@ -206,12 +194,12 @@ function updateModelSelectionOptions(value: GenerationControlOptions) { } structuredOptions.value = { ...structuredOptions.value, ...value } } - function toBackendConfig(): DataProcessConfig { const options = processType.value === 'unstructured' ? unstructuredOptions.value : structuredOptions.value const common = { + ...sourceConfigForBackend(sourceMode.value, externalSource), preprocess_options: [...options.preprocessOptions], semantic_enrichment: options.semanticEnrichment, dataset_split: { ...options.datasetSplit }, @@ -247,7 +235,6 @@ function toBackendConfig(): DataProcessConfig { qa_pairs_per_row: structuredOptions.value.qaPairsPerRow, } } - function taskPayload() { return { name: task.name.trim(), @@ -256,12 +243,10 @@ function taskPayload() { config: toBackendConfig(), } } - async function persistWorkflowStep(step: StepId) { if (!taskId.value) return await updateDataProcessWorkflowStep(taskId.value, step as DataProcessWorkflowStep) } - async function saveTaskConfiguration() { if (isRegeneration.value) { const regenerated = await prepareRegeneration(taskPayload()) @@ -277,19 +262,6 @@ async function saveTaskConfiguration() { return saved } -function externalPayload(): DataProcessExternalSourcePayload { - return { - type: 'postgresql', - url: externalSource.url.trim(), - auth_mode: externalSource.authMode, - username: externalSource.username || undefined, - password: externalSource.password || undefined, - limit: externalSource.limit, - query: externalSource.query?.trim() || undefined, - file_name: externalSource.fileName || 'external-data.jsonl', - } -} - function mapPreviewItem(item: DataProcessPreviewItem): PreviewItem { const sourceLocator = item.quality_score?.source_locator return { @@ -378,7 +350,7 @@ async function beforeStartGeneration() { ) } watch( - [() => task.name, () => task.description, processType, structuredOptions, unstructuredOptions, externalSource], + [() => task.name, () => task.description, processType, sourceMode, structuredOptions, unstructuredOptions, externalSource], () => { if (!hydrating.value) dirty.value = true }, @@ -396,6 +368,13 @@ watch(processType, (nextType, previousType) => { ElMessage.info('处理类型已变更,请重新上传或拉取匹配的源数据') }) +watch(processType, (nextType) => { + if (nextType === 'unstructured' && sourceMode.value === 'external') { + sourceMode.value = 'local' + externalConnected.value = false + } +}) + watch(generationOptionsSignature, (currentSignature, previousSignature) => { if (hydrating.value) return if (currentSignature === previousSignature) return @@ -453,6 +432,16 @@ function updateExternalSource(value: ExternalDataSource) { externalConnected.value = false } +async function updateSourceMode(value: SourceMode) { + const previous = sourceMode.value + sourceMode.value = value + externalConnected.value = false + dirty.value = true + externalPulling.value = true + try { await saveTaskConfiguration() } catch { sourceMode.value = previous } + finally { externalPulling.value = false } +} + async function handleTestConnection() { if (!taskId.value) { ElMessage.error('任务尚未创建,请返回模型选择步骤后重试') @@ -464,7 +453,8 @@ async function handleTestConnection() { } externalPulling.value = true try { - const result = await testDataProcessExternalSource(taskId.value, externalPayload()) + await saveTaskConfiguration() + const result = await testDataProcessExternalSource(taskId.value, externalSourcePayload(externalSource)) externalConnected.value = result.connected if (result.connected) ElMessage.success(result.message || '数据源连接测试成功') else ElMessage.warning(result.message || '数据源连接失败') @@ -490,7 +480,8 @@ async function handlePullData() { } externalPulling.value = true try { - const response = await pullDataProcessExternalSource(taskId.value, externalPayload()) + await saveTaskConfiguration() + const response = await pullDataProcessExternalSource(taskId.value, externalSourcePayload(externalSource)) const newFiles: UploadedDataFile[] = [] for (const file of response.files) { const content = await loadCanonicalSourceContent(taskId.value, file.id) @@ -620,6 +611,12 @@ function applyPreviewProgress(progress: DataProcessPreviewProgress) { async function completePreviewWorkspace() { previewItems.value = await loadAllPreviewItems() + for (const file of uploadedFiles.value) { + if (!file.content && file.sourceFileId && file.fileFormat?.replace('.', '') === 'pdf') { + file.content = await loadCanonicalSourceContent(taskId.value as string, file.sourceFileId) + .catch(() => file.content) + } + } const configSignature = buildPreviewConfigSignature() const previewCounts = new Map() for (const item of previewItems.value) { @@ -678,7 +675,7 @@ async function nextFromUpload() { return } if (uploadedFiles.value.length === 0) { - ElMessage.warning(processType.value === 'external' ? '请先拉取至少一个数据源' : '请上传至少一个源数据文件') + ElMessage.warning(sourceMode.value === 'external' ? '请先拉取至少一个数据源' : '请上传至少一个源数据文件') return } const failedUploads = uploadedFiles.value.filter((file) => file.status === 'failed') @@ -992,6 +989,9 @@ async function initializeExistingWorkflow() { const sourceTask = await loadRegenerationSource() if (!sourceTask) return taskId.value = String(sourceTask.id) + const restoredSource = restoreExternalSourceConfig(sourceTask.config || {}) + sourceMode.value = restoredSource.mode + Object.assign(externalSource, restoredSource.source) if (!isWorkflowResume.value) return if (sourceTask.status === 'completed' && sourceTask.results_confirmed !== false) { allowLeave = true @@ -1102,6 +1102,7 @@ onMounted(() => { { :preview-building="previewBuilding" :source-uploading="sourceUploading" @update:external-source="updateExternalSource" + @update:source-mode="updateSourceMode" @file-change="handleFileChange" @remove-file="handleRemoveFile" @use-sample="useSampleFile" @@ -1154,6 +1156,7 @@ onMounted(() => { :preview-items="previewItems" :regenerating-result-id="regeneratingResultId" :bulk-regeneration="bulkRegeneration" + :output-type="activeOutputType" @update:field="updateResultField" @regenerate:all="regenerateAllResults" @regenerate:item="regenerateResult" @@ -1179,8 +1182,8 @@ onMounted(() => { {{ primaryActionLabel }} diff --git a/frontend/src/views/data-process/DataProcessDetailView.vue b/frontend/src/views/data-process/DataProcessDetailView.vue index 2a0ce81..ef0919b 100644 --- a/frontend/src/views/data-process/DataProcessDetailView.vue +++ b/frontend/src/views/data-process/DataProcessDetailView.vue @@ -53,7 +53,7 @@ const resultCellTooltipOptions = { } as const let resultFilterTimer: ReturnType | null = null -const editForm = reactive({ instruction: '', input: '', output: '' }) +const editForm = reactive({ instruction: '', input: '', output: '', chosen: '', rejected: '' }) const publishForm = reactive({ dataset_name: '', dataset_type: 'train', @@ -69,6 +69,8 @@ const processTypeMap: Record = { } const configLabelMap: Record = { + source_mode: '数据来源', + external_source: '外部数据源配置', preprocess_options: '预处理规则', dataset_split: '数据集划分', generation_model_id: '数据生成模型', @@ -212,6 +214,7 @@ const outputDatasetName = computed(() => ( )) const outputDatasetId = computed(() => detail.value?.output_dataset_id || null) const hasPublishedOutputs = computed(() => outputDatasets.value.length > 0) +const isDpoOutput = computed(() => detail.value?.config?.output_type === 'dpo') const hasCurrentPublishedDataset = computed(() => ( Boolean(outputDatasetId.value) && outputDatasets.value.some((dataset) => String(dataset.id) === String(outputDatasetId.value)) @@ -294,6 +297,18 @@ const configRows = computed(() => { }) function formatConfigValue(key: string, value: unknown) { + if (key === 'source_mode') return value === 'external' ? '外部数据源' : '本地上传' + if (key === 'external_source' && value && typeof value === 'object') { + const source = value as Record + return [ + source.type, + source.url, + source.ssl_mode ? `SSL ${source.ssl_mode}` : '', + source.limit ? `上限 ${source.limit} 条` : '', + source.connect_timeout_seconds ? `连接 ${source.connect_timeout_seconds}s` : '', + source.statement_timeout_seconds ? `查询 ${source.statement_timeout_seconds}s` : '', + ].filter(Boolean).join(' · ') + } if (key === 'generation_model_id') { const snapshot = detail.value?.config?.generation_model_snapshot if (snapshot && typeof snapshot === 'object' && !Array.isArray(snapshot)) { @@ -305,7 +320,7 @@ function formatConfigValue(key: string, value: unknown) { return chunkMethodLabelMap[value] || value } if (key === 'output_type') { - return value === 'reasoning' ? '思维链回答' : '标准回答' + return value === 'reasoning' ? '思维链回答' : value === 'dpo' ? 'DPO 偏好对' : '标准回答' } if (key === 'reasoning_detail') { return value === 'detailed' ? '详细推理' : '普通推理' @@ -461,14 +476,26 @@ function openResultEditor(result: DataProcessResult) { instruction: result.instruction, input: result.input, output: result.output, + chosen: result.chosen || result.output || '', + rejected: result.rejected || '', }) editDialogVisible.value = true } async function saveResult() { if (!editingResult.value) return - if (!editForm.instruction.trim() || !editForm.output.trim()) { - ElMessage.warning('Instruction 和 Output 不能为空') + const invalid = isDpoOutput.value + ? !editForm.instruction.trim() + || !editForm.chosen.trim() + || !editForm.rejected.trim() + || editForm.chosen.trim() === editForm.rejected.trim() + : !editForm.instruction.trim() || !editForm.output.trim() + if (invalid) { + ElMessage.warning( + isDpoOutput.value + ? 'Instruction、Chosen、Rejected 均不能为空,且两个回答不能相同' + : 'Instruction 和 Output 不能为空', + ) return } @@ -478,6 +505,8 @@ async function saveResult() { instruction: editForm.instruction, input: editForm.input, output: editForm.output, + chosen: editForm.chosen, + rejected: editForm.rejected, expected_updated_at: editingResult.value.updated_at, }) replaceResult(updated) @@ -530,6 +559,7 @@ function openPublishDialog() { if (!detail.value) return publishForm.dataset_name = `${detail.value.name}-数据集` publishForm.split = configuredSplit() + publishForm.format = isDpoOutput.value ? 'dpo' : 'alpaca_jsonl' publishDialogVisible.value = true } @@ -808,7 +838,11 @@ onBeforeUnmount(() => { - + +