diff --git a/backend/app/api/v1/endpoints/health.py b/backend/app/api/v1/endpoints/health.py index c5676c1..8a01289 100644 --- a/backend/app/api/v1/endpoints/health.py +++ b/backend/app/api/v1/endpoints/health.py @@ -1,7 +1,6 @@ from fastapi import APIRouter from app.core.logging import get_logger -from app.db.platform_store import get_platform_store router = APIRouter() logger = get_logger(__name__) @@ -10,5 +9,9 @@ logger = get_logger(__name__) @router.get("/health") async def health_check() -> dict[str, object]: logger.info("health check requested") - return {"code": 0, "message": "ok", "data": get_platform_store().health_metrics()} + return { + "code": 0, + "message": "ok", + "data": {"cpu_percent": 0.0, "memory_percent": 0.0, "disk_percent": 0.0}, + } diff --git a/backend/app/api/v1/endpoints/platform.py b/backend/app/api/v1/endpoints/platform.py index 88ae3fa..c676e88 100644 --- a/backend/app/api/v1/endpoints/platform.py +++ b/backend/app/api/v1/endpoints/platform.py @@ -15,7 +15,7 @@ from app.core.auth import filter_accessible_resource_ids, get_current_user, has_ from app.core.config import get_settings from app.db.platform_store import get_platform_store from app.modules.compute_gateway.client import ComputeNodeClient -from app.modules.compute_gateway.sync import poll_compute_jobs_once +from app.modules.compute_gateway.sync import fetch_eval_result_content, poll_compute_jobs_once router = APIRouter() @@ -33,6 +33,62 @@ def _select_first_online_node(store: Any) -> dict[str, Any] | None: return None +def _build_messages_payload(payload: dict[str, Any]) -> dict[str, Any]: + """Convert frontend inference payload to compute API messages format. + + Accepts both: + - OpenAI-style: {messages: [{role, content}, ...], temperature, ...} + - Frontend-style: {user_question, system_prompt, temperature, ...} + """ + if payload.get("messages"): + messages = payload["messages"] + # messages already in OpenAI format; pass through with optional system prompt + if payload.get("system_prompt") and not any(m.get("role") == "system" for m in messages): + messages = [{"role": "system", "content": payload["system_prompt"]}] + list(messages) + else: + messages = [] + if payload.get("system_prompt"): + messages.append({"role": "system", "content": payload["system_prompt"]}) + question = payload.get("user_question") or payload.get("question") or "" + if question: + messages.append({"role": "user", "content": question}) + return { + "messages": messages, + "temperature": float(payload.get("temperature", 0.7)), + "top_p": float(payload.get("top_p", 0.95)), + "max_new_tokens": int(payload.get("max_tokens", 2048)), + "do_sample": bool(payload.get("do_sample", True)), + } + + +async def _stream_chat_proxy(payload: dict[str, Any]) -> StreamingResponse: + """Common SSE streaming proxy: convert payload → forward to compute node → stream back.""" + store = get_platform_store() + node = _select_first_online_node(store) + if not node: + return StreamingResponse( + iter(['data: {"error": "no online compute node available for inference"}\n\n']), + media_type="text/event-stream", + ) + client = ComputeNodeClient(node["api_base_url"]) + compute_payload = _build_messages_payload(payload) + + async def stream_proxy(): + async with httpx.AsyncClient(timeout=300) as http: + url = f"{node['api_base_url'].rstrip('/')}{client.route_prefix}/inference/chat/stream" + try: + async with http.stream("POST", url, json=compute_payload, headers=client.headers()) as resp: + if resp.status_code >= 400: + yield f'data: {{"error": "compute node returned {resp.status_code}"}}\n\n'.encode() + return + async for chunk in resp.aiter_bytes(): + yield chunk + except Exception as exc: + yield f'data: {{"error": "stream proxy failed: {exc}"}}\n\n'.encode() + + return StreamingResponse(stream_proxy(), media_type="text/event-stream") + + def fail(status_code: int, message: str) -> HTTPException: return HTTPException(status_code=status_code, detail={"code": status_code, "message": message, "data": None}) @@ -856,6 +912,7 @@ async def upload_dataset_files( ) -> dict[str, Any]: created: list[dict[str, Any]] = [] compute_sync: list[dict[str, Any]] = [] + pending_sync: list[tuple[str, str, bytes]] = [] store = get_platform_store() try: store.dataset(dataset_id) @@ -867,16 +924,18 @@ async def upload_dataset_files( content = raw.decode("utf-8", errors="replace") created_file = store.add_dataset_file(conn, dataset_id, file.filename or "upload.jsonl", content) created.append(created_file) - if sync_to_compute: - compute_sync.extend( - await _sync_dataset_file_to_compute_nodes( - store, - dataset_id, - created_file["id"], - created_file["name"], - raw, - ) + pending_sync.append((created_file["id"], created_file["name"], raw)) + if sync_to_compute: + for file_id, file_name, raw in pending_sync: + compute_sync.extend( + await _sync_dataset_file_to_compute_nodes( + store, + dataset_id, + file_id, + file_name, + raw, ) + ) return ok({"files": created, "compute_sync": compute_sync}) @@ -1212,7 +1271,23 @@ async def model_eval_list(current_user: dict = Depends(get_current_user)) -> dic @router.get("/model-eval/{task_id}") async def model_eval_detail(task_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]: try: - task = get_platform_store().eval_task(task_id) + store = get_platform_store() + task = store.eval_task(task_id) + if task.get("compute_job_id") and task.get("compute_node_id") and task.get("status") in {"queued", "running", "completed"}: + node = next( + (n for n in store.compute_nodes() if n["id"] == task.get("compute_node_id")), + None, + ) + if node: + try: + client = ComputeNodeClient(node["api_base_url"]) + job = await client.get_job(task["compute_job_id"]) + result_content = None + if job.get("status") == "completed" and not task.get("samples"): + result_content = await fetch_eval_result_content(client, node, job) + task = store.apply_eval_job_result(task_id, job, result_content) + except Exception: + pass except KeyError: raise fail(404, "eval task not found") if not has_resource_access("eval", task_id, current_user, "read"): @@ -1222,8 +1297,149 @@ async def model_eval_detail(task_id: str, current_user: dict = Depends(get_curre @router.post("/model-eval/start") async def model_eval_start(payload: dict[str, Any] = Body(...)) -> dict[str, Any]: - task = get_platform_store().create_eval_task(payload) - return ok({"task_id": task["id"], **task}) + """Start an evaluation task: submit eval job to compute node.""" + store = get_platform_store() + # 1. Create eval task record + task = store.create_eval_task({**payload, "status": "pending"}) + + # 2. Resolve model path (supports both regular models and trained models) + model_id = str(payload.get("model_id", "")) + model_path = "" + adapter_path = payload.get("adapter_path", "") + try: + db_model = store.model(model_id) + model_path = db_model.get("path", "") + except KeyError: + # Try trained_models table (IDs prefixed with tm_) + trained = next((m for m in store.trained_models() if m["id"] == model_id), None) + if trained: + merged_path = trained.get("merged_path", "") + base_path = trained.get("base_model_path", "") + if trained.get("merged") and merged_path: + # Merged model: use merged_path as model, no adapter needed + model_path = merged_path + elif base_path: + # Unmerged: use base model + adapter checkpoint + model_path = base_path + if merged_path: + adapter_path = merged_path + else: + model_path = merged_path or base_path + if not model_path: + store.update_eval_task(task["id"], {"status": "failed", "error": "model not found or no path"}) + return ok({"task_id": task["id"], "status": "failed", "error": "model not found or no path"}) + + # 3. Resolve dataset file + dataset_id = str(payload.get("dataset_id", "")) + dataset_path = "" + try: + ds_files = store.training_dataset_files(dataset_id) + if ds_files: + dataset_path = ds_files[0].get("local_path") or ds_files[0].get("name", "") + except Exception: + pass + if not dataset_path: + # Try to get file content and sync to compute + try: + ds = store.dataset(dataset_id) + for f in ds.get("files", []): + if f.get("content"): + dataset_path = f.get("name", f"dataset_{dataset_id}.jsonl") + break + except KeyError: + pass + if not dataset_path: + store.update_eval_task(task["id"], {"status": "failed", "error": "dataset not found or no files"}) + return ok({"task_id": task["id"], "status": "failed", "error": "dataset not found or no files"}) + + # 4. Resolve dimension config + dimension_id = str(payload.get("dimension_id", "")) + dimension_cfg: dict[str, Any] = {} + if dimension_id: + try: + dim = store.dimension(dimension_id) + # Resolve eval model API config + eval_model_name = dim.get("eval_model", "") + api_url = "" + api_key = "" + if eval_model_name: + try: + eval_model = store.model(eval_model_name) if eval_model_name.startswith("m_") else store.model_by_name(eval_model_name) + api_url = eval_model.get("api_url", "") + api_key = eval_model.get("api_key", "") + except (KeyError, Exception): + pass + dimension_cfg = { + "type": dim.get("type", ""), + "eval_model": eval_model_name, + "eval_method": dim.get("eval_method", ""), + "eval_prompt": dim.get("eval_prompt", ""), + "api_url": api_url, + "api_key": api_key, + "score_min": dim.get("score_min", 0), + "score_max": dim.get("score_max", 5), + "pass_threshold": dim.get("pass_threshold", 3), + } + except KeyError: + pass + + # 5. Select compute node + node = _select_first_online_node(store) + if not node: + store.update_eval_task(task["id"], {"status": "failed", "error": "no online compute node"}) + return ok({"task_id": task["id"], "status": "failed", "error": "no online compute node"}) + + # 6. Build eval job payload + output_dir = f"/data/yg-ft/outputs/{task['id']}" + job_payload = { + "id": f"eval_{task['id']}", + "name": task.get("eval_task_name", task["id"]), + "engine": "eval", + "model_name_or_path": model_path, + "adapter_name_or_path": adapter_path, + "template": payload.get("template", "qwen"), + "dataset_path": dataset_path, + "output_dir": output_dir, + "basic_metrics": payload.get("basic_metrics", {}), + "dimension": dimension_cfg, + "gpus": [int(payload.get("gpu_id", 0))], + "temperature": payload.get("temperature", 0.1), + "max_new_tokens": payload.get("max_new_tokens", 512), + "compute_node_id": node["id"], + } + + # 7. Submit to compute node via create_job (uses engine="eval" path) + try: + client = ComputeNodeClient(node["api_base_url"]) + # Sync dataset file to compute node if needed + if not dataset_path.startswith("/"): + try: + ds_files = store.training_dataset_files(dataset_id) + if ds_files and ds_files[0].get("content"): + upload_result = await client.upload_file( + ds_files[0].get("name", "eval_data.jsonl"), + ds_files[0]["content"].encode("utf-8"), + f"datasets/{dataset_id}/{ds_files[0].get('name', 'eval_data.jsonl')}", + resource_type="dataset", + resource_id=dataset_id, + ) + job_payload["dataset_path"] = upload_result.get("local_path", dataset_path) + except Exception: + pass + + job = await client.create_job(job_payload) + store.update_eval_task(task["id"], { + "status": "running", + "compute_job_id": job.get("id"), + "compute_node_id": node["id"], + "output_dir": output_dir, + }) + if job.get("status") in {"queued", "running"}: + store.mark_inference_loaded(node["id"]) + return ok({"task_id": task["id"], "status": "running", "job": job}) + except Exception as exc: + store.update_eval_task(task["id"], {"status": "failed", "error": str(exc)}) + return ok({"task_id": task["id"], "status": "failed", "error": str(exc)}) @router.delete("/model-eval/{task_id}") @@ -1298,8 +1514,27 @@ async def model_compare_detail(task_id: str) -> dict[str, Any]: raise fail(404, "compare task not found") +async def _unload_from_compute_node() -> dict[str, Any]: + """Best-effort unload the inference model from the first online compute node.""" + store = get_platform_store() + # Clear all inference tracking — only one model can be loaded at a time + for node in store.compute_nodes(): + store.mark_inference_unloaded(node["id"]) + node = _select_first_online_node(store) + if not node: + return {"unloaded": False, "error": "no online compute node"} + try: + client = ComputeNodeClient(node["api_base_url"]) + result = await client._request("POST", "/inference/unload", json_data={}) + return result + except Exception as exc: + return {"unloaded": False, "error": str(exc)} + + @router.delete("/model-compare/{task_id}") async def model_compare_delete(task_id: str) -> dict[str, Any]: + # 删除前先释放算力节点上的模型 + await _unload_from_compute_node() get_platform_store().delete_compare_task(task_id) return ok({"deleted": task_id}) @@ -1329,26 +1564,56 @@ async def model_compare_update_load_status(task_id: str, payload: dict[str, Any] @router.post("/model-compare/{task_id}/load") async def model_compare_load(task_id: str) -> dict[str, Any]: + """真正加载模型到算力节点(不再使用假 PID/端口)。""" try: - task = get_platform_store().compare_task(task_id) + store = get_platform_store() + task = store.compare_task(task_id) models = task.get("models") or [] if isinstance(models, str): try: models = json.loads(models) except json.JSONDecodeError: models = [] - loaded_models = [ - { - "model_id": item.get("model_id"), - "model_name": item.get("model_name"), - "status": "ready", - "pid": 45000 + index, - "port": item.get("port") or 18000 + index, + # 选取在线算力节点 + node = _select_first_online_node(store) + if not node: + return ok({"status": "failed", "error": "no online compute node"}) + client = ComputeNodeClient(node["api_base_url"]) + loaded_models = [] + for item in models: + if not isinstance(item, dict): + continue + model_path = item.get("model_path", "") + if not model_path: + # 尝试从模型库获取路径 + model_id = item.get("model_id", "") + try: + db_model = store.model(model_id) + model_path = db_model.get("path", "") + except KeyError: + pass + if not model_path: + loaded_models.append({**item, "status": "error", "error": "model_path not found"}) + continue + # 真正调用算力节点加载模型 + load_payload = { + "model_name_or_path": model_path, + "template": item.get("template", "qwen"), } - for index, item in enumerate(models) - if isinstance(item, dict) - ] - return ok(get_platform_store().update_compare_task(task_id, {"status": "loaded", "load_status": {"loaded_models": loaded_models}})) + if item.get("adapter_path"): + load_payload["adapter_name_or_path"] = item["adapter_path"] + try: + result = await client._request("POST", "/inference/load", json_data=load_payload) + if result.get("loaded"): + store.mark_inference_loaded(node["id"]) + loaded_models.append({**item, "status": "ready"}) + else: + loaded_models.append({**item, "status": "error", "error": result.get("error", "load failed")}) + except Exception as exc: + loaded_models.append({**item, "status": "error", "error": str(exc)}) + status = "loaded" if any(m.get("status") == "ready" for m in loaded_models) else "failed" + updated = store.update_compare_task(task_id, {"status": status, "load_status": {"loaded_models": loaded_models}}) + return ok(updated) except KeyError: raise fail(404, "compare task not found") @@ -1356,7 +1621,11 @@ async def model_compare_load(task_id: str) -> dict[str, Any]: @router.post("/model-compare/{task_id}/unload") async def model_compare_unload(task_id: str) -> dict[str, Any]: try: - return ok(get_platform_store().update_compare_task(task_id, {"status": "pending", "load_status": {"loaded_models": []}})) + store = get_platform_store() + # 真正释放算力节点上的模型资源 + unload_result = await _unload_from_compute_node() + updated = store.update_compare_task(task_id, {"status": "pending", "load_status": {"loaded_models": []}}) + return ok({"task": updated, "unload": unload_result}) except KeyError: raise fail(404, "compare task not found") @@ -1368,18 +1637,23 @@ async def model_compare_start_model(task_id: str, payload: dict[str, Any] = Body @router.post("/model-compare/chat-with-port") async def model_compare_chat_with_port(payload: dict[str, Any] = Body(...)) -> dict[str, Any]: - question = "" - for message in payload.get("messages") or []: - if message.get("role") == "user": - question = str(message.get("content") or "") - content = f"当前后端已收到推理请求:{question[:120]}" - return ok({"response": content, "content": content}) + """Proxy non-streaming chat to the compute node running the inference model.""" + store = get_platform_store() + node = _select_first_online_node(store) + if not node: + return ok({"response": "no online compute node available for inference", "request": payload}) + try: + client = ComputeNodeClient(node["api_base_url"]) + result = await client._request("POST", "/inference/chat", json_data=_build_messages_payload(payload)) + return ok(result) + except Exception as exc: + return ok({"response": f"inference failed: {exc}", "request": payload}) @router.post("/model-compare/stream-chat") -async def model_compare_stream_chat(payload: dict[str, Any] = Body(...)) -> dict[str, Any]: - question = payload.get("user_question") or payload.get("question") or "" - return ok({"response": f"当前后端已收到流式推理请求:{str(question)[:120]}"}) +async def model_compare_stream_chat(payload: dict[str, Any] = Body(...)) -> StreamingResponse: + """Stream chat from the compute node (SSE proxy).""" + return await _stream_chat_proxy(payload) @router.post("/model-chat/batch") @@ -1396,7 +1670,7 @@ async def model_chat_local(payload: dict[str, Any] = Body(...)) -> dict[str, Any return ok({"response": "no online compute node available for inference", "request": payload}) try: client = ComputeNodeClient(node["api_base_url"]) - result = await client._request("POST", "/inference/chat", json_data=payload) + result = await client._request("POST", "/inference/chat", json_data=_build_messages_payload(payload)) return ok(result) except Exception as exc: return ok({"response": f"inference failed: {exc}", "request": payload}) @@ -1405,28 +1679,15 @@ async def model_chat_local(payload: dict[str, Any] = Body(...)) -> dict[str, Any @router.post("/model-chat/local/chat/stream") async def model_chat_local_stream(payload: dict[str, Any] = Body(...)) -> StreamingResponse: """Stream chat from the compute node.""" - store = get_platform_store() - node = _select_first_online_node(store) - if not node: - return StreamingResponse( - iter(['data: {"error": "no online compute node"}\n\n']), - media_type="text/event-stream", - ) - client = ComputeNodeClient(node["api_base_url"]) - - async def stream_proxy(): - async with httpx.AsyncClient(timeout=300) as http: - url = f"{node['api_base_url'].rstrip('/')}/modelTF/inference/chat/stream" - async with http.stream("POST", url, json=payload, headers=client.headers()) as resp: - async for chunk in resp.aiter_bytes(): - yield chunk - - return StreamingResponse(stream_proxy(), media_type="text/event-stream") + return await _stream_chat_proxy(payload) @router.post("/model-chat/local/preload") async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[str, Any]: """Load a model on the compute node for inference.""" + model_path = (payload.get("model_name_or_path") or "").strip() + if not model_path: + return ok({"loaded": False, "error": "model_name_or_path is required"}) store = get_platform_store() node = _select_first_online_node(store) if not node: @@ -1434,6 +1695,8 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[ try: client = ComputeNodeClient(node["api_base_url"]) result = await client._request("POST", "/inference/load", json_data=payload) + if result.get("loaded"): + store.mark_inference_loaded(node["id"]) return ok(result) except Exception as exc: return ok({"loaded": False, "error": str(exc)}) @@ -1443,6 +1706,9 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[ async def model_chat_local_unload() -> dict[str, Any]: """Unload the inference model from the compute node.""" store = get_platform_store() + # Clear all inference tracking + for n in store.compute_nodes(): + store.mark_inference_unloaded(n["id"]) node = _select_first_online_node(store) if not node: return ok({"unloaded": False, "error": "no online compute node"}) @@ -1472,6 +1738,9 @@ async def model_chat_local_status() -> dict[str, Any]: @router.post("/model-chat/trained/preload") async def model_chat_trained_preload(payload: dict[str, Any] = Body(...)) -> dict[str, Any]: """Load a trained model (base + adapter) on the compute node for inference.""" + model_path = (payload.get("model_name_or_path") or "").strip() + if not model_path: + return ok({"loaded": False, "error": "model_name_or_path is required"}) store = get_platform_store() node = _select_first_online_node(store) if not node: @@ -1479,6 +1748,8 @@ async def model_chat_trained_preload(payload: dict[str, Any] = Body(...)) -> dic try: client = ComputeNodeClient(node["api_base_url"]) result = await client._request("POST", "/inference/load", json_data=payload) + if result.get("loaded"): + store.mark_inference_loaded(node["id"]) return ok(result) except Exception as exc: return ok({"loaded": False, "error": str(exc)}) @@ -1517,6 +1788,16 @@ async def update_compute_node(node_id: str, payload: dict[str, Any] = Body(...)) raise fail(400, str(exc)) +@router.delete("/compute/nodes/{node_id}") +async def delete_compute_node(node_id: str) -> dict[str, Any]: + try: + return ok(get_platform_store().delete_compute_node(node_id)) + except KeyError: + raise fail(404, "compute node not found") + except ValueError as exc: + raise fail(400, str(exc)) + + @router.post("/compute/nodes/{node_id}/test-connection") async def test_compute_node(node_id: str) -> dict[str, Any]: store = get_platform_store() diff --git a/backend/app/db/platform_store.py b/backend/app/db/platform_store.py index a9f67ca..7cad6af 100644 --- a/backend/app/db/platform_store.py +++ b/backend/app/db/platform_store.py @@ -88,6 +88,22 @@ def parse_size_bytes(value: Any) -> int: return max(0, round(amount * _SIZE_UNIT_BYTES[unit])) +def count_dataset_records(content: str) -> int: + text = (content or "").strip() + if not text: + return 0 + + try: + value = json.loads(text) + if isinstance(value, list): + return len(value) + return 1 + except (TypeError, ValueError, json.JSONDecodeError): + pass + + return len([line for line in text.splitlines() if line.strip()]) + + def version_number(value: Any, default: int = 0) -> int: try: number = int(value) @@ -296,6 +312,7 @@ class PlatformStore: # request (notably expensive against the remote PostgreSQL instance). # TCP keepalive 让操作系统持续保活连接,抵抗远程库空闲静默断连。 pool_kwargs = { + "connect_timeout": 5, "keepalives": 1, "keepalives_idle": 10, "keepalives_interval": 5, @@ -320,6 +337,19 @@ class PlatformStore: self._pool.open() self.ensure_schema() self.ensure_seed_data() + # Track which compute nodes have an active inference model loaded + self._inference_nodes: set[str] = set() + + # ── inference node tracking ──────────────────────────────────── + + def mark_inference_loaded(self, node_id: str) -> None: + self._inference_nodes.add(node_id) + + def mark_inference_unloaded(self, node_id: str) -> None: + self._inference_nodes.discard(node_id) + + def is_inference_loaded(self, node_id: str) -> bool: + return node_id in self._inference_nodes @contextmanager def connect(self) -> Iterator["PgConnection"]: @@ -1298,6 +1328,7 @@ class PlatformStore: file_size_bytes = int(file_row.get("size_bytes") or 0) if file_size_bytes <= 0: file_size_bytes = parse_size_bytes(file_row.get("size")) + file_record_count = int(file_row.get("record_count") or 0) decoded_files.append( { "id": file_row["id"], @@ -1306,7 +1337,7 @@ class PlatformStore: "size_bytes": file_size_bytes, **dataset_file_version_summary(file_row), "create_time": file_row["create_time"], - "record_count": int(file_row.get("record_count") or 0), + "record_count": file_record_count, "split": metadata.get("file_split"), } ) @@ -1324,6 +1355,9 @@ class PlatformStore: total_size_bytes = int(row.get("size_bytes") or 0) if total_size_bytes <= 0: total_size_bytes = parse_size_bytes(row.get("size")) + total_record_count = sum(int(item.get("record_count") or 0) for item in decoded_files) + if not decoded_files: + total_record_count = int(row.get("record_count") or row.get("count") or 0) current_version_nos = sorted( { int(item["current_version_no"]) @@ -1333,6 +1367,8 @@ class PlatformStore: ) return { **dict(row), + "count": total_record_count, + "record_count": total_record_count, "size_bytes": total_size_bytes, "current_version_no": ( current_version_nos[0] if len(current_version_nos) == 1 else None @@ -1407,7 +1443,7 @@ class PlatformStore: version_id = f"{file_id}_v1" size_bytes = len(content.encode("utf-8")) size = f"{size_bytes} B" - record_count = len([line for line in content.splitlines() if line.strip()]) + record_count = count_dataset_records(content) version = { "id": version_id, "version": 1, @@ -1440,10 +1476,18 @@ class PlatformStore: ) conn.execute( """UPDATE datasets - SET count=count+?, record_count=record_count+?, - size_bytes=size_bytes+?, size=((size_bytes+?)::text || ' B') + SET count=stats.record_count, + record_count=stats.record_count, + size_bytes=stats.size_bytes, + size=(stats.size_bytes::text || ' B') + FROM ( + SELECT COALESCE(SUM(record_count), 0) AS record_count, + COALESCE(SUM(size_bytes), 0) AS size_bytes + FROM dataset_files + WHERE dataset_id=? + ) stats WHERE id=?""", - (record_count, record_count, size_bytes, size_bytes, dataset_id), + (dataset_id, dataset_id), ) return { "id": file_id, @@ -1548,19 +1592,57 @@ class PlatformStore: row = conn.execute("SELECT * FROM dataset_files WHERE id=?", (file_id,)).fetchone() if not row: raise KeyError(file_id) + content = payload.get("content", "") + size_bytes = len(content.encode("utf-8")) + record_count = count_dataset_records(content) versions = json_loads(row["versions"], []) version = { "id": f"{file_id}_v{len(versions) + 1}", "version": len(versions) + 1, + "version_no": len(versions) + 1, "create_time": utcnow(), "description": payload.get("description", "online edit"), + "size_bytes": size_bytes, + "record_count": record_count, } versions.append(version) conn.execute( - "UPDATE dataset_files SET content=?, active_version_id=?, versions=? WHERE id=?", - (payload.get("content", ""), version["id"], json_dumps(versions), file_id), + """ + UPDATE dataset_files + SET content=?, active_version_id=?, current_version_id=?, versions=?, + size_bytes=?, size=?, record_count=?, version_no=? + WHERE id=? + """, + ( + content, + version["id"], + version["id"], + json_dumps(versions), + size_bytes, + f"{size_bytes} B", + record_count, + version["version_no"], + file_id, + ), ) - return {"version": version, "content": payload.get("content", "")} + conn.execute( + """UPDATE datasets + SET count=stats.record_count, + record_count=stats.record_count, + size_bytes=stats.size_bytes, + size=(stats.size_bytes::text || ' B') + FROM ( + SELECT dataset_id, + COALESCE(SUM(record_count), 0) AS record_count, + COALESCE(SUM(size_bytes), 0) AS size_bytes + FROM dataset_files + WHERE dataset_id=(SELECT dataset_id FROM dataset_files WHERE id=?) + GROUP BY dataset_id + ) stats + WHERE datasets.id=stats.dataset_id""", + (file_id,), + ) + return {"version": version, "content": content} def activate_file_version(self, file_id: str, version_id: str) -> dict[str, Any]: with self.connect() as conn: @@ -2114,10 +2196,63 @@ class PlatformStore: ) return self.eval_task(task_id) + def update_eval_task(self, task_id: str, updates: dict[str, Any]) -> dict[str, Any]: + """Update fields in an eval task's payload without replacing the whole record.""" + task = self.eval_task(task_id) + merged = {**task, **updates} + with self.connect() as conn: + conn.execute( + "UPDATE eval_tasks SET payload=?, status=? WHERE id=?", + (json_dumps(merged), merged.get("status", task.get("status", "pending")), task_id), + ) + return self.eval_task(task_id) + def delete_eval_task(self, task_id: str) -> None: with self.connect() as conn: conn.execute("DELETE FROM eval_tasks WHERE id=?", (task_id,)) + def running_eval_tasks(self) -> list[dict[str, Any]]: + """Return eval tasks that have been submitted to a compute node and are still running.""" + return [ + task for task in self.eval_tasks() + if task.get("compute_job_id") and task.get("status") in {"queued", "running"} + ] + + def apply_eval_job_result(self, task_id: str, job: dict[str, Any], result_content: dict[str, Any] | None = None) -> dict[str, Any]: + """Sync a compute job status/result back to an eval task.""" + task = self.eval_task(task_id) + job_status = str(job.get("status", "")) + status_map = {"queued": "running", "running": "running", "completed": "completed", + "failed": "failed", "stopped": "stopped"} + new_status = status_map.get(job_status, job_status or task.get("status", "pending")) + updates: dict[str, Any] = { + "status": new_status, + "progress": int(job.get("progress", 0)), + "output_dir": job.get("output_dir", task.get("output_dir", "")), + } + # On completion, populate results from eval_results.json content + if new_status == "completed" and result_content: + updates.update({ + "overall_score": result_content.get("overall_score", 0), + "overall_score_max": result_content.get("overall_score_max", 100), + "overall_evaluation": result_content.get("overall_evaluation", ""), + "improvement_suggestions": result_content.get("improvement_suggestions", []), + "dimension_summary": result_content.get("dimension_summary", []), + "samples": result_content.get("samples", []), + "sample_count": result_content.get("sample_count", 0), + "completed_count": result_content.get("completed_count", 0), + "passed_count": result_content.get("passed_count", 0), + "basic_metrics": result_content.get("basic_metrics", {}), + "score": result_content.get("overall_score", 0), + "completed_time": utcnow(), + }) + elif new_status in {"failed", "stopped"}: + updates.update({ + "error": job.get("error") or task.get("error") or "", + "completed_time": utcnow(), + }) + return self.update_eval_task(task_id, updates) + def dimensions(self) -> list[dict[str, Any]]: with self.connect() as conn: rows = conn.execute("SELECT * FROM eval_dimensions ORDER BY create_time DESC").fetchall() @@ -2284,6 +2419,15 @@ class PlatformStore: ).fetchall() return {int(row["gpu_index"]) for row in rows} + def _node_gpu_indexes(self, conn: PgConnection, node: dict[str, Any]) -> set[int]: + rows = conn.execute("SELECT gpu_index FROM gpus WHERE node_id=?", (node["id"],)).fetchall() + if rows: + return {int(row["gpu_index"]) for row in rows} + return set(range(max(0, int(node.get("gpu_count") or 0)))) + + def _node_capacity(self, node: dict[str, Any]) -> int: + return max(1, int(node.get("max_parallel_jobs") or 1), int(node.get("gpu_count") or 0)) + def _schedule_node_locked(self, conn: PgConnection, payload: dict[str, Any]) -> dict[str, Any]: requested = payload.get("requested_node_id") or payload.get("compute_node_id") requested_gpus = [int(item) for item in payload.get("gpus") or []] @@ -2291,13 +2435,15 @@ class PlatformStore: candidates = [ n for n in nodes - if n["enabled"] and n["scheduler_status"] == "online" and n["current_running_jobs"] < n["max_parallel_jobs"] + if n["enabled"] and n["scheduler_status"] == "online" and n["current_running_jobs"] < self._node_capacity(n) ] if requested_gpus: + requested_gpu_set = set(requested_gpus) candidates = [ node for node in candidates - if not set(requested_gpus).intersection(self._active_gpu_indexes(conn, node["id"])) + if requested_gpu_set.issubset(self._node_gpu_indexes(conn, node)) + and not requested_gpu_set.intersection(self._active_gpu_indexes(conn, node["id"])) ] if requested: selected = next((n for n in candidates if n["id"] == requested), None) @@ -2312,8 +2458,8 @@ class PlatformStore: reason = "disabled" elif node["scheduler_status"] != "online": reason = f"status={node['scheduler_status']}" - elif node["current_running_jobs"] >= node["max_parallel_jobs"]: - reason = f"capacity full {node['current_running_jobs']}/{node['max_parallel_jobs']}" + elif node["current_running_jobs"] >= self._node_capacity(node): + reason = f"capacity full {node['current_running_jobs']}/{self._node_capacity(node)}" else: reason = "not selected" reasons.append(f"{node['code']}({reason})") @@ -2664,6 +2810,24 @@ class PlatformStore: ) return next(node for node in self.compute_nodes() if node["id"] == node_id) + def delete_compute_node(self, node_id: str) -> dict[str, Any]: + with self.connect() as conn: + node = conn.execute("SELECT * FROM compute_nodes WHERE id=?", (node_id,)).fetchone() + if not node: + raise KeyError(node_id) + active = conn.execute( + """ + SELECT COUNT(*) AS cnt + FROM fine_tune_tasks + WHERE compute_node_id=? AND status IN ('syncing','queued','running') + """, + (node_id,), + ).fetchone() + if active and int(active["cnt"] or 0) > 0: + raise ValueError("compute node has active training tasks") + conn.execute("DELETE FROM compute_nodes WHERE id=?", (node_id,)) + return {"deleted": node_id} + def update_compute_node_health(self, node_id: str, health: dict[str, Any], success: bool, error: str | None = None) -> dict[str, Any]: current = next((n for n in self.compute_nodes() if n["id"] == node_id), None) if not current: @@ -2749,6 +2913,11 @@ class PlatformStore: ) busy = task is not None and task.get("status") == "running" reserved = task is not None and task.get("status") in {"syncing", "queued"} + # Also mark GPU as busy if an inference model is loaded on this node + inference_busy = self.is_inference_loaded(row["node_id"]) + if inference_busy and not busy: + busy = True + reserved = False memory_used = round(row["memory_total_gb"] * (0.72 if busy else 0.18 if reserved else 0.04), 1) gpu_percent = 86 if busy else 22 if reserved else 3 memory_total = float(row["memory_total_gb"] or 0) @@ -2827,11 +2996,12 @@ class PlatformStore: } def health_metrics(self) -> dict[str, float]: - info = self.system_info() + # Health checks must stay lightweight. The Docker healthcheck and page + # refresh probes should not wait on dashboard/GPU/database aggregation. return { - "cpu_percent": info["cpu"]["percent"], - "memory_percent": info["memory"]["percent"], - "disk_percent": info["disk"]["percent"], + "cpu_percent": 0.0, + "memory_percent": 0.0, + "disk_percent": 0.0, } def queue(self) -> list[dict[str, Any]]: diff --git a/backend/app/modules/compute_gateway/client.py b/backend/app/modules/compute_gateway/client.py index 8d1e821..aeeaae6 100644 --- a/backend/app/modules/compute_gateway/client.py +++ b/backend/app/modules/compute_gateway/client.py @@ -207,7 +207,8 @@ class ComputeNodeClient: "resource_id": resource_id or "", } files = {"file": (filename, content)} - async with httpx.AsyncClient(timeout=max(self.timeout, 60), headers=self.headers()) as client: + timeout = httpx.Timeout(max(self.timeout, 60), connect=self.timeout) + async with httpx.AsyncClient(timeout=timeout, headers=self.headers()) as client: response = await client.post( _join_url(self.api_base_url, f"{self.route_prefix}/compute/files/upload"), data=data, diff --git a/backend/app/modules/compute_gateway/sync.py b/backend/app/modules/compute_gateway/sync.py index 0e2720a..27f51e1 100644 --- a/backend/app/modules/compute_gateway/sync.py +++ b/backend/app/modules/compute_gateway/sync.py @@ -10,6 +10,24 @@ def _node_for_task(task: dict[str, Any]) -> dict[str, Any] | None: return next((node for node in get_platform_store().compute_nodes() if node["id"] == task.get("compute_node_id")), None) +async def fetch_eval_result_content(client: ComputeNodeClient, node: dict[str, Any], job: dict[str, Any]) -> dict[str, Any] | None: + output_dir = job.get("output_dir") + if not output_dir: + return None + full_path = f"{str(output_dir).rstrip('/')}/eval_results.json" + data_root = "/data/yg-ft/" + if full_path.startswith(data_root): + full_path = full_path[len(data_root):] + rel_path = full_path.lstrip("/") + import httpx + url = f"{node['api_base_url'].rstrip('/')}/modelTF/compute/files/read" + async with httpx.AsyncClient(timeout=30, headers=client.headers()) as http: + response = await http.get(url, params={"path": rel_path}) + response.raise_for_status() + payload = response.json() + return payload if isinstance(payload, dict) else None + + async def poll_compute_jobs_once() -> dict[str, Any]: store = get_platform_store() synced: list[dict[str, Any]] = [] @@ -48,4 +66,34 @@ async def poll_compute_jobs_once() -> dict[str, Any]: standalone_synced.append(store.sync_model_merge_job(record["id"], job)) except Exception as exc: # noqa: BLE001 - keep polling other jobs failed.append({"job_id": record["id"], "error": str(exc)}) - return {"synced": len(synced) + len(standalone_synced), "failed": failed, "items": synced, "standalone": standalone_synced} + + # ── Eval job sync ──────────────────────────────────────────────── + eval_synced = 0 + for eval_task in store.running_eval_tasks(): + node = next( + (item for item in store.compute_nodes() if item["id"] == eval_task.get("compute_node_id")), + None, + ) + if not node: + failed.append({"eval_task_id": eval_task["id"], "error": "compute node not found"}) + continue + try: + client = ComputeNodeClient(node["api_base_url"]) + job = await client.get_job(eval_task["compute_job_id"]) + result_content = None + # Try to read eval_results.json from the job output directory + if job.get("status") == "completed" and job.get("output_dir"): + try: + result_content = await fetch_eval_result_content(client, node, job) + except Exception: + pass + store.apply_eval_job_result(eval_task["id"], job, result_content) + # If job completed, un-mark inference loaded + if job.get("status") in {"completed", "failed", "stopped"}: + store.mark_inference_unloaded(node["id"]) + eval_synced += 1 + except Exception as exc: # noqa: BLE001 + failed.append({"eval_task_id": eval_task["id"], "error": str(exc)}) + + return {"synced": len(synced) + len(standalone_synced) + eval_synced, "failed": failed, + "items": synced, "standalone": standalone_synced, "eval_synced": eval_synced} diff --git a/backend/pyproject.toml b/backend/pyproject.toml index a1d4e4f..02cd526 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -10,6 +10,7 @@ dependencies = [ "pydantic>=2.7.0", "sqlalchemy>=2.0.30", "psycopg[binary]>=3.2.1", + "psycopg-pool>=3.2.1", "alembic>=1.13.1", "redis>=5.0.4", "httpx>=0.27.0", diff --git a/backend/requirements.txt b/backend/requirements.txt index 900f72b..04e346a 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -4,6 +4,7 @@ python-multipart>=0.0.9 pydantic>=2.7.0 sqlalchemy>=2.0.30 psycopg[binary]>=3.2.1 +psycopg-pool>=3.2.1 alembic>=1.13.1 redis>=5.0.4 httpx>=0.27.0 diff --git a/compute/api/main.py b/compute/api/main.py index f96dca1..d88d2fd 100644 --- a/compute/api/main.py +++ b/compute/api/main.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json import os import math import hashlib @@ -449,6 +450,29 @@ def create_app() -> FastAPI: accelerator_errors, accelerator_warnings, accelerator = _validate_training_accelerator(payload) errors.extend(accelerator_errors) warnings.extend(accelerator_warnings) + elif engine == "eval": + # Eval engine: validate model path and dataset path + if not payload.get("model_name_or_path"): + errors.append("model_name_or_path is required for eval") + else: + path_checks.append(_check_path_item({ + "name": "model_name_or_path", + "path": payload.get("model_name_or_path", ""), + "type": "any", + "required": True, + })) + if payload.get("dataset_path"): + path_checks.append(_check_path_item({ + "name": "dataset_path", + "path": payload.get("dataset_path", ""), + "type": "file", + "required": True, + })) + else: + errors.append("dataset_path is required for eval") + if shutil.which("python") is None: + errors.append("python runtime not found") + elif engine == "smoke": warnings.append("smoke engine skips model and dataset path checks") @@ -688,8 +712,6 @@ def create_app() -> FastAPI: infer_backend=payload.get("infer_backend", "huggingface"), infer_dtype=payload.get("infer_dtype", "auto"), ) - if not result.get("loaded"): - raise HTTPException(status_code=500, detail=result.get("error", "model load failed")) return result @app.post(f"{route_prefix}/inference/unload") @@ -813,6 +835,22 @@ def create_app() -> FastAPI: "checksum_sha256": checksum, } + @app.get(f"{route_prefix}/compute/files/read") + async def read_file(path: str = Query(...)) -> JSONResponse: + """Read a text file from within YG_FT_DATA_ROOT. Used by the backend + to fetch eval results and other job outputs.""" + data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft")) + target = (data_root / path.lstrip("/\\")).resolve() + if not _path_inside(data_root, target): + raise HTTPException(status_code=400, detail="path must stay inside YG_FT_DATA_ROOT") + if not target.is_file(): + raise HTTPException(status_code=404, detail="file not found") + try: + content = target.read_text(encoding="utf-8") + return JSONResponse(json.loads(content) if content.strip().startswith("{") else {"content": content}) + except Exception as exc: + raise HTTPException(status_code=500, detail=str(exc)) + @app.get(f"{route_prefix}/compute/files/{{file_id}}/download") async def download_file(file_id: str) -> FileResponse: upload_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft")) / "uploads" diff --git a/compute/engines/llama_factory/adapter.py b/compute/engines/llama_factory/adapter.py index 3a7042e..4c0b7b0 100644 --- a/compute/engines/llama_factory/adapter.py +++ b/compute/engines/llama_factory/adapter.py @@ -204,6 +204,31 @@ def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA- command.extend(["--quantization_bit", str(quantization_bit)]) return LlamaFactoryCommand(command=command, work_dir=str(Path(llama_factory_home)), env={}) + if engine == "eval": + output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'eval-job')}" + eval_config_path = str(Path(output_dir) / "eval_config.json") + eval_config = { + "model_name_or_path": config.get("model_name_or_path", ""), + "adapter_name_or_path": config.get("adapter_name_or_path", ""), + "template": config.get("template", "qwen"), + "dataset_path": config.get("dataset_path", ""), + "output_dir": output_dir, + "basic_metrics": config.get("basic_metrics", {}), + "dimension": config.get("dimension", {}), + "temperature": config.get("temperature", 0.1), + "top_p": config.get("top_p", 0.95), + "max_new_tokens": config.get("max_new_tokens", 512), + "infer_backend": config.get("infer_backend", "huggingface"), + "infer_dtype": config.get("infer_dtype", "auto"), + } + Path(output_dir).mkdir(parents=True, exist_ok=True) + Path(eval_config_path).write_text(json.dumps(eval_config, ensure_ascii=False, indent=2), encoding="utf-8") + return LlamaFactoryCommand( + command=["python", "-u", "-m", "compute.engines.llama_factory.eval_runner", "--config", eval_config_path], + work_dir="/app", + env={}, + ) + errors = validate_config(config) if errors: raise ValueError("; ".join(errors)) diff --git a/compute/engines/llama_factory/eval_runner.py b/compute/engines/llama_factory/eval_runner.py new file mode 100644 index 0000000..3c974a8 --- /dev/null +++ b/compute/engines/llama_factory/eval_runner.py @@ -0,0 +1,480 @@ +""" +Evaluation runner — executes model evaluation as a subprocess job. + +Usage: + python -m compute.engines.llama_factory.eval_runner --config + +The config JSON is written by the compute API before spawning this subprocess. +Results are written to ``output_dir/eval_results.json`` and progress is printed +to stdout (captured as job logs). +""" + +from __future__ import annotations + +import json +import math +import re +import sys +import time +from difflib import SequenceMatcher +from pathlib import Path +from typing import Any + + +def _load_dataset(path: str) -> list[dict[str, Any]]: + """Load a JSON or JSONL dataset file. + + Supports common field names used across the platform: + * ``instruction`` + ``input`` + ``output`` (Alpaca-style) + * ``question`` + ``answer`` + * ``messages`` (ShareGPT-style – the last assistant message is treated as reference) + """ + file_path = Path(path) + text = file_path.read_text(encoding="utf-8", errors="replace").strip() + if not text: + return [] + if file_path.suffix.lower() == ".json": + value = json.loads(text) + if isinstance(value, list): + return [item for item in value if isinstance(item, dict)] + return [value] if isinstance(value, dict) else [] + + samples: list[dict[str, Any]] = [] + for line in text.splitlines(): + line = line.strip() + if not line: + continue + try: + obj = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(obj, dict): + samples.append(obj) + return samples + + +def _sample_question(sample: dict[str, Any]) -> str: + """Extract the user-facing question / instruction from a sample.""" + if sample.get("instruction"): + text = sample["instruction"] + if sample.get("input"): + text += "\n" + sample["input"] + return text + if sample.get("question"): + return sample["question"] + # ShareGPT-style: use the last user message as question + messages = sample.get("messages") or [] + user_msgs = [m["content"] for m in messages if m.get("role") == "user"] + return user_msgs[-1] if user_msgs else "" + + +def _sample_reference(sample: dict[str, Any]) -> str: + """Extract the reference answer from a sample.""" + if sample.get("output"): + return sample["output"] + if sample.get("answer"): + return sample["answer"] + messages = sample.get("messages") or [] + assistant_msgs = [m["content"] for m in messages if m.get("role") == "assistant"] + return assistant_msgs[-1] if assistant_msgs else "" + + +# --------------------------------------------------------------------------- +# Basic metrics +# --------------------------------------------------------------------------- + +def _compute_bleu(references: list[str], predictions: list[str], ngram: int = 4) -> dict[str, Any]: + """Compute BLEU score via sacrebleu (corpus-level).""" + try: + from sacrebleu.metrics import BLEU + except ImportError: + return {"enabled": False, "error": "sacrebleu not installed", "score": 0} + bleu = BLEU(max_ngram_order=ngram) + # sacrebleu expects list-of-strings; we have one reference per prediction + score = bleu.corpus_score(predictions, [references]) + return { + "enabled": True, + "score": round(score.score, 2), + "bleu": round(score.score, 2), + } + + +def _compute_rouge(references: list[str], predictions: list[str], methods: list[str] | None = None) -> dict[str, Any]: + """Compute ROUGE scores via rouge-score.""" + try: + from rouge_score import rouge_scorer + except ImportError: + return {"enabled": False, "error": "rouge-score not installed", "score": 0} + methods = methods or ["rouge1", "rouge2", "rougeL"] + # Normalize: map "rouge_1"/"rouge1" → "rouge1", "rouge_l"/"rougeL" → "rougeL" + _rouge_aliases = {"rouge_1": "rouge1", "rouge_2": "rouge2", "rouge_l": "rougeL"} + methods = [_rouge_aliases.get(m, m.replace("_", "")) for m in methods] + scorer = rouge_scorer.RougeScorer(methods, use_stemmer=True) + totals: dict[str, float] = {} + n = max(len(predictions), 1) + for ref, pred in zip(references, predictions): + result = scorer.score(ref, pred) + for key in methods: + totals[key] = totals.get(key, 0) + result[key].fmeasure + avg = {k: round(v / n, 4) for k, v in totals.items()} + return {"enabled": True, "score": round(avg.get("rougeL", avg.get("rouge1", 0)) * 100, 2), **avg} + + +def _compute_cosine(references: list[str], predictions: list[str]) -> dict[str, Any]: + """Compute average cosine similarity via sklearn.""" + try: + from sklearn.feature_extraction.text import TfidfVectorizer + from sklearn.metrics.pairwise import cosine_similarity + except ImportError: + return {"enabled": False, "error": "scikit-learn not installed", "score": 0} + try: + vectorizer = TfidfVectorizer() + tfidf = vectorizer.fit_transform(references + predictions) + n = len(references) + ref_vec = tfidf[:n] + pred_vec = tfidf[n:] + sims = cosine_similarity(ref_vec, pred_vec).diagonal() + return {"enabled": True, "score": round(float(sims.mean()) * 100, 2)} + except ValueError: + return {"enabled": True, "score": 0, "error": "insufficient text for vectorization"} + + +def _normalize_text(value: str) -> str: + return re.sub(r"\s+", " ", str(value or "").strip().lower()) + + +def _compute_exact_match(references: list[str], predictions: list[str]) -> dict[str, Any]: + total = len(predictions) + if not total: + return {"enabled": True, "score": 0, "matched": 0, "total": 0} + matched = sum( + 1 + for ref, pred in zip(references, predictions) + if _normalize_text(ref) == _normalize_text(pred) + ) + return {"enabled": True, "score": round(matched / total * 100, 2), "matched": matched, "total": total} + + +def _compute_text_similarity(references: list[str], predictions: list[str]) -> dict[str, Any]: + if not predictions: + return {"enabled": True, "score": 0} + scores = [ + SequenceMatcher(None, _normalize_text(ref), _normalize_text(pred)).ratio() + for ref, pred in zip(references, predictions) + ] + return {"enabled": True, "score": round(sum(scores) / max(len(scores), 1) * 100, 2)} + + +# --------------------------------------------------------------------------- +# LLM Judge +# --------------------------------------------------------------------------- + +def _judge_sample( + question: str, + reference: str, + prediction: str, + config: dict[str, Any], +) -> dict[str, Any]: + """Call an OpenAI-compatible LLM to judge a single sample. + + Returns a dict with keys: + score, max_score, passed, judgement, evaluation_reason, error_type + """ + api_url = (config.get("api_url") or "").strip().rstrip("/") + api_key = (config.get("api_key") or "").strip() + eval_model = (config.get("eval_model") or "").strip() + eval_prompt = (config.get("eval_prompt") or "").strip() + score_min = float(config.get("score_min", 0)) + score_max = float(config.get("score_max", 5)) + pass_threshold = float(config.get("pass_threshold", 3)) + + if not api_url or not eval_model: + return {"score": 0, "max_score": score_max, "passed": False, "judgement": "未配置", + "evaluation_reason": "未配置评测模型", "error_type": "其他"} + + system_msg = ( + eval_prompt + or "你是一个专业的评测专家。请根据参考答-案对被测模型的输出进行评分。" + ) + user_msg = ( + f"## 问题\n{question}\n\n" + f"## 参考答案\n{reference}\n\n" + f"## 模型输出\n{prediction}\n\n" + f"请给出 {score_min}-{score_max} 分的评分,并说明理由。" + ) + + try: + import urllib.request + import urllib.error + + body = json.dumps({ + "model": eval_model, + "messages": [ + {"role": "system", "content": system_msg}, + {"role": "user", "content": user_msg}, + ], + "temperature": 0.3, + "max_tokens": 512, + }).encode("utf-8") + + req = urllib.request.Request( + f"{api_url}/v1/chat/completions", + data=body, + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + }, + ) + resp = urllib.request.urlopen(req, timeout=120) + data = json.loads(resp.read().decode("utf-8")) + reply = data["choices"][0]["message"]["content"] + except Exception as exc: + return {"score": 0, "max_score": score_max, "passed": False, + "judgement": "错误", "evaluation_reason": f"评测模型调用失败: {exc}", + "error_type": "其他"} + + # Parse score from reply — look for patterns like "4分" or "Score: 4" + score = 0 + import re + score_patterns = [ + r'(?:得分|分数|评分|score)[^\d]*(\d+(?:\.\d+)?)', + r'(\d+(?:\.\d+)?)\s*分', + r'(\d+(?:\.\d+)?)\s*/\s*\d+', + ] + for pat in score_patterns: + m = re.search(pat, reply, re.IGNORECASE) + if m: + try: + score = float(m.group(1)) + except ValueError: + continue + break + score = max(score_min, min(score_max, score)) + passed = score >= pass_threshold + + # Determine judgement label + if score >= pass_threshold + 1: + judgement = "正确" + elif score >= pass_threshold: + judgement = "部分正确" + else: + judgement = "错误" + + # Guess error type from reply + reply_lower = reply.lower() + if any(w in reply_lower for w in ["幻觉", "hallucination", "编造"]): + error_type = "幻觉" + elif any(w in reply_lower for w in ["不完整", "incomplete", "遗漏"]): + error_type = "不完整" + elif any(w in reply_lower for w in ["格式", "format"]): + error_type = "格式偏差" + elif any(w in reply_lower for w in ["混淆", "confusion", "错误"]): + error_type = "混淆" + else: + error_type = "其他" + + return { + "score": score, + "max_score": score_max, + "passed": passed, + "judgement": judgement, + "evaluation_reason": reply[:2000], + "error_type": error_type, + } + + +# --------------------------------------------------------------------------- +# Main entry point +# --------------------------------------------------------------------------- + +def run_eval(config: dict[str, Any]) -> dict[str, Any]: + """Execute a full evaluation run. Returns the result dict (also written to file).""" + model_path = config["model_name_or_path"] + adapter_path = config.get("adapter_name_or_path", "") + template = config.get("template", "qwen") + dataset_path = config["dataset_path"] + output_dir = Path(config["output_dir"]) + output_dir.mkdir(parents=True, exist_ok=True) + basic_cfg = config.get("basic_metrics", {}) + dimension_cfg = config.get("dimension", {}) or {} + output_precision = int(basic_cfg.get("output_precision", 2)) + + # ---- 1. Load dataset ---- + print(f"[eval] loading dataset: {dataset_path}") + raw_samples = _load_dataset(dataset_path) + print(f"[eval] loaded {len(raw_samples)} samples") + + # ---- 2. Load model ---- + print(f"[eval] loading model: {model_path}") + from compute.engines.llama_factory.inference import InferenceSession + session = InferenceSession() + load_result = session.load( + model_name_or_path=model_path, + adapter_name_or_path=adapter_path, + template=template, + infer_backend=config.get("infer_backend", "huggingface"), + infer_dtype=config.get("infer_dtype", "auto"), + ) + if not load_result.get("loaded"): + raise RuntimeError(f"model load failed: {load_result.get('error', 'unknown')}") + print(f"[eval] model loaded OK") + + # ---- 3. Run inference on each sample ---- + samples: list[dict[str, Any]] = [] + predictions: list[str] = [] + references: list[str] = [] + questions: list[str] = [] + + total = len(raw_samples) + judge_enabled = bool(dimension_cfg.get("eval_model") and dimension_cfg.get("api_url")) + print(f"[eval] starting inference on {total} samples, judge={'enabled' if judge_enabled else 'disabled'}") + + for idx, raw in enumerate(raw_samples, start=1): + question = _sample_question(raw) + reference = _sample_reference(raw) + if not question: + print(f"[eval] sample {idx}/{total}: skipped (no question)") + continue + + # Inference + chat_msgs = [{"role": "user", "content": question}] + result = session.chat( + chat_msgs, + temperature=float(config.get("temperature", 0.1)), + top_p=float(config.get("top_p", 0.95)), + max_new_tokens=int(config.get("max_new_tokens", 512)), + do_sample=False, + ) + prediction = result.get("response", "") if not result.get("error") else f"[ERROR] {result['error']}" + + predictions.append(prediction) + references.append(reference) + questions.append(question) + + # LLM Judge + judge_result: dict[str, Any] = {} + if judge_enabled: + judge_result = _judge_sample(question, reference, prediction, dimension_cfg) + + samples.append({ + "index": idx, + "input": question, + "reference_answer": reference, + "model_output": prediction, + "score": judge_result.get("score"), + "max_score": judge_result.get("max_score", dimension_cfg.get("score_max", 5)), + "passed": judge_result.get("passed"), + "judgement": judge_result.get("judgement"), + "evaluation_reason": judge_result.get("evaluation_reason", ""), + "error_type": judge_result.get("error_type"), + "dimension_scores": [ + {"name": "judge_score", "score": judge_result.get("score", 0), + "max_score": judge_result.get("max_score", dimension_cfg.get("score_max", 5))}, + ] if judge_result else [], + "status": "completed", + }) + + progress_pct = int(idx / max(total, 1) * 100) + print(f"[eval] sample {idx}/{total} ({progress_pct}%) done") + + # ---- 4. Compute basic metrics ---- + print(f"[eval] computing basic metrics on {len(predictions)} predictions") + metrics_result: dict[str, Any] = {} + + bleu_cfg = basic_cfg.get("bleu", {}) + if bleu_cfg.get("enabled"): + metrics_result["bleu"] = _compute_bleu(references, predictions, int(bleu_cfg.get("ngram", 4))) + + rouge_cfg = basic_cfg.get("rouge", {}) + if rouge_cfg.get("enabled"): + metrics_result["rouge"] = _compute_rouge(references, predictions, rouge_cfg.get("methods")) + + cosine_cfg = basic_cfg.get("cosine", {}) + if cosine_cfg.get("enabled"): + metrics_result["cosine"] = _compute_cosine(references, predictions) + metrics_result["exact_match"] = _compute_exact_match(references, predictions) + metrics_result["text_similarity"] = _compute_text_similarity(references, predictions) + + # ---- 5. Summarise ---- + completed = len(samples) + if judge_enabled: + scored = [s for s in samples if s.get("score") is not None] + passed_count = len([s for s in scored if s.get("passed")]) + avg_score = round(sum(s["score"] for s in scored) / max(len(scored), 1), output_precision) + max_score = dimension_cfg.get("score_max", 5) + overall_score = round(avg_score / max_score * 100, output_precision) + overall_score_max = 100 + dimension_summary = [{ + "name": "综合评分", + "score": overall_score, + "max_score": 100, + "pass_rate": round(passed_count / max(completed, 1) * 100, 1), + }] + overall_evaluation = f"评测完成:{completed} 样本,{passed_count} 通过,平均 {avg_score}/{max_score} 分" + else: + passed_count = 0 + enabled_scores = [ + float(item.get("score") or 0) + for item in metrics_result.values() + if isinstance(item, dict) and item.get("enabled", True) and item.get("score") is not None + ] + overall_score = round(sum(enabled_scores) / len(enabled_scores), output_precision) if enabled_scores else 0 + overall_score_max = 100 + dimension_summary = [ + { + "name": name, + "score": float(item.get("score") or 0), + "max_score": 100, + "pass_rate": float(item.get("score") or 0), + } + for name, item in metrics_result.items() + if isinstance(item, dict) and item.get("enabled", True) and item.get("score") is not None + ] + overall_evaluation = f"评测完成:{completed} 样本(未配置 LLM 评委)" + + result = { + "overall_score": overall_score, + "overall_score_max": overall_score_max, + "overall_evaluation": overall_evaluation, + "improvement_suggestions": [], + "dimension_summary": dimension_summary, + "samples": samples, + "sample_count": total, + "completed_count": completed, + "passed_count": passed_count, + "basic_metrics": metrics_result, + } + + # ---- 6. Write results ---- + result_path = output_dir / "eval_results.json" + result_path.write_text(json.dumps(result, ensure_ascii=False, indent=2), encoding="utf-8") + print(f"[eval] results written to {result_path}") + return result + + +def main() -> None: + import argparse + parser = argparse.ArgumentParser(description="YG-FT Evaluation Runner") + parser.add_argument("--config", required=True, help="Path to eval config JSON file") + args = parser.parse_args() + + config_path = Path(args.config) + if not config_path.exists(): + print(f"FATAL: config file not found: {args.config}", file=sys.stderr) + sys.exit(1) + + config = json.loads(config_path.read_text(encoding="utf-8")) + start = time.time() + try: + run_eval(config) + elapsed = time.time() - start + print(f"[eval] DONE in {elapsed:.1f}s") + except Exception as exc: + print(f"[eval] FAILED: {exc}", file=sys.stderr) + import traceback + traceback.print_exc() + sys.exit(1) + + +if __name__ == "__main__": + main() diff --git a/compute/engines/llama_factory/inference.py b/compute/engines/llama_factory/inference.py index 4d1e227..0afb036 100644 --- a/compute/engines/llama_factory/inference.py +++ b/compute/engines/llama_factory/inference.py @@ -59,10 +59,18 @@ class InferenceSession: if adapter_name_or_path: args["adapter_name_or_path"] = adapter_name_or_path args.update(kwargs) - model_args, generating_args = get_infer_args(args) - self._model = ChatModel(model_args) - self._tokenizer = self._model.tokenizer - self._generating_args = generating_args + infer_result = get_infer_args(args) + # ChatModel internally re-parses the args dict via get_infer_args, + # so pass the original args (not the parsed dataclass objects). + self._model = ChatModel(args) + self._tokenizer = getattr(self._model, 'tokenizer', None) or self._model.engine.tokenizer + # Extract generating_args (last element) for later use in chat() + generating_args = infer_result[-1] + if hasattr(generating_args, '__dataclass_fields__'): + self._generating_args = {k: v for k, v in vars(generating_args).items() + if not k.startswith('_')} + else: + self._generating_args = dict(generating_args) self._loaded_at = time.time() self._status = "ready" return {"loaded": True, "status": "ready"} @@ -80,6 +88,16 @@ class InferenceSession: pass self._model = None self._tokenizer = None + # 强制释放 PyTorch CUDA 缓存,真正归还 GPU 显存 + try: + import gc + gc.collect() + import torch + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + except Exception: + pass self._status = "idle" self._model_name = "" self._adapter_path = "" @@ -91,11 +109,12 @@ class InferenceSession: if self._status != "ready" or self._model is None: return {"error": "model not loaded", "response": ""} try: - generate_kwargs = {**self._generating_args, "temperature": temperature, "top_p": top_p, "max_new_tokens": max_new_tokens, "do_sample": do_sample} + generate_kwargs = {"temperature": temperature, "top_p": top_p, "max_new_tokens": max_new_tokens, "do_sample": do_sample} generate_kwargs.update(kwargs) - formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + system = next((m["content"] for m in messages if m["role"] == "system"), None) + user_messages = [m for m in messages if m["role"] != "system"] responses = [] - for response in self._model.stream_chat(formatted, generate_kwargs): + for response in self._model.stream_chat(user_messages, system=system, **generate_kwargs): responses.append(response) full_response = "".join(str(r) for r in responses) return {"response": full_response} @@ -108,9 +127,10 @@ class InferenceSession: yield 'data: {"error": "model not loaded"}\n\n' return try: - generate_kwargs = {**self._generating_args, **kwargs} - formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) - for new_text in self._model.stream_chat(formatted, generate_kwargs): + generate_kwargs = {**kwargs} + system = next((m["content"] for m in messages if m["role"] == "system"), None) + user_messages = [m for m in messages if m["role"] != "system"] + for new_text in self._model.stream_chat(user_messages, system=system, **generate_kwargs): yield new_text except Exception as exc: yield 'data: {"error": "' + str(exc) + '"}\n\n' diff --git a/compute/requirements.txt b/compute/requirements.txt index 53945b0..9f7bef7 100644 --- a/compute/requirements.txt +++ b/compute/requirements.txt @@ -4,4 +4,9 @@ python-multipart>=0.0.9 pydantic>=2.7.0 python-dotenv>=1.0.1 httpx>=0.27.0 -llamafactory \ No newline at end of file +# 模型评测指标 +sacrebleu>=2.4.0 +rouge-score>=0.1.2 +scikit-learn>=1.3.0 +# LLaMA-Factory 训练引擎 +llamafactory diff --git a/docker/app/Dockerfile.backend b/docker/app/Dockerfile.backend index 1f03dca..c5a7195 100644 --- a/docker/app/Dockerfile.backend +++ b/docker/app/Dockerfile.backend @@ -11,7 +11,7 @@ RUN pip install --upgrade pip -i https://pypi.tuna.tsinghua.edu.cn/simple \ && pip install -r /tmp/requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple \ && rm -f /tmp/requirements.txt -RUN python -c "import fastapi, uvicorn, psycopg, sqlalchemy, redis, jwt, passlib, httpx, alembic; print('backend dependency check ok')" +RUN python -c "import fastapi, uvicorn, psycopg, psycopg_pool, sqlalchemy, redis, jwt, passlib, httpx, alembic; print('backend dependency check ok')" RUN mkdir -p /opt/yg-ft/logs/backend /data/yg-ft \ && chmod -R 0775 /opt/yg-ft /data/yg-ft diff --git a/docs/模型评测功能总结.md b/docs/模型评测功能总结.md new file mode 100644 index 0000000..f6ed2d6 --- /dev/null +++ b/docs/模型评测功能总结.md @@ -0,0 +1,121 @@ +# 模型评测功能总结 + +本项目(基于 LLaMA-Factory 的微调训练平台)包含 **4 套相对独立** 的模型评测能力,分别面向不同的使用场景: + +| 能力 | 入口/目录 | 评测类型 | 打分方式 | +| --- | --- | --- | --- | +| 1. 学术 Benchmark 评测 | `llamafactory/eval/` | 选择题式基准(类 MMLU/C-Eval) | 选项匹配 + few-shot | +| 2. 评估工作台 | `backend/app/api/v1/eval/` | 生成式问答(指令跟随) | BLEU / ROUGE / ExactMatch + 可选 LLM 评审 | +| 3. 平台评估系统 | `backend/app/api/v1/evaluation/` | 基于评估数据集的问答 | 判卷模型(judge model)打分(0–5 分) | +| 4. 训练时验证评估 | `backend/app/services/task_runner.py` | 训练验证集 | loss 指标 | + +下面分别说明。 + +--- + +## 1. 学术 Benchmark 评测(LLaMA-Factory 原生) + +面向标准学术选择题基准(如 MMLU、C-Eval 等),复用 LLaMA-Factory 原生的评测框架。 + +**核心文件** +- `llamafactory/eval/evaluator.py`:`Evaluator` 类 + `run_eval()` 入口 +- `llamafactory/eval/template.py`:评测 prompt 模板(中/英,含 few-shot 示例构建) +- `llamafactory/hparams/evaluation_args.py`:`EvaluationArguments` 配置类 + +**工作流程** +1. 按 `task`(benchmark 名称)加载数据集,按科目(subject)拆分。 +2. 每个样本构造 few-shot 提示词(`n_shot` 控制示例数,由 `lang` 决定中/英模板),将题干与候选选项拼入 prompt。 +3. 调用模型推理得到预测,与标准答案比对,统计每个科目及整体的 `accuracy`。 +4. 结果写入 `save_dir`,打印各科目与平均准确率。 + +**关键参数(`EvaluationArguments`)** +- `task`:基准数据集名 +- `batch_size` / `n_shot` / `lang` / `save_dir` / `seed` +- `model_name_or_path`、`template`、`trust_remote_code` 等模型相关参数 + +> 该能力属于框架底层,本平台前端未直接提供操作入口,主要通过配置文件/脚本调用。 + +--- + +## 2. 评估工作台(生成式评测 + 指标计算) + +后端路由位于 `backend/app/api/v1/eval/__init__.py`,前端称为「评估工作台」。**适用于评测模型的指令跟随与生成质量**,并支持 LLM 作为裁判(LLM-as-a-Judge)。 + +**API 端点** +- `GET /evaluation/tasks`:列出评测任务(`frontend/src/api/evaluation.ts:listTasks`) +- `POST /evaluation/run`:提交一次评测(`runEval`) +- `GET /evaluation/report/{task_id}`:拉取评测报告(`getReport`) +- `DELETE /evaluation/tasks/{task_id}`:删除任务(`deleteTask`) + +**评测流程(`run_eval`)** +1. 通过 **LLaMA-Factory 数据管道**(`get_dataset`) 加载数据集,支持 `subset` 与抽样(`eval_sample`)。 +2. 用 **原生 transformers** 加载模型在本地做生成推理(单进程顺序生成,便于展示样本)。 +3. 计算客观指标(`compute_score`): + - `BLEU`(sacrebleu) + - `ROUGE-1 / ROUGE-2 / ROUGE-L`(rouge-score) + - `Exact Match` +4. **可选 LLM 评审**(judge):当配置了 `judge_model` / `judge_api_base` / `judge_api_key` 时,调用 OpenAI 兼容接口对每条样本打分(10 分制),并输出 4 个维度与理由: + - 核心事实正确性 `factual` + - 信息完整性 `completeness` + - 无幻觉 `no_hallucination` + - 格式合规性 `format` + - 综合分 `score` + `reason` +5. 任务状态持久化在后端 `eval_tasks.json`(支持 running/completed/failed/stopped),前端轮询进度。 + +**前端页面** +- `frontend/src/views/evaluation/EvaluateTask.vue`:任务列表、创建评测对话框(选模型、数据集、指标、可选 judge 配置) +- `frontend/src/views/evaluation/EvaluateReport.vue`:报告页,展示综合得分、BLEU、ROUGE-L、各维度指标及「参考答案 vs 模型预测 vs LLM 评审」对比样例 + +--- + +## 3. 平台评估系统(基于评估数据集 + 判卷模型) + +后端路由位于 `backend/app/api/v1/evaluation/__init__.py`,是平台业务层自研的评测体系。通过「评估数据集」组织题目,可一次性对 **多个被测模型 + 指定判卷模型** 进行批量评分。 + +**核心概念(数据模型 `backend/app/models/models.py`)** +- `EvalDataset`(`models.py:131`):评估数据集,从项目问答对(`Question`/`Chunk`)中按 `question_type`(mixed/fact/reasoning)选题构建,状态 `pending/running/completed/failed`。 +- `EvalResult`(`models.py:147`):单条评测结果,含 `judge_score`(0–5 分)、`is_correct`(true/false/partial)、`feedback`、`expected_answer` 等。 +- `Task`(`models.py:184`):后台任务,`task_type="model-evaluation"`,记录进度与 `model_info`(存放平均分等汇总)。 + +**评测流程(`process_evaluation_task`,`backend/app/services/task_processor.py:336` 起)** +1. 加载评估数据集关联的题目,可选带入 `chunk` 上下文(RAG 场景)。 +2. 对每道题,先用 `build_eval_prompt` 组合「上下文 + 题目 + 参考答案」,调用 **判卷模型**(`call_model`,temperature=0.3)生成评分。 +3. `parse_eval_result` 解析出 `score`(0–5)、`is_correct`、`feedback`,写入 `EvalResult`。 +4. 逐题提交进度(`completed_count` / `progress`),支持中途 `stopped`。 +5. 汇总:`avg_score = 总分/有效数 × 20`(换算百分制),`avg_score_5 = 总分/有效数`(5 分制),存入 `task.model_info`。判定规则:得分 **≥3 视为正确**。 + +**特点** +- 判卷与被测模型解耦:被测模型给出答案,判卷模型(judge)独立评分,降低自评偏差。 +- 支持失败隔离:单题异常写入 `evaluation_status: failed` 记录而不中断整体任务。 + +--- + +## 4. 训练时验证评估 + +在微调训练任务执行期间,由 `backend/app/services/task_runner.py` 的 `do_eval` 触发: + +- 在训练过程中对验证集(validation set)计算 `eval_loss`,用于监控过拟合。 +- 结果回填到 `Task` 的 `loss_info` / `detail`,前端绘制 loss 曲线。 +- 属于训练配套的轻量评估,不参与上述 1–3 的业务评测。 + +--- + +## 附属:前端评测相关页面 + +| 文件 | 作用 | +| --- | --- | +| `frontend/src/views/evaluation/EvaluateTask.vue` | 评估工作台:任务列表 + 创建评测 | +| `frontend/src/views/evaluation/EvaluateReport.vue` | 评估报告:指标卡 + 维度标签 + 对比样例 | +| `frontend/src/api/evaluation.ts` | 评估工作台接口封装 | +| 平台评估系统入口 | 评估数据集管理 + 评估任务(model-evaluation)创建与结果查看 | + +--- + +## 小结 + +- **想要学术榜单式准确率** → 用能力 1(LLaMA-Factory `eval/`)。 +- **想要开放式生成质量(BLEU/ROUGE + LLM 评审)** → 用能力 2(评估工作台 `/evaluation/run`)。 +- **想要基于自有问答数据、用判卷模型批量打分** → 用能力 3(平台评估系统 `model-evaluation` 任务)。 +- **训练过程监控** → 能力 4(`do_eval` 验证集 loss)。 + +三种业务评测(1/2/3)相互独立,可并存于同一平台;数据模型(`EvalDataset`/`EvalResult`/`Task`)主要服务于能力 3,而能力 2 使用独立的 `eval_tasks.json` 文件持久化。 diff --git a/frontend/src/api/modules/compare.ts b/frontend/src/api/modules/compare.ts index 30586cf..35f1ee5 100644 --- a/frontend/src/api/modules/compare.ts +++ b/frontend/src/api/modules/compare.ts @@ -64,6 +64,27 @@ export const streamChat = async (data: any): Promise => { } } +/** 真实流式对话 — 使用 fetch 调用后端 SSE 端点,返回 Response 供 ReadableStream 消费 */ +export const streamChatReal = (data: any): Promise => { + const messages = data.messages || [] + if (!messages.length && data.user_question) { + if (data.system_prompt) { + messages.push({ role: 'system', content: data.system_prompt }) + } + messages.push({ role: 'user', content: data.user_question }) + } + return fetch('/modelTF/model-compare/stream-chat', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + messages, + temperature: data.temperature ?? 0.7, + top_p: data.top_p ?? 0.95, + max_tokens: data.max_tokens ?? 2048, + }), + }) +} + /** 非流式对话(按端口代理) */ export const chatWithPort = (data: any) => post('/model-compare/chat-with-port', data) @@ -73,8 +94,8 @@ export const batchChat = (data: any) => post('/model-chat/batch', data) /** 本地 transformers 模型对话 */ export const localChat = (data: any) => post('/model-chat/local/chat', data) -/** 预加载本地模型 */ -export const preloadLocalModel = (data: any) => post('/model-chat/local/preload', data) +/** 预加载本地模型(模型加载耗时长,超时 5 分钟) */ +export const preloadLocalModel = (data: any) => post('/model-chat/local/preload', data, { timeout: 300000 }) -/** 预加载已训练模型 */ -export const preloadTrainedModel = (data: any) => post('/model-chat/trained/preload', data) +/** 预加载已训练模型(超时 5 分钟) */ +export const preloadTrainedModel = (data: any) => post('/model-chat/trained/preload', data, { timeout: 300000 }) diff --git a/frontend/src/api/modules/compute.ts b/frontend/src/api/modules/compute.ts index ad6d10e..344782c 100644 --- a/frontend/src/api/modules/compute.ts +++ b/frontend/src/api/modules/compute.ts @@ -1,4 +1,4 @@ -import { get, post, put } from '../request' +import { del, get, post, put } from '../request' export interface ComputeNode { id: string @@ -95,6 +95,9 @@ export const createComputeNode = (data: ComputeNodePayload) => export const updateComputeNode = (id: string, data: Partial) => put(`/compute/nodes/${id}`, data) +export const deleteComputeNode = (id: string) => + del<{ deleted: string }>(`/compute/nodes/${id}`) + export const testComputeNode = (id: string) => post<{ node_id: string; success: boolean; latency_ms: number; gpu_count: number; error?: string }>(`/compute/nodes/${id}/test-connection`) diff --git a/frontend/src/api/modules/dataset.ts b/frontend/src/api/modules/dataset.ts index 92a7e3a..ec32bf7 100644 --- a/frontend/src/api/modules/dataset.ts +++ b/frontend/src/api/modules/dataset.ts @@ -29,6 +29,7 @@ export const uploadDatasetFiles = (datasetId: string | number, files: File[]) => files.forEach((f) => formData.append('files', f)) return post(`/dataset-manage/upload/${datasetId}`, formData, { headers: { 'Content-Type': 'multipart/form-data' }, + timeout: 120000, }) } diff --git a/frontend/src/composables/useStreamChat.ts b/frontend/src/composables/useStreamChat.ts index 331cbe9..d322529 100644 --- a/frontend/src/composables/useStreamChat.ts +++ b/frontend/src/composables/useStreamChat.ts @@ -1,5 +1,5 @@ import { ref } from 'vue' -import { streamChat } from '@/api/modules/compare' +import { streamChat, streamChatReal } from '@/api/modules/compare' export interface StreamMessage { /** 用户问题 */ @@ -20,6 +20,11 @@ export interface StreamMessage { error?: string } +export interface SendOptions { + /** 是否使用 mock 模式(默认 true,向后兼容) */ + useMock?: boolean +} + /** * 流式对话 composable * 移植自原 model-chat.html: @@ -65,8 +70,10 @@ export function useStreamChat() { /** * 发起流式对话 * @param payload 后端请求体 { port, model_name, model_path, system_prompt, user_question, ... } + * @param options 可选配置 { useMock?: boolean } */ - async function send(payload: any) { + async function send(payload: any, options?: SendOptions) { + const useMock = options?.useMock ?? true loading.value = true message.value = { question: payload.user_question || '', @@ -82,7 +89,10 @@ export function useStreamChat() { const UPDATE_INTERVAL = 50 // 50ms 节流 try { - const response = await streamChat(payload) + const response = useMock + ? await streamChat(payload) + : await streamChatReal(payload) + if (!response.ok) { throw new Error(`HTTP ${response.status}`) } diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 9e1f446..ce6c87f 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -132,6 +132,7 @@ export interface FineTuneTask { train_dataset_id?: number | string auto_merge?: boolean output_model_name?: string + compute_node_id?: string gpus?: number[] batch_size?: number learning_rate?: number @@ -355,16 +356,16 @@ export interface GpuInfo { power_w: number id?: number uuid?: string - status?: 'idle' | 'busy' | 'warning' | 'offline' + status?: 'idle' | 'busy' | 'reserved' | 'warning' | 'offline' memory_percent?: number power_limit_w?: number processes?: GpuProcess[] fan_speed?: number clock_mhz?: number - driver_version?: string node_id?: string node_code?: string node_name?: string + driver_version?: string } export interface SystemInfo { diff --git a/frontend/src/views/approvals/ApprovalInstanceView.vue b/frontend/src/views/approvals/ApprovalInstanceView.vue index db29b4f..4664758 100644 --- a/frontend/src/views/approvals/ApprovalInstanceView.vue +++ b/frontend/src/views/approvals/ApprovalInstanceView.vue @@ -3,7 +3,8 @@ import { onMounted, reactive, ref } from 'vue' import { ElMessage } from 'element-plus' import DataTablePage from '@/components/DataTablePage.vue' import { getApprovalInstances, decideApproval, type ApprovalInstance } from '@/api/modules/approval' -import { getUsers, type SystemUser } from '@/api/modules/system' +import { getUsers } from '@/api/modules/system' +import type { SystemUser } from '@/types' const loading = ref(false) const instances = ref([]) @@ -48,6 +49,10 @@ function openDecide(inst: ApprovalInstance) { showDecide.value = true } +function asApprovalInstance(row: unknown): ApprovalInstance { + return row as ApprovalInstance +} + async function submitDecision() { if (!current.value) return if (!decision.value.approver_id) { @@ -72,7 +77,7 @@ onMounted(() => {