diff --git a/backend/app/api/v1/endpoints/platform.py b/backend/app/api/v1/endpoints/platform.py index cc76daf..1206d47 100644 --- a/backend/app/api/v1/endpoints/platform.py +++ b/backend/app/api/v1/endpoints/platform.py @@ -48,6 +48,21 @@ def _training_diagnostics(errors: list[str], warnings: list[str] | None = None, text = "\n".join(source_items).lower() diagnostics: list[dict[str, str]] = [] rules = [ + ( + ["api 模型", "api模型", "api model"], + "API 模型不能用于本地训练", + "当前选择的基座模型为 API 类型,LLaMA-Factory 需要本地可访问的模型路径。请在模型管理中创建或选择模型来源为「本地」且配置了算力节点路径的模型。", + ), + ( + ["未配置算力节点", "未配置.*路径", "模型.*路径"], + "模型缺少算力节点路径", + "请在模型管理中编辑该模型,设置模型路径为算力节点可访问的本地目录。", + ), + ( + ["不支持本地训练", "not trainable"], + "模型不可用于训练", + "当前选择的模型不支持作为 LLaMA-Factory 训练基座。请确认模型来源为本地、路径已配置且模型目录在算力节点上存在。", + ), ( ["dataset columns missing", "keyerror", "history", "instruction", "input", "output", "messages"], "训练数据字段不匹配", diff --git a/backend/app/db/platform_store.py b/backend/app/db/platform_store.py index 0d872bf..ab67a5c 100644 --- a/backend/app/db/platform_store.py +++ b/backend/app/db/platform_store.py @@ -153,23 +153,40 @@ def llama_dataset_keys(dataset_key: str, file_names: list[str]) -> list[str]: def llama_dataset_info(dataset_key: str, file_names: list[str], formatting: str = "alpaca") -> dict[str, Any]: result: dict[str, Any] = {} + fmt = str(formatting).lower() for key, file_name in zip(llama_dataset_keys(dataset_key, file_names), file_names): - if formatting == "sharegpt": + if fmt == "sharegpt": result[key] = { "file_name": file_name, "formatting": "sharegpt", "columns": {"messages": "messages"}, } - continue - result[key] = { - "file_name": file_name, - "formatting": "alpaca", - "columns": { - "prompt": "instruction", - "query": "input", - "response": "output", - }, - } + elif fmt == "dpo": + result[key] = { + "file_name": file_name, + "formatting": "dpo", + "columns": { + "prompt": "system", + "chosen": "chosen", + "rejected": "rejected", + }, + } + elif fmt in {"cpt", "pt", "pretrain"}: + result[key] = { + "file_name": file_name, + "formatting": "cpt", + "columns": {"prompt": "text"}, + } + else: + result[key] = { + "file_name": file_name, + "formatting": "alpaca", + "columns": { + "prompt": "instruction", + "query": "input", + "response": "output", + }, + } return result @@ -432,7 +449,7 @@ class PlatformStore: (status, progress, completed_at, row["id"]), ) - def _ensure_trained_model(self, conn: PgConnection, task: dict[str, Any]) -> None: + def _ensure_trained_model(self, conn: PgConnection, task: dict[str, Any], job: dict[str, Any] | None = None) -> None: name = task.get("output_model_name") or f"{task['name']}-lora" exists = conn.execute("SELECT id FROM trained_models WHERE name=?", (name,)).fetchone() if exists: @@ -443,8 +460,8 @@ class PlatformStore: conn.execute( """ INSERT INTO trained_models - (id, name, train_methods, base_model_path, create_time, merged, merging, merged_path) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) + (id, name, train_methods, base_model_path, create_time, merged, merging, merged_path, artifact_dir) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) """, ( trained_model_id, @@ -455,20 +472,52 @@ class PlatformStore: 0, 0, output_dir, + output_dir, ), ) + # Use real artifact data from compute node when available + artifacts = (job or {}).get("artifacts") or [] + if artifacts: + total_size = sum(int(a.get("size_bytes") or a.get("size", 0)) for a in artifacts) + checksums = [a.get("checksum_sha256", "") for a in artifacts if a.get("checksum_sha256")] + combined_checksum = checksums[0] if len(checksums) == 1 else "" + # Register individual artifact files + for artifact in artifacts[:50]: # limit to 50 file entries + artifact_path = artifact.get("path") or artifact.get("name", "") + abs_path = artifact_path if artifact_path.startswith("/") else f"{output_dir.rstrip('/')}/{artifact_path.lstrip('/')}" + self._upsert_model_artifact( + conn, + trained_model_id, + "trained_model", + "adapter_file", + abs_path, + int(artifact.get("size_bytes") or artifact.get("size", 0)), + artifact.get("checksum_sha256", ""), + { + "task_id": task.get("id"), + "train_method": task.get("train_method", "lora"), + "base_model": task.get("base_model"), + "artifact_name": artifact.get("name", ""), + }, + task.get("compute_job_id"), + ) + else: + total_size = 0 + combined_checksum = "" + # Register the top-level adapter directory entry self._upsert_model_artifact( conn, trained_model_id, "trained_model", "adapter", output_dir, - 0, - "", + total_size, + combined_checksum, { "task_id": task.get("id"), "train_method": task.get("train_method", "lora"), "base_model": task.get("base_model"), + "file_count": len(artifacts), }, task.get("compute_job_id"), ) @@ -1079,24 +1128,29 @@ class PlatformStore: def create_model(self, payload: dict[str, Any]) -> dict[str, Any]: model_id = payload.get("id") or new_id("m") + model_source = payload.get("model_source", "local") + path = payload.get("path", "") + # Automatically determine can_train: local models with a path can be trained + can_train = 1 if (model_source != "api" and path and str(path).strip()) else 0 with self.connect() as conn: conn.execute( """ INSERT INTO models - (id, name, type, purpose, model_source, description, path, api_url, api_key, online_model_name, create_time) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + (id, name, type, purpose, model_source, description, path, api_url, api_key, online_model_name, can_train, create_time) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) """, ( model_id, payload["name"], payload.get("type", "LLM"), payload.get("purpose", "training"), - payload.get("model_source", "local"), + model_source, payload.get("description"), - payload.get("path"), + path, payload.get("api_url"), payload.get("api_key"), payload.get("online_model_name"), + can_train, utcnow(), ), ) @@ -1105,23 +1159,28 @@ class PlatformStore: def update_model(self, model_id: str, payload: dict[str, Any]) -> dict[str, Any]: current = self.model(model_id) merged = {**current, **payload} + # Recompute can_train when relevant fields change + model_source = merged.get("model_source", "local") + path = merged.get("path", "") + can_train = 1 if (model_source != "api" and path and str(path).strip()) else 0 with self.connect() as conn: conn.execute( """ UPDATE models - SET name=?, type=?, purpose=?, model_source=?, description=?, path=?, api_url=?, api_key=?, online_model_name=? + SET name=?, type=?, purpose=?, model_source=?, description=?, path=?, api_url=?, api_key=?, online_model_name=?, can_train=? WHERE id=? """, ( merged["name"], merged.get("type", "LLM"), merged.get("purpose", "training"), - merged.get("model_source", "local"), + model_source, merged.get("description"), - merged.get("path"), + path, merged.get("api_url"), merged.get("api_key"), merged.get("online_model_name"), + can_train, model_id, ), ) @@ -1717,6 +1776,23 @@ class PlatformStore: dataset_id = str(task.get("train_dataset_id") or task.get("dataset_id") or "") with self.connect() as conn: model = conn.execute("SELECT * FROM models WHERE id=?", (base_model_id,)).fetchone() + if not model: + raise RuntimeError(f"base model not found: {base_model_id}") + # P0-1: Reject non-trainable models (API models or models without local path) + if not model.get("can_train"): + model_source = model.get("model_source") or "unknown" + model_path = model.get("path") or "" + if model_source == "api": + raise RuntimeError( + f"模型 '{model['name']}' 为 API 模型,不能作为 LLaMA-Factory 本地训练基座,请选择本地路径模型" + ) + if not model_path or not str(model_path).strip(): + raise RuntimeError( + f"模型 '{model['name']}' 未配置算力节点可访问路径,请先在模型管理中设置模型本地路径" + ) + raise RuntimeError( + f"模型 '{model['name']}' 不支持本地训练(source={model_source}),请选择其他模型" + ) dataset = conn.execute("SELECT * FROM datasets WHERE id=?", (dataset_id,)).fetchone() files = conn.execute( """SELECT id, name, size, active_version_id, create_time, metadata @@ -1767,6 +1843,24 @@ class PlatformStore: training_keys = runtime_keys[: len(training_files)] validation_keys = runtime_keys[len(training_files) :] dataset_format = str(task.get("dataset_format") or (dataset and dataset.get("formatting")) or "alpaca").lower() + # P0-2: Validate dataset content against declared format + train_type = str(task.get("train_type", task.get("train_method", ""))).upper() + expected_format = { + "DPO": "dpo", + "CPT": "cpt", + }.get(train_type) + if expected_format: + dataset_format = expected_format + format_errors: list[str] = [] + for file_entry in training_files: + content = file_entry.get("content") or "" + if content: + from app.modules.data_process.dataset_format import validate_dataset_format + file_errors = validate_dataset_format(dataset_format, content=content) + if file_errors: + format_errors.extend(file_errors) + if format_errors: + raise RuntimeError("数据集格式校验失败:\n" + "\n".join(f" - {e}" for e in format_errors[:10])) health_detail = node.get("health_detail") or {} dataset_root = str(health_detail.get("dataset_root") or f"{node['data_root'].rstrip('/')}/datasets") output_root = str(health_detail.get("output_root") or f"{node['data_root'].rstrip('/')}/outputs") @@ -1860,7 +1954,21 @@ class PlatformStore: self._sync_gpu_allocations(conn, payload, job, status) self._upsert_checkpoints(conn, task_id, job.get("checkpoints") or []) if status == "completed": - self._ensure_trained_model(conn, payload) + self._ensure_trained_model(conn, payload, job) + # P0-4: Persist failure info for diagnosis + if status in {"failed", "stopped"}: + failure_reason = job.get("error") or job.get("message") or "compute job failed" + log_snippet = job.get("log_snippet") or "" + conn.execute( + "UPDATE fine_tune_tasks SET failure_reason = ? WHERE id = ?", + (failure_reason[:2000], task_id), + ) + # Store last log snippet if available (max 8KB) + if log_snippet: + conn.execute( + "UPDATE fine_tune_tasks SET payload = ? WHERE id = ?", + (json_dumps({**payload, "last_log_snippet": log_snippet[:8192]}), task_id), + ) return self.task(task_id) def running_compute_tasks(self) -> list[dict[str, Any]]: diff --git a/backend/app/db/sql/003_model_path_governance.sql b/backend/app/db/sql/003_model_path_governance.sql new file mode 100644 index 0000000..0a1b795 --- /dev/null +++ b/backend/app/db/sql/003_model_path_governance.sql @@ -0,0 +1,18 @@ +-- 003_model_path_governance +-- 模型路径治理:增加 can_train 标识,区分本地可训练模型与 API / 远程模型。 +-- 训练预检阶段依赖该字段拦截不适合 LLaMA-Factory 本地训练的基座模型。 + +-- 1. models 表增加 can_train(默认 0,后设搬迁为 1 的规则如下) +ALTER TABLE models ADD COLUMN IF NOT EXISTS can_train INTEGER NOT NULL DEFAULT 0; + +-- 2. 将已有模型按规则推定 can_train: +-- - path 非空 且 model_source != 'api' → 可训练 (1) +-- - 其余 → 不可训练 (0) +UPDATE models +SET can_train = CASE + WHEN path IS NOT NULL AND path != '' AND model_source IS NOT NULL AND model_source != 'api' THEN 1 + ELSE 0 +END; + +-- 3. 给 trained_models 增加 artifact_dir(训练产物目录扫描结果目录) +ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS artifact_dir TEXT; diff --git a/backend/app/modules/compute_gateway/sync.py b/backend/app/modules/compute_gateway/sync.py index 17e97c9..0e2720a 100644 --- a/backend/app/modules/compute_gateway/sync.py +++ b/backend/app/modules/compute_gateway/sync.py @@ -27,6 +27,13 @@ async def poll_compute_jobs_once() -> dict[str, Any]: store.record_training_log_metrics(task["id"], str(logs.get("content") or "")) except Exception: pass + # P0-4: Force-fetch last log snippet when job reaches terminal state + if job.get("status") in {"failed", "stopped"}: + try: + last_logs = await client.job_logs(task["compute_job_id"], tail_lines=200) + job["log_snippet"] = str(last_logs.get("content") or "")[:8192] + except Exception: + pass synced.append(store.apply_compute_job(task["id"], job)) except Exception as exc: # noqa: BLE001 - keep polling other jobs failed.append({"task_id": task["id"], "error": str(exc)}) diff --git a/backend/app/modules/data_process/dataset_format.py b/backend/app/modules/data_process/dataset_format.py new file mode 100644 index 0000000..fb0b68b --- /dev/null +++ b/backend/app/modules/data_process/dataset_format.py @@ -0,0 +1,148 @@ +"""Dataset format validation for Alpaca, ShareGPT, DPO, CPT formats. + +Used by the training preflight flow to validate that uploaded dataset files +conform to the declared format before submitting to the compute node. +""" + +from __future__ import annotations + +import json +from typing import Any + + +def _load_sample(path: str | None, content: str | None = None, max_samples: int = 20) -> list[dict[str, Any]]: + """Load up to max_samples records from JSONL file path or raw content string.""" + try: + if content is not None: + text = content.strip() + elif path: + with open(path, "r", encoding="utf-8") as fh: + text = fh.read().strip() + else: + return [] + except Exception: + return [] + + if not text: + return [] + + lines = text.splitlines()[:max_samples] + records: list[dict[str, Any]] = [] + for line in lines: + line = line.strip() + if not line: + continue + try: + record = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(record, dict): + records.append(record) + return records + + +def _check_alpaca(records: list[dict[str, Any]]) -> list[str]: + """Validate Alpaca format: requires 'instruction' field.""" + errors: list[str] = [] + if not records: + errors.append("Alpaca 格式数据集无有效记录") + return errors + missing_instruction = sum(1 for r in records if not r.get("instruction")) + if missing_instruction: + errors.append( + f"Alpaca 格式要求每条记录包含 instruction 字段," + f"前{len(records)}条中有{missing_instruction}条缺失" + ) + return errors + + +def _check_sharegpt(records: list[dict[str, Any]]) -> list[str]: + """Validate ShareGPT format: requires 'messages' (list of dicts with role/content).""" + errors: list[str] = [] + if not records: + errors.append("ShareGPT 格式数据集无有效记录") + return errors + bad = 0 + for r in records: + messages = r.get("messages") + if not isinstance(messages, list) or not messages: + bad += 1 + continue + for msg in messages: + if not isinstance(msg, dict) or "role" not in msg or "content" not in msg: + bad += 1 + break + if bad: + errors.append( + f"ShareGPT 格式要求每条记录包含 messages 列表," + f"每条消息需有 role 和 content 字段,前{len(records)}条中有{bad}条不符合" + ) + return errors + + +def _check_dpo(records: list[dict[str, Any]]) -> list[str]: + """Validate DPO format: requires 'chosen' and 'rejected' fields.""" + errors: list[str] = [] + if not records: + errors.append("DPO 格式数据集无有效记录") + return errors + missing_chosen = sum(1 for r in records if not r.get("chosen")) + missing_rejected = sum(1 for r in records if not r.get("rejected")) + if missing_chosen: + errors.append(f"DPO 格式要求 chosen 字段,前{len(records)}条中有{missing_chosen}条缺失") + if missing_rejected: + errors.append(f"DPO 格式要求 rejected 字段,前{len(records)}条中有{missing_rejected}条缺失") + return errors + + +def _check_cpt(records: list[dict[str, Any]]) -> list[str]: + """Validate CPT format: requires 'text' field, should NOT have instruction/output.""" + errors: list[str] = [] + if not records: + errors.append("CPT 格式数据集无有效记录") + return errors + missing_text = sum(1 for r in records if not r.get("text")) + has_instruction = sum(1 for r in records if r.get("instruction") or r.get("output")) + if missing_text: + errors.append(f"CPT 格式要求 text 字段,前{len(records)}条中有{missing_text}条缺失") + if has_instruction: + errors.append( + f"CPT 格式不应包含 instruction/output 字段(疑似 Alpaca 格式)," + f"前{len(records)}条中有{has_instruction}条包含此类字段" + ) + return errors + + +FORMAT_VALIDATORS = { + "alpaca": _check_alpaca, + "alpaca_jsonl": _check_alpaca, + "sharegpt": _check_sharegpt, + "dpo": _check_dpo, + "cpt": _check_cpt, + "pt": _check_cpt, +} + + +def validate_dataset_format( + dataset_format: str, + content: str | None = None, + path: str | None = None, + max_samples: int = 20, +) -> list[str]: + """Validate dataset content against expected format. + + Args: + dataset_format: One of 'alpaca', 'sharegpt', 'dpo', 'cpt'. + content: Raw file content (JSONL text). Mutually exclusive with path. + path: File path to read content from. + max_samples: Maximum records to sample for validation. + + Returns: + List of error messages (empty if valid). + """ + fmt = str(dataset_format).lower().strip() + validator = FORMAT_VALIDATORS.get(fmt) + if not validator: + return [f"不支持的数据集格式: {dataset_format},支持的格式: {', '.join(sorted(FORMAT_VALIDATORS))}"] + records = _load_sample(path=path, content=content, max_samples=max_samples) + return validator(records)