fix: 推理/评测结果同步、数据集统计与算力节点管理增强

后端:
- 抽取 fetch_eval_result_content 复用函数,model_eval_detail 直接应用评测任务结果
- health 接口移除数据库依赖,返回静态指标
- 数据集: count_dataset_records JSON 感知计数; 文件统计改为从 dataset_files 聚合重算; 在线编辑记录 size/record_count/version_no; 上传同步批处理
- 算力节点: 调度支持 requested GPU 子集校验与容量计算; 新增 delete_compute_node(含活动任务保护)及 DELETE 接口; 连接池 connect_timeout
- 评测任务落库 basic_metrics/score/completed_time, failed/stopped 记录 error

评测引擎:
- _load_dataset 支持 JSON/JSONL 文件
- 新增 exact match 与文本相似度指标, 余弦相似度去掉 2 样本限制

前端:
- 算力节点列表「维护」改为「删除」(带确认弹窗), compute.ts 新增 deleteComputeNode
- 数据集上传超时调整为 120s; FineTuneTask 增加 compute_node_id; GpuInfo 状态增加 reserved
This commit is contained in:
wuyongtao
2026-08-03 15:49:21 +08:00
parent cc08b164d0
commit 5cc306eb0a
15 changed files with 402 additions and 112 deletions

View File

@@ -1,7 +1,6 @@
from fastapi import APIRouter from fastapi import APIRouter
from app.core.logging import get_logger from app.core.logging import get_logger
from app.db.platform_store import get_platform_store
router = APIRouter() router = APIRouter()
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -10,5 +9,9 @@ logger = get_logger(__name__)
@router.get("/health") @router.get("/health")
async def health_check() -> dict[str, object]: async def health_check() -> dict[str, object]:
logger.info("health check requested") 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},
}

View File

@@ -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.core.config import get_settings
from app.db.platform_store import get_platform_store from app.db.platform_store import get_platform_store
from app.modules.compute_gateway.client import ComputeNodeClient 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() router = APIRouter()
@@ -915,6 +915,7 @@ async def upload_dataset_files(
) -> dict[str, Any]: ) -> dict[str, Any]:
created: list[dict[str, Any]] = [] created: list[dict[str, Any]] = []
compute_sync: list[dict[str, Any]] = [] compute_sync: list[dict[str, Any]] = []
pending_sync: list[tuple[str, str, bytes]] = []
store = get_platform_store() store = get_platform_store()
try: try:
store.dataset(dataset_id) store.dataset(dataset_id)
@@ -926,16 +927,18 @@ async def upload_dataset_files(
content = raw.decode("utf-8", errors="replace") content = raw.decode("utf-8", errors="replace")
created_file = store.add_dataset_file(conn, dataset_id, file.filename or "upload.jsonl", content) created_file = store.add_dataset_file(conn, dataset_id, file.filename or "upload.jsonl", content)
created.append(created_file) created.append(created_file)
if sync_to_compute: pending_sync.append((created_file["id"], created_file["name"], raw))
compute_sync.extend( if sync_to_compute:
await _sync_dataset_file_to_compute_nodes( for file_id, file_name, raw in pending_sync:
store, compute_sync.extend(
dataset_id, await _sync_dataset_file_to_compute_nodes(
created_file["id"], store,
created_file["name"], dataset_id,
raw, file_id,
) file_name,
raw,
) )
)
return ok({"files": created, "compute_sync": compute_sync}) return ok({"files": created, "compute_sync": compute_sync})
@@ -1273,8 +1276,7 @@ async def model_eval_detail(task_id: str, current_user: dict = Depends(get_curre
try: try:
store = get_platform_store() store = get_platform_store()
task = store.eval_task(task_id) task = store.eval_task(task_id)
# If the eval job completed on a compute node, try to load results if task.get("compute_job_id") and task.get("compute_node_id") and task.get("status") in {"queued", "running", "completed"}:
if task.get("result_artifact_path"):
node = next( node = next(
(n for n in store.compute_nodes() if n["id"] == task.get("compute_node_id")), (n for n in store.compute_nodes() if n["id"] == task.get("compute_node_id")),
None, None,
@@ -1283,11 +1285,10 @@ async def model_eval_detail(task_id: str, current_user: dict = Depends(get_curre
try: try:
client = ComputeNodeClient(node["api_base_url"]) client = ComputeNodeClient(node["api_base_url"])
job = await client.get_job(task["compute_job_id"]) job = await client.get_job(task["compute_job_id"])
artifacts = job.get("artifacts") or [] result_content = None
for art in artifacts: if job.get("status") == "completed" and not task.get("samples"):
if art.get("name") == "eval_results.json": result_content = await fetch_eval_result_content(client, node, job)
task["_result_artifact"] = art task = store.apply_eval_job_result(task_id, job, result_content)
break
except Exception: except Exception:
pass pass
except KeyError: except KeyError:
@@ -1790,6 +1791,16 @@ async def update_compute_node(node_id: str, payload: dict[str, Any] = Body(...))
raise fail(400, str(exc)) 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") @router.post("/compute/nodes/{node_id}/test-connection")
async def test_compute_node(node_id: str) -> dict[str, Any]: async def test_compute_node(node_id: str) -> dict[str, Any]:
store = get_platform_store() store = get_platform_store()

View File

@@ -88,6 +88,22 @@ def parse_size_bytes(value: Any) -> int:
return max(0, round(amount * _SIZE_UNIT_BYTES[unit])) 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: def version_number(value: Any, default: int = 0) -> int:
try: try:
number = int(value) number = int(value)
@@ -296,6 +312,7 @@ class PlatformStore:
# request (notably expensive against the remote PostgreSQL instance). # request (notably expensive against the remote PostgreSQL instance).
# TCP keepalive 让操作系统持续保活连接,抵抗远程库空闲静默断连。 # TCP keepalive 让操作系统持续保活连接,抵抗远程库空闲静默断连。
pool_kwargs = { pool_kwargs = {
"connect_timeout": 5,
"keepalives": 1, "keepalives": 1,
"keepalives_idle": 30, "keepalives_idle": 30,
"keepalives_interval": 10, "keepalives_interval": 10,
@@ -1311,6 +1328,7 @@ class PlatformStore:
file_size_bytes = int(file_row.get("size_bytes") or 0) file_size_bytes = int(file_row.get("size_bytes") or 0)
if file_size_bytes <= 0: if file_size_bytes <= 0:
file_size_bytes = parse_size_bytes(file_row.get("size")) file_size_bytes = parse_size_bytes(file_row.get("size"))
file_record_count = int(file_row.get("record_count") or 0)
decoded_files.append( decoded_files.append(
{ {
"id": file_row["id"], "id": file_row["id"],
@@ -1319,7 +1337,7 @@ class PlatformStore:
"size_bytes": file_size_bytes, "size_bytes": file_size_bytes,
**dataset_file_version_summary(file_row), **dataset_file_version_summary(file_row),
"create_time": file_row["create_time"], "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"), "split": metadata.get("file_split"),
} }
) )
@@ -1337,6 +1355,9 @@ class PlatformStore:
total_size_bytes = int(row.get("size_bytes") or 0) total_size_bytes = int(row.get("size_bytes") or 0)
if total_size_bytes <= 0: if total_size_bytes <= 0:
total_size_bytes = parse_size_bytes(row.get("size")) 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( current_version_nos = sorted(
{ {
int(item["current_version_no"]) int(item["current_version_no"])
@@ -1346,6 +1367,8 @@ class PlatformStore:
) )
return { return {
**dict(row), **dict(row),
"count": total_record_count,
"record_count": total_record_count,
"size_bytes": total_size_bytes, "size_bytes": total_size_bytes,
"current_version_no": ( "current_version_no": (
current_version_nos[0] if len(current_version_nos) == 1 else None current_version_nos[0] if len(current_version_nos) == 1 else None
@@ -1420,7 +1443,7 @@ class PlatformStore:
version_id = f"{file_id}_v1" version_id = f"{file_id}_v1"
size_bytes = len(content.encode("utf-8")) size_bytes = len(content.encode("utf-8"))
size = f"{size_bytes} B" size = f"{size_bytes} B"
record_count = len([line for line in content.splitlines() if line.strip()]) record_count = count_dataset_records(content)
version = { version = {
"id": version_id, "id": version_id,
"version": 1, "version": 1,
@@ -1453,10 +1476,18 @@ class PlatformStore:
) )
conn.execute( conn.execute(
"""UPDATE datasets """UPDATE datasets
SET count=count+?, record_count=record_count+?, SET count=stats.record_count,
size_bytes=size_bytes+?, size=((size_bytes+?)::text || ' B') 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=?""", WHERE id=?""",
(record_count, record_count, size_bytes, size_bytes, dataset_id), (dataset_id, dataset_id),
) )
return { return {
"id": file_id, "id": file_id,
@@ -1561,19 +1592,57 @@ class PlatformStore:
row = conn.execute("SELECT * FROM dataset_files WHERE id=?", (file_id,)).fetchone() row = conn.execute("SELECT * FROM dataset_files WHERE id=?", (file_id,)).fetchone()
if not row: if not row:
raise KeyError(file_id) 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"], []) versions = json_loads(row["versions"], [])
version = { version = {
"id": f"{file_id}_v{len(versions) + 1}", "id": f"{file_id}_v{len(versions) + 1}",
"version": len(versions) + 1, "version": len(versions) + 1,
"version_no": len(versions) + 1,
"create_time": utcnow(), "create_time": utcnow(),
"description": payload.get("description", "online edit"), "description": payload.get("description", "online edit"),
"size_bytes": size_bytes,
"record_count": record_count,
} }
versions.append(version) versions.append(version)
conn.execute( 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]: def activate_file_version(self, file_id: str, version_id: str) -> dict[str, Any]:
with self.connect() as conn: with self.connect() as conn:
@@ -2173,6 +2242,14 @@ class PlatformStore:
"sample_count": result_content.get("sample_count", 0), "sample_count": result_content.get("sample_count", 0),
"completed_count": result_content.get("completed_count", 0), "completed_count": result_content.get("completed_count", 0),
"passed_count": result_content.get("passed_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) return self.update_eval_task(task_id, updates)
@@ -2342,6 +2419,15 @@ class PlatformStore:
).fetchall() ).fetchall()
return {int(row["gpu_index"]) for row in rows} 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]: 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 = payload.get("requested_node_id") or payload.get("compute_node_id")
requested_gpus = [int(item) for item in payload.get("gpus") or []] requested_gpus = [int(item) for item in payload.get("gpus") or []]
@@ -2349,13 +2435,15 @@ class PlatformStore:
candidates = [ candidates = [
n n
for n in nodes 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: if requested_gpus:
requested_gpu_set = set(requested_gpus)
candidates = [ candidates = [
node node
for node in candidates 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: if requested:
selected = next((n for n in candidates if n["id"] == requested), None) selected = next((n for n in candidates if n["id"] == requested), None)
@@ -2370,8 +2458,8 @@ class PlatformStore:
reason = "disabled" reason = "disabled"
elif node["scheduler_status"] != "online": elif node["scheduler_status"] != "online":
reason = f"status={node['scheduler_status']}" reason = f"status={node['scheduler_status']}"
elif 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']}/{node['max_parallel_jobs']}" reason = f"capacity full {node['current_running_jobs']}/{self._node_capacity(node)}"
else: else:
reason = "not selected" reason = "not selected"
reasons.append(f"{node['code']}({reason})") reasons.append(f"{node['code']}({reason})")
@@ -2722,6 +2810,24 @@ class PlatformStore:
) )
return next(node for node in self.compute_nodes() if node["id"] == node_id) 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]: 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) current = next((n for n in self.compute_nodes() if n["id"] == node_id), None)
if not current: if not current:
@@ -2890,11 +2996,12 @@ class PlatformStore:
} }
def health_metrics(self) -> dict[str, float]: 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 { return {
"cpu_percent": info["cpu"]["percent"], "cpu_percent": 0.0,
"memory_percent": info["memory"]["percent"], "memory_percent": 0.0,
"disk_percent": info["disk"]["percent"], "disk_percent": 0.0,
} }
def queue(self) -> list[dict[str, Any]]: def queue(self) -> list[dict[str, Any]]:

View File

@@ -207,7 +207,8 @@ class ComputeNodeClient:
"resource_id": resource_id or "", "resource_id": resource_id or "",
} }
files = {"file": (filename, content)} 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( response = await client.post(
_join_url(self.api_base_url, f"{self.route_prefix}/compute/files/upload"), _join_url(self.api_base_url, f"{self.route_prefix}/compute/files/upload"),
data=data, data=data,

View File

@@ -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) 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]: async def poll_compute_jobs_once() -> dict[str, Any]:
store = get_platform_store() store = get_platform_store()
synced: list[dict[str, Any]] = [] synced: list[dict[str, Any]] = []
@@ -66,18 +84,7 @@ async def poll_compute_jobs_once() -> dict[str, Any]:
# Try to read eval_results.json from the job output directory # Try to read eval_results.json from the job output directory
if job.get("status") == "completed" and job.get("output_dir"): if job.get("status") == "completed" and job.get("output_dir"):
try: try:
full_path = f"{job['output_dir'].rstrip('/')}/eval_results.json" result_content = await fetch_eval_result_content(client, node, job)
# Convert absolute path to relative (strip YG_FT_DATA_ROOT prefix)
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
settings_path = f"{node['api_base_url'].rstrip('/')}/modelTF/compute/files/read"
async with httpx.AsyncClient(timeout=30, headers=client.headers()) as http:
read_resp = await http.get(settings_path, params={"path": rel_path})
if read_resp.status_code == 200:
result_content = read_resp.json()
except Exception: except Exception:
pass pass
store.apply_eval_job_result(eval_task["id"], job, result_content) store.apply_eval_job_result(eval_task["id"], job, result_content)

View File

@@ -13,30 +13,42 @@ from __future__ import annotations
import json import json
import math import math
import re
import sys import sys
import time import time
from difflib import SequenceMatcher
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
def _load_jsonl(path: str) -> list[dict[str, Any]]: def _load_dataset(path: str) -> list[dict[str, Any]]:
"""Load a JSONL dataset file. Each line must be a JSON object. """Load a JSON or JSONL dataset file.
Supports common field names used across the platform: Supports common field names used across the platform:
* ``instruction`` + ``input`` + ``output`` (Alpaca-style) * ``instruction`` + ``input`` + ``output`` (Alpaca-style)
* ``question`` + ``answer`` * ``question`` + ``answer``
* ``messages`` (ShareGPT-style the last assistant message is treated as reference) * ``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]] = [] samples: list[dict[str, Any]] = []
with open(path, encoding="utf-8") as fh: for line in text.splitlines():
for line in fh: line = line.strip()
line = line.strip() if not line:
if not line: continue
continue try:
try: obj = json.loads(line)
obj = json.loads(line) except json.JSONDecodeError:
except json.JSONDecodeError: continue
continue if isinstance(obj, dict):
samples.append(obj) samples.append(obj)
return samples return samples
@@ -115,8 +127,6 @@ def _compute_cosine(references: list[str], predictions: list[str]) -> dict[str,
from sklearn.metrics.pairwise import cosine_similarity from sklearn.metrics.pairwise import cosine_similarity
except ImportError: except ImportError:
return {"enabled": False, "error": "scikit-learn not installed", "score": 0} return {"enabled": False, "error": "scikit-learn not installed", "score": 0}
if len(predictions) < 2:
return {"enabled": True, "score": 0, "error": "need at least 2 samples for corpus cosine"}
try: try:
vectorizer = TfidfVectorizer() vectorizer = TfidfVectorizer()
tfidf = vectorizer.fit_transform(references + predictions) tfidf = vectorizer.fit_transform(references + predictions)
@@ -129,6 +139,32 @@ def _compute_cosine(references: list[str], predictions: list[str]) -> dict[str,
return {"enabled": True, "score": 0, "error": "insufficient text for vectorization"} 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 # LLM Judge
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -265,7 +301,7 @@ def run_eval(config: dict[str, Any]) -> dict[str, Any]:
# ---- 1. Load dataset ---- # ---- 1. Load dataset ----
print(f"[eval] loading dataset: {dataset_path}") print(f"[eval] loading dataset: {dataset_path}")
raw_samples = _load_jsonl(dataset_path) raw_samples = _load_dataset(dataset_path)
print(f"[eval] loaded {len(raw_samples)} samples") print(f"[eval] loaded {len(raw_samples)} samples")
# ---- 2. Load model ---- # ---- 2. Load model ----
@@ -356,6 +392,8 @@ def run_eval(config: dict[str, Any]) -> dict[str, Any]:
cosine_cfg = basic_cfg.get("cosine", {}) cosine_cfg = basic_cfg.get("cosine", {})
if cosine_cfg.get("enabled"): if cosine_cfg.get("enabled"):
metrics_result["cosine"] = _compute_cosine(references, predictions) 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 ---- # ---- 5. Summarise ----
completed = len(samples) completed = len(samples)
@@ -375,9 +413,23 @@ def run_eval(config: dict[str, Any]) -> dict[str, Any]:
overall_evaluation = f"评测完成:{completed} 样本,{passed_count} 通过,平均 {avg_score}/{max_score}" overall_evaluation = f"评测完成:{completed} 样本,{passed_count} 通过,平均 {avg_score}/{max_score}"
else: else:
passed_count = 0 passed_count = 0
overall_score = 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 overall_score_max = 100
dimension_summary = [] 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 评委)" overall_evaluation = f"评测完成:{completed} 样本(未配置 LLM 评委)"
result = { result = {

View File

@@ -1,4 +1,4 @@
import { get, post, put } from '../request' import { del, get, post, put } from '../request'
export interface ComputeNode { export interface ComputeNode {
id: string id: string
@@ -95,6 +95,9 @@ export const createComputeNode = (data: ComputeNodePayload) =>
export const updateComputeNode = (id: string, data: Partial<ComputeNode>) => export const updateComputeNode = (id: string, data: Partial<ComputeNode>) =>
put<ComputeNode>(`/compute/nodes/${id}`, data) put<ComputeNode>(`/compute/nodes/${id}`, data)
export const deleteComputeNode = (id: string) =>
del<{ deleted: string }>(`/compute/nodes/${id}`)
export const testComputeNode = (id: string) => export const testComputeNode = (id: string) =>
post<{ node_id: string; success: boolean; latency_ms: number; gpu_count: number; error?: string }>(`/compute/nodes/${id}/test-connection`) post<{ node_id: string; success: boolean; latency_ms: number; gpu_count: number; error?: string }>(`/compute/nodes/${id}/test-connection`)

View File

@@ -29,6 +29,7 @@ export const uploadDatasetFiles = (datasetId: string | number, files: File[]) =>
files.forEach((f) => formData.append('files', f)) files.forEach((f) => formData.append('files', f))
return post(`/dataset-manage/upload/${datasetId}`, formData, { return post(`/dataset-manage/upload/${datasetId}`, formData, {
headers: { 'Content-Type': 'multipart/form-data' }, headers: { 'Content-Type': 'multipart/form-data' },
timeout: 120000,
}) })
} }

View File

@@ -132,6 +132,7 @@ export interface FineTuneTask {
train_dataset_id?: number | string train_dataset_id?: number | string
auto_merge?: boolean auto_merge?: boolean
output_model_name?: string output_model_name?: string
compute_node_id?: string
gpus?: number[] gpus?: number[]
batch_size?: number batch_size?: number
learning_rate?: number learning_rate?: number
@@ -355,7 +356,7 @@ export interface GpuInfo {
power_w: number power_w: number
id?: number id?: number
uuid?: string uuid?: string
status?: 'idle' | 'busy' | 'warning' | 'offline' status?: 'idle' | 'busy' | 'reserved' | 'warning' | 'offline'
memory_percent?: number memory_percent?: number
power_limit_w?: number power_limit_w?: number
processes?: GpuProcess[] processes?: GpuProcess[]

View File

@@ -1,12 +1,12 @@
<script setup lang="ts"> <script setup lang="ts">
import { computed, onMounted, onUnmounted, reactive, ref, watch } from 'vue' import { computed, onMounted, onUnmounted, reactive, ref, watch } from 'vue'
import { useRoute, useRouter } from 'vue-router' import { useRoute, useRouter } from 'vue-router'
import { ElMessage } from 'element-plus' import { ElMessage, ElMessageBox } from 'element-plus'
import { import {
checkNodeReplicaDrift, checkNodeReplicaDrift,
createComputeNode, createComputeNode,
deleteComputeNode,
disableComputeNode, disableComputeNode,
drainComputeNode,
enableComputeNode, enableComputeNode,
getComputeGpus, getComputeGpus,
getComputeNodes, getComputeNodes,
@@ -129,11 +129,10 @@ async function changeTab(name: string | number) {
await router.replace({ path: '/compute', query: { tab: String(name) } }) await router.replace({ path: '/compute', query: { tab: String(name) } })
} }
async function handleNodeAction(action: 'enable' | 'disable' | 'drain' | 'test', node: ComputeNode) { async function handleNodeAction(action: 'enable' | 'disable' | 'test', node: ComputeNode) {
const nodeId = String(node.id) const nodeId = String(node.id)
if (action === 'enable') await enableComputeNode(nodeId) if (action === 'enable') await enableComputeNode(nodeId)
if (action === 'disable') await disableComputeNode(nodeId) if (action === 'disable') await disableComputeNode(nodeId)
if (action === 'drain') await drainComputeNode(nodeId)
if (action === 'test') { if (action === 'test') {
const result = await testComputeNode(nodeId) const result = await testComputeNode(nodeId)
if (result.success) { if (result.success) {
@@ -145,6 +144,27 @@ async function handleNodeAction(action: 'enable' | 'disable' | 'drain' | 'test',
await load() await load()
} }
async function handleDeleteNode(node: ComputeNode) {
try {
await ElMessageBox.confirm(
`确定删除算力节点「${node.name || node.code}」吗?节点删除后,其 GPU 设备和资源副本记录也会一并移除。`,
'删除算力节点',
{ type: 'warning', confirmButtonText: '删除', cancelButtonText: '取消' },
)
} catch {
return
}
try {
await deleteComputeNode(String(node.id))
ElMessage.success('算力节点已删除')
if (selectedNodeId.value === node.id) selectedNodeId.value = ''
await load({ showButtonLoading: true })
} catch (err: any) {
const message = err?.response?.data?.detail?.message || err?.response?.data?.message || '删除算力节点失败'
ElMessage.error(message)
}
}
async function handleReplicaDriftCheck() { async function handleReplicaDriftCheck() {
if (!selectedNodeId.value) return if (!selectedNodeId.value) return
checkingReplicas.value = true checkingReplicas.value = true
@@ -355,7 +375,7 @@ onUnmounted(() => {
<el-button size="small" @click="handleNodeAction('test', asComputeNode(row))">测试</el-button> <el-button size="small" @click="handleNodeAction('test', asComputeNode(row))">测试</el-button>
<el-button v-if="row.enabled" size="small" @click="handleNodeAction('disable', asComputeNode(row))">停用</el-button> <el-button v-if="row.enabled" size="small" @click="handleNodeAction('disable', asComputeNode(row))">停用</el-button>
<el-button v-else size="small" type="primary" @click="handleNodeAction('enable', asComputeNode(row))">启用</el-button> <el-button v-else size="small" type="primary" @click="handleNodeAction('enable', asComputeNode(row))">启用</el-button>
<el-button size="small" type="warning" plain @click="handleNodeAction('drain', asComputeNode(row))">维护</el-button> <el-button size="small" type="danger" plain @click="handleDeleteNode(asComputeNode(row))">删除</el-button>
</template> </template>
</el-table-column> </el-table-column>
</el-table> </el-table>

View File

@@ -50,6 +50,20 @@ const rules: FormRules = {
} }
/** 处理文件选择(替换模式:新文件覆盖旧文件) */ /** 处理文件选择(替换模式:新文件覆盖旧文件) */
function parseDatasetRecordValues(text: string, fileName: string): unknown[] {
const content = text.trim()
if (!content) return []
if (fileName.toLowerCase().endsWith('.json')) {
const parsed = JSON.parse(content)
return Array.isArray(parsed) ? parsed : [parsed]
}
return content
.split('\n')
.map((line) => line.trim())
.filter(Boolean)
.map((line) => JSON.parse(line))
}
async function handleFileChange(uploadFile: UploadFile) { async function handleFileChange(uploadFile: UploadFile) {
const raw = uploadFile.raw const raw = uploadFile.raw
if (!raw) return if (!raw) return
@@ -70,25 +84,19 @@ async function handleFileChange(uploadFile: UploadFile) {
async function analyzeFile(file: File) { async function analyzeFile(file: File) {
try { try {
const text = await file.text() const text = await file.text()
const lines = text.trim().split('\n').filter(Boolean) const records = parseDatasetRecordValues(text, file.name)
fileCount.value = lines.length fileCount.value = records.length
// Alpaca 格式校验:每行 JSON 须含 instruction 字段 // Alpaca 格式校验:每行 JSON 须含 instruction 字段
let validCount = 0 const validCount = records.filter(
for (const line of lines) { (obj) => obj && typeof obj === 'object' && 'instruction' in obj,
try { ).length
const obj = JSON.parse(line) if (validCount > 0 && validCount === records.length) {
if (obj.instruction !== undefined) validCount++
} catch {
// 非 JSON 行(如纯 JSONL 多行结构)
}
}
if (validCount > 0 && validCount === lines.length) {
formatValid.value = true formatValid.value = true
formatMessage.value = `符合 Alpaca 格式(含 instruction 字段)` formatMessage.value = `符合 Alpaca 格式(含 instruction 字段)`
} else if (validCount > 0) { } else if (validCount > 0) {
formatValid.value = true formatValid.value = true
formatMessage.value = `部分符合 Alpaca 格式(${validCount}/${lines.length}` formatMessage.value = `部分符合 Alpaca 格式(${validCount}/${records.length}`
} else { } else {
formatValid.value = false formatValid.value = false
formatMessage.value = '未检测到标准 Alpaca 格式(缺少 instruction 字段),仍可上传' formatMessage.value = '未检测到标准 Alpaca 格式(缺少 instruction 字段),仍可上传'

View File

@@ -1,9 +1,10 @@
<script setup lang="ts"> <script setup lang="ts">
import { computed, onMounted, ref } from 'vue' import { computed, onMounted, onUnmounted, ref } from 'vue'
import { useRoute } from 'vue-router' import { useRoute } from 'vue-router'
import PageCard from '@/components/PageCard.vue' import PageCard from '@/components/PageCard.vue'
import ModelStatusTag from '@/components/ModelStatusTag.vue' import ModelStatusTag from '@/components/ModelStatusTag.vue'
import { getEvalDetail } from '@/api/modules/eval' import { getEvalDetail } from '@/api/modules/eval'
import { usePolling } from '@/composables/usePolling'
import type { EvalSampleResult, EvalTaskDetail } from '@/types' import type { EvalSampleResult, EvalTaskDetail } from '@/types'
const route = useRoute() const route = useRoute()
@@ -16,6 +17,7 @@ const keyword = ref('')
const judgementFilter = ref('') const judgementFilter = ref('')
const currentPage = ref(1) const currentPage = ref(1)
const pageSize = ref(10) const pageSize = ref(10)
const ACTIVE_STATUSES = new Set(['pending', 'queued', 'running'])
const filteredSamples = computed(() => { const filteredSamples = computed(() => {
const normalizedKeyword = keyword.value.trim().toLowerCase() const normalizedKeyword = keyword.value.trim().toLowerCase()
@@ -74,8 +76,8 @@ function resetPage() {
currentPage.value = 1 currentPage.value = 1
} }
async function loadDetail() { async function loadDetail(options: { silent?: boolean } = {}) {
loading.value = true if (!options.silent) loading.value = true
loadError.value = '' loadError.value = ''
try { try {
detail.value = await getEvalDetail(taskId) detail.value = await getEvalDetail(taskId)
@@ -83,11 +85,29 @@ async function loadDetail() {
detail.value = null detail.value = null
loadError.value = '评测详情加载失败,请稍后重试。' loadError.value = '评测详情加载失败,请稍后重试。'
} finally { } finally {
loading.value = false if (!options.silent) loading.value = false
} }
} }
onMounted(loadDetail) const { start: startPolling, stop: stopPolling } = usePolling(
async () => {
await loadDetail({ silent: true })
if (!ACTIVE_STATUSES.has(String(detail.value?.status || ''))) {
stopPolling()
}
},
5000,
{ immediate: false },
)
onMounted(async () => {
await loadDetail()
if (ACTIVE_STATUSES.has(String(detail.value?.status || ''))) {
startPolling()
}
})
onUnmounted(stopPolling)
</script> </script>
<template> <template>
@@ -113,7 +133,7 @@ onMounted(loadDetail)
<i class="fa fa-exclamation-circle" aria-hidden="true" /> <i class="fa fa-exclamation-circle" aria-hidden="true" />
<h2>无法加载评测详情</h2> <h2>无法加载评测详情</h2>
<p>{{ loadError }}</p> <p>{{ loadError }}</p>
<el-button type="primary" @click="loadDetail">重新加载</el-button> <el-button type="primary" @click="() => loadDetail()">重新加载</el-button>
</div> </div>
<template v-else-if="detail"> <template v-else-if="detail">

View File

@@ -24,14 +24,16 @@ const leaderboard = ref([
{ rank: 3, name: 'Qwen-Max', score: 85.3 }, { rank: 3, name: 'Qwen-Max', score: 85.3 },
]) ])
async function loadEvalList() { const ACTIVE_STATUSES = new Set(['pending', 'queued', 'running'])
evalLoading.value = true
async function loadEvalList(options: { silent?: boolean } = {}) {
if (!options.silent) evalLoading.value = true
try { try {
evalList.value = (await getEvalList()) || [] evalList.value = (await getEvalList()) || []
} catch { } catch {
evalList.value = [] evalList.value = []
} finally { } finally {
evalLoading.value = false if (!options.silent) evalLoading.value = false
} }
} }
@@ -56,14 +58,21 @@ function handleViewDetail(row: any) {
} }
const { start: startPolling, stop: stopPolling } = usePolling( const { start: startPolling, stop: stopPolling } = usePolling(
() => loadEvalList(), async () => {
await loadEvalList({ silent: true })
if (!evalList.value.some((item) => ACTIVE_STATUSES.has(String(item.status || '')))) {
stopPolling()
}
},
5000, 5000,
{ immediate: false }, { immediate: false },
) )
onMounted(async () => { onMounted(async () => {
await loadEvalList() await loadEvalList()
startPolling() if (evalList.value.some((item) => ACTIVE_STATUSES.has(String(item.status || '')))) {
startPolling()
}
}) })
onUnmounted(() => { onUnmounted(() => {

View File

@@ -37,7 +37,7 @@ const models = ref<ModelItem[]>([])
const datasets = ref<DatasetItem[]>([]) const datasets = ref<DatasetItem[]>([])
const gpus = ref<GpuInfo[]>([]) const gpus = ref<GpuInfo[]>([])
const computeNodes = ref<Array<{ id: string; scheduler_status?: string }>>([]) const computeNodes = ref<Array<{ id: string; scheduler_status?: string }>>([])
const selectedGpuId = ref<number | null>(null) const selectedGpuKeys = ref<string[]>([])
/** Only show GPUs from nodes that are online or draining */ /** Only show GPUs from nodes that are online or draining */
const availableGpus = computed(() => { const availableGpus = computed(() => {
@@ -74,7 +74,13 @@ const selectedModel = computed(() => models.value.find((model) => model.id === f
const modelDialogTitle = computed(() => selectedModel.value?.name || '') const modelDialogTitle = computed(() => selectedModel.value?.name || '')
/** 训练命令与提交载荷共用同一份表单模型。 */ /** 训练命令与提交载荷共用同一份表单模型。 */
const selectedGpuIds = computed(() => (selectedGpuId.value != null ? [selectedGpuId.value] : [])) const selectedGpus = computed(() =>
selectedGpuKeys.value
.map((key) => availableGpus.value.find((gpu) => gpuKey(gpu) === key))
.filter((gpu): gpu is GpuInfo => Boolean(gpu)),
)
const selectedComputeNodeId = computed(() => selectedGpus.value[0]?.node_id)
const selectedGpuIds = computed(() => selectedGpus.value.map((gpu) => Number(gpu.id)))
const commandPreview = computed(() => buildFineTuneCommand(form, selectedGpuIds.value)) const commandPreview = computed(() => buildFineTuneCommand(form, selectedGpuIds.value))
const remoteCommandPreview = computed(() => { const remoteCommandPreview = computed(() => {
@@ -83,9 +89,32 @@ const remoteCommandPreview = computed(() => {
return preflightResult.value?.preview?.command_text || '' return preflightResult.value?.preview?.command_text || ''
}) })
/** GPU 单选切换(每次只选中一张 GPU */ function gpuKey(gpu: GpuInfo) {
function toggleGpu(gpuId: number) { return `${gpu.node_id || 'local'}:${gpu.id ?? gpu.uuid ?? gpu.name}`
selectedGpuId.value = selectedGpuId.value === gpuId ? null : gpuId }
function isGpuUnavailable(gpu: GpuInfo) {
return gpu.status === 'busy' || gpu.status === 'reserved' || gpu.status === 'offline'
}
function isGpuSelected(gpu: GpuInfo) {
return selectedGpuKeys.value.includes(gpuKey(gpu))
}
/** GPU 多选切换:单个任务只允许选择同一算力节点内的空闲卡。 */
function toggleGpu(gpu: GpuInfo) {
if (isGpuUnavailable(gpu) || gpu.id == null) return
const key = gpuKey(gpu)
if (isGpuSelected(gpu)) {
selectedGpuKeys.value = selectedGpuKeys.value.filter((item) => item !== key)
return
}
if (selectedComputeNodeId.value && gpu.node_id && selectedComputeNodeId.value !== gpu.node_id) {
selectedGpuKeys.value = [key]
ElMessage.info('已切换到新的算力节点,之前选择的 GPU 已清空')
return
}
selectedGpuKeys.value = [...selectedGpuKeys.value, key]
} }
function gpuUsageWidth(percent: number) { function gpuUsageWidth(percent: number) {
@@ -178,8 +207,8 @@ async function loadGpus() {
const [sys, nodes] = await Promise.all([getSystemInfo(), getComputeNodes().catch(() => [])]) const [sys, nodes] = await Promise.all([getSystemInfo(), getComputeNodes().catch(() => [])])
gpus.value = sys?.gpu || [] gpus.value = sys?.gpu || []
computeNodes.value = nodes || [] computeNodes.value = nodes || []
// Default select first available GPU const firstIdle = availableGpus.value.find((gpu) => !isGpuUnavailable(gpu) && gpu.id != null)
if (availableGpus.value.length > 0) selectedGpuId.value = availableGpus.value[0].id ?? null if (firstIdle) selectedGpuKeys.value = [gpuKey(firstIdle)]
} catch { } catch {
gpus.value = [] gpus.value = []
} }
@@ -189,8 +218,8 @@ async function handleSubmit() {
if (!formRef.value) return if (!formRef.value) return
await formRef.value.validate(async (valid) => { await formRef.value.validate(async (valid) => {
if (!valid) return if (!valid) return
if (selectedGpuId.value == null) { if (!selectedGpuIds.value.length) {
ElMessage.warning('请选择一 GPU') ElMessage.warning('请至少选择一张空闲 GPU')
return return
} }
submitting.value = true submitting.value = true
@@ -207,7 +236,7 @@ async function handleSubmit() {
return return
} }
const payload = buildFineTunePayload(form, selectedGpuIds.value) const payload = buildFineTunePayload(form, selectedGpuIds.value, selectedComputeNodeId.value)
const preflight = await runPreflight(payload) const preflight = await runPreflight(payload)
if (!preflight?.valid) { if (!preflight?.valid) {
ElMessage.error('训练预检未通过,请先处理预检问题') ElMessage.error('训练预检未通过,请先处理预检问题')
@@ -232,7 +261,7 @@ async function handleSubmit() {
}) })
} }
async function runPreflight(payload = buildFineTunePayload(form, selectedGpuIds.value)) { async function runPreflight(payload = buildFineTunePayload(form, selectedGpuIds.value, selectedComputeNodeId.value)) {
preflightLoading.value = true preflightLoading.value = true
try { try {
const result = await preflightFineTune(payload) const result = await preflightFineTune(payload)
@@ -261,8 +290,8 @@ async function handlePreflightClick() {
if (!formRef.value) return if (!formRef.value) return
await formRef.value.validate(async (valid) => { await formRef.value.validate(async (valid) => {
if (!valid) return if (!valid) return
if (selectedGpuId.value == null) { if (!selectedGpuIds.value.length) {
ElMessage.warning('请选择一 GPU') ElMessage.warning('请至少选择一张空闲 GPU')
return return
} }
await runPreflight() await runPreflight()
@@ -297,12 +326,16 @@ onMounted(() => {
<el-divider content-position="left">训练配置</el-divider> <el-divider content-position="left">训练配置</el-divider>
<el-form-item label="GPU 硬件"> <el-form-item label="GPU 硬件">
<div class="gpu-list"> <div class="gpu-list">
<div class="gpu-selection-summary">
已选择 {{ selectedGpuIds.length }} GPU
<template v-if="selectedGpus[0]?.node_code"> · {{ selectedGpus[0].node_code }}</template>
</div>
<div <div
v-for="gpu in availableGpus" v-for="gpu in availableGpus"
:key="gpu.id" :key="gpuKey(gpu)"
class="gpu-card" class="gpu-card"
:class="{ active: selectedGpuId === gpu.id, 'is-busy': gpu.gpu_percent > 80 }" :class="{ active: isGpuSelected(gpu), 'is-busy': isGpuUnavailable(gpu), 'is-disabled': isGpuUnavailable(gpu) }"
@click="toggleGpu(gpu.id!)" @click="toggleGpu(gpu)"
> >
<div class="gpu-card-top"> <div class="gpu-card-top">
<div class="gpu-title"> <div class="gpu-title">
@@ -312,7 +345,7 @@ onMounted(() => {
</span> </span>
<span class="gpu-name">{{ gpu.name }}</span> <span class="gpu-name">{{ gpu.name }}</span>
</div> </div>
<span class="gpu-usage">{{ gpu.gpu_percent }}%</span> <span class="gpu-usage">{{ isGpuUnavailable(gpu) ? gpu.status : `${gpu.gpu_percent}%` }}</span>
</div> </div>
<div class="gpu-usage-bar"> <div class="gpu-usage-bar">
<span :style="{ width: gpuUsageWidth(gpu.gpu_percent) }" /> <span :style="{ width: gpuUsageWidth(gpu.gpu_percent) }" />
@@ -578,6 +611,13 @@ onMounted(() => {
width: 100%; width: 100%;
} }
.gpu-selection-summary {
grid-column: 1 / -1;
color: #64748b;
font-size: 12px;
line-height: 20px;
}
.gpu-card { .gpu-card {
border: 1px solid #e5e7eb; border: 1px solid #e5e7eb;
border-radius: 6px; border-radius: 6px;
@@ -627,6 +667,11 @@ onMounted(() => {
background: #dc2626; background: #dc2626;
} }
} }
&.is-disabled {
cursor: not-allowed;
opacity: 0.72;
}
} }
.gpu-card-top { .gpu-card-top {

View File

@@ -65,6 +65,7 @@ export function createDefaultFineTuneForm(): FineTuneFormModel {
export function buildFineTunePayload( export function buildFineTunePayload(
form: FineTuneFormModel, form: FineTuneFormModel,
gpus: number[], gpus: number[],
computeNodeId?: string,
): Omit<FineTuneStartPayload, 'task_id'> { ): Omit<FineTuneStartPayload, 'task_id'> {
return { return {
name: form.name, name: form.name,
@@ -77,6 +78,7 @@ export function buildFineTunePayload(
train_dataset_id: form.train_dataset_id, train_dataset_id: form.train_dataset_id,
auto_merge: form.train_type === 'SFT' && form.auto_merge, auto_merge: form.train_type === 'SFT' && form.auto_merge,
output_model_name: form.name, output_model_name: form.name,
compute_node_id: computeNodeId,
batch_size: form.batch_size, batch_size: form.batch_size,
learning_rate: form.learning_rate, learning_rate: form.learning_rate,
n_epochs: form.n_epochs, n_epochs: form.n_epochs,