feat: 新增 compute_gateway、compute_poller、agent 模块,重构前端 dist

- 新增 backend/app/modules/compute_gateway(client/sync)计算网关模块
- 新增 backend/app/workers/compute_poller 计算轮询 worker
- 新增 compute/agent/process_manager 进程管理器
- 新增 scripts/ 脚本目录
- 更新 Docker 部署配置(app/compute/nginx)
- 更新后端平台 API、数据库 SQL、core 配置
- 更新前端多个视图组件及 API 模块
- 重构 frontend/dist 构建产物(新 hash)
- 更新多项文档

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
wuyongtao
2026-07-22 17:32:59 +08:00
parent 1e438164c1
commit 836343b29e
180 changed files with 3352 additions and 263 deletions

View File

@@ -25,5 +25,57 @@ compute/
## 运行模式
- 默认 `COMPUTE_EXECUTION_MODE=real`Compute API 只暴露健康检查和接口契约;真实训练执行器完成前,创建作业会返回未实现错误
- 默认 `COMPUTE_EXECUTION_MODE=real`Compute API 会通过 `compute.agent.process_manager.ProcessManager` 启动真实 `llamafactory-cli train` 子进程,并将日志写入 `TRAINING_LOG_ROOT`
- 真实模式下 GPU 发现优先使用宿主机 `nvidia-smi`。如果部署环境暂时无法调用 `nvidia-smi`,可通过 `COMPUTE_GPU_COUNT``COMPUTE_GPU_NAME``COMPUTE_GPU_MEMORY_GB``COMPUTE_GPU_POWER_LIMIT_W` 声明兼容 GPU 清单,便于应用侧先完成节点登记和联调。
- 仅隔离联调时可设置 `COMPUTE_EXECUTION_MODE=simulator`,启用内存状态机和合成 GPU/日志数据。该模式不得作为生产运行路径。
- 服务间鉴权默认开启:设置 `COMPUTE_AUTH_ENABLED=true` 和一致的 `COMPUTE_SERVICE_TOKEN`,应用侧会通过 `X-Compute-Token` 调用 Compute API。
- 真实训练作业会登记到 `TRAINING_LOG_ROOT/compute-jobs.json`。Compute API 重启后会恢复作业索引,继续提供状态、停止和日志查询。
- 同一算力节点内按 GPU ID 做轻量锁定;已有运行中作业占用的 GPU 不允许再次提交,避免同机多 GPU 场景下误复用。
真实执行前提:
- 镜像或宿主机环境中 `llamafactory-cli` 可执行。
- `LLAMA_FACTORY_HOME` 指向 LLaMA-Factory 工作目录。
- 基座模型路径和数据集名称/目录已经在算力服务器本地可访问。
- 应用侧训练任务中的 GPU、模型、数据集配置能映射到当前节点本地路径。
## 应用侧接入
应用平台通过“算力节点”页面维护每台 GPU 服务器的 `Compute API``File Gateway` 地址。点击连接测试时Backend API 会主动调用:
```text
GET /modelTF/v1/compute/health
GET /modelTF/compute/resources/gpus
```
连接成功后,应用侧会同步节点健康信息、能力标签和 GPU 清单到 PostgreSQL。多节点阶段仍按“每台算力服务器 = 单机多 GPU 节点”管理,每台服务器都部署 Compute API、Agent、File Gateway 契约和 LLaMA-Factory。
训练闭环:
```text
Frontend 创建/启动训练
-> Backend API 选择 compute_nodes 节点
-> Backend API POST /modelTF/compute/jobs 到目标 Compute API
-> Compute API 启动 llamafactory-cli 子进程
-> Backend Worker 定时 GET /modelTF/compute/jobs/{id}
-> Backend API 同步 fine_tune_tasks 状态、进度、PID、日志路径和产物索引
```
## 当前接口能力
日志接口:
```text
GET /modelTF/compute/jobs/{job_id}/logs?tail_lines=200
GET /modelTF/compute/jobs/{job_id}/logs?offset=0&limit=500
```
返回 `content``metrics``total_lines``offset``limit``has_more``next_offset`,用于前端增量刷新和日志平台采集。
文件导入:
```text
POST /modelTF/compute/files/import-local
```
该接口用于应用侧调度前把算力服务器本地可访问的模型/数据集路径导入到 `YG_FT_DATA_ROOT` 内部。目标路径会校验不能逃逸出 `YG_FT_DATA_ROOT`,源路径必须已存在于算力服务器本地或挂载目录。

View File

@@ -0,0 +1,246 @@
from __future__ import annotations
import os
import json
import contextlib
import signal
import subprocess
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
TERMINAL_STATUSES = {"completed", "failed", "stopped"}
@dataclass
class ManagedProcess:
id: str
name: str
command: list[str]
work_dir: str
log_path: Path
output_dir: str
gpus: list[int]
process: subprocess.Popen[Any] | None
created_at: float
pid: int | None = None
status: str = "running"
progress: int = 5
artifacts: list[dict[str, Any]] = field(default_factory=list)
class ProcessManager:
def __init__(self, log_root: str) -> None:
self.log_root = Path(log_root)
self.log_root.mkdir(parents=True, exist_ok=True)
self.registry_path = self.log_root / "compute-jobs.json"
self.jobs: dict[str, ManagedProcess] = {}
self._load_registry()
def create_job(self, payload: dict[str, Any], command: list[str], work_dir: str) -> dict[str, Any]:
job_id = str(payload.get("id") or f"job_{int(time.time() * 1000)}")
if job_id in self.jobs and self.jobs[job_id].status not in TERMINAL_STATUSES:
raise ValueError(f"job {job_id} is already running")
output_dir = str(payload.get("output_dir") or f"/data/yg-ft/outputs/{payload.get('name', job_id)}")
Path(output_dir).mkdir(parents=True, exist_ok=True)
log_path = self.log_root / f"{job_id}.log"
env = os.environ.copy()
gpus = [int(item) for item in payload.get("gpus") or []]
locked = self.locked_gpus()
conflict = sorted(set(gpus).intersection(locked))
if conflict:
raise ValueError(f"gpu already locked: {conflict}")
if gpus:
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(item) for item in gpus)
env.update({str(k): str(v) for k, v in payload.get("env", {}).items()})
cwd = work_dir if Path(work_dir).exists() else None
with log_path.open("ab") as log_file:
log_file.write(f"[INFO] starting job_id={job_id} command={' '.join(command)}\n".encode("utf-8"))
process = subprocess.Popen(
command,
cwd=cwd,
env=env,
stdout=log_file,
stderr=subprocess.STDOUT,
)
managed = ManagedProcess(
id=job_id,
name=str(payload.get("name") or job_id),
command=command,
work_dir=work_dir,
log_path=log_path,
output_dir=output_dir,
gpus=gpus,
process=process,
created_at=time.time(),
pid=process.pid,
progress=10,
)
self.jobs[job_id] = managed
data = self.serialize(managed)
self._save_registry()
return data
def get_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
return self.serialize(job)
def list_jobs(self) -> list[dict[str, Any]]:
return [self.serialize(job) for job in self.jobs.values()]
def stop_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
if job.status not in TERMINAL_STATUSES:
try:
if job.process is not None and os.name == "nt":
job.process.terminate()
elif job.pid is not None:
os.kill(job.pid, signal.SIGTERM)
if job.process is not None:
job.process.wait(timeout=10)
except Exception:
if job.process is not None:
job.process.kill()
elif job.pid is not None:
with contextlib.suppress(Exception):
os.kill(job.pid, signal.SIGKILL)
job.status = "stopped"
job.progress = min(job.progress, 99)
data = self.serialize(job)
self._save_registry()
return data
def logs(self, job_id: str) -> str:
job = self.jobs.get(job_id)
if not job or not job.log_path.exists():
return ""
return job.log_path.read_text(encoding="utf-8", errors="replace")
def serialize(self, job: ManagedProcess) -> dict[str, Any]:
code = job.process.poll() if job.process is not None else None
if job.status not in TERMINAL_STATUSES:
if job.process is None and job.pid is not None and not self._pid_alive(job.pid):
job.status = "failed"
job.progress = min(job.progress, 99)
code = -1
elif code is None:
job.status = "running"
elapsed = max(0, int(time.time() - job.created_at))
job.progress = min(95, max(job.progress, 10 + elapsed // 6))
elif code == 0:
job.status = "completed"
job.progress = 100
job.artifacts = self._collect_artifacts(job.output_dir)
else:
job.status = "failed"
job.progress = min(job.progress, 99)
self._save_registry()
return {
"id": job.id,
"name": job.name,
"status": job.status,
"progress": job.progress,
"pid": job.pid,
"gpus": job.gpus,
"created_at": job.created_at,
"command": job.command,
"work_dir": job.work_dir,
"output_dir": job.output_dir,
"log_file": str(job.log_path),
"artifacts": job.artifacts,
"return_code": code,
}
def locked_gpus(self) -> set[int]:
locked: set[int] = set()
for job in self.jobs.values():
status = self.serialize(job)["status"]
if status in {"queued", "running"}:
locked.update(job.gpus)
return locked
def _collect_artifacts(self, output_dir: str) -> list[dict[str, Any]]:
root = Path(output_dir)
if not root.exists():
return []
artifacts: list[dict[str, Any]] = []
for path in root.rglob("*"):
if path.is_file():
artifacts.append(
{
"path": str(path),
"name": path.name,
"size": path.stat().st_size,
}
)
return artifacts[:200]
def _save_registry(self) -> None:
items = []
for job in self.jobs.values():
items.append(
{
"id": job.id,
"name": job.name,
"command": job.command,
"work_dir": job.work_dir,
"log_path": str(job.log_path),
"output_dir": job.output_dir,
"gpus": job.gpus,
"pid": job.pid,
"created_at": job.created_at,
"status": job.status,
"progress": job.progress,
"artifacts": job.artifacts,
}
)
self.registry_path.write_text(json.dumps(items, ensure_ascii=False, indent=2), encoding="utf-8")
def _load_registry(self) -> None:
if not self.registry_path.exists():
return
try:
items = json.loads(self.registry_path.read_text(encoding="utf-8"))
except json.JSONDecodeError:
return
for item in items if isinstance(items, list) else []:
if not isinstance(item, dict):
continue
pid = item.get("pid")
status = item.get("status", "failed")
if status not in TERMINAL_STATUSES and pid and not self._pid_alive(int(pid)):
status = "failed"
job = ManagedProcess(
id=str(item["id"]),
name=str(item.get("name") or item["id"]),
command=[str(part) for part in item.get("command") or []],
work_dir=str(item.get("work_dir") or ""),
log_path=Path(item.get("log_path") or self.log_root / f"{item['id']}.log"),
output_dir=str(item.get("output_dir") or ""),
gpus=[int(gpu) for gpu in item.get("gpus") or []],
process=None,
pid=int(pid) if pid else None,
created_at=float(item.get("created_at") or time.time()),
status=status,
progress=int(item.get("progress") or 0),
artifacts=item.get("artifacts") or [],
)
self.jobs[job.id] = job
def _pid_alive(self, pid: int) -> bool:
if pid <= 0:
return False
try:
os.kill(pid, 0)
return True
except OSError:
return False

View File

@@ -2,12 +2,17 @@ from __future__ import annotations
import os
import math
import hashlib
import shutil
import subprocess
import time
from pathlib import Path
from typing import Any
from fastapi import FastAPI, HTTPException
from fastapi import FastAPI, File, Form, HTTPException, Query, Request, UploadFile
from fastapi.responses import FileResponse, JSONResponse
from compute.agent.process_manager import ProcessManager
from compute.engines.llama_factory.adapter import build_command, parse_log_line
@@ -15,6 +20,20 @@ def create_app() -> FastAPI:
app = FastAPI(title="YG Fine-Tune Compute API")
jobs: dict[str, dict[str, Any]] = {}
route_prefix = os.getenv("MODELTF_ROUTE_PREFIX", "/modelTF").rstrip("/") or "/modelTF"
process_manager = ProcessManager(os.getenv("TRAINING_LOG_ROOT", "/opt/yg-ft/logs/training"))
@app.middleware("http")
async def compute_token_auth(request: Request, call_next):
token = os.getenv("COMPUTE_SERVICE_TOKEN", "")
auth_enabled = os.getenv("COMPUTE_AUTH_ENABLED", "true").lower() == "true"
public_paths = {f"{route_prefix}/health", "/health"}
if auth_enabled and token and request.url.path not in public_paths:
header_token = request.headers.get("x-compute-token", "")
auth_header = request.headers.get("authorization", "")
bearer_token = auth_header.removeprefix("Bearer ").strip() if auth_header.startswith("Bearer ") else ""
if header_token != token and bearer_token != token:
return JSONResponse({"detail": "invalid compute service token"}, status_code=401)
return await call_next(request)
def now() -> float:
return time.time()
@@ -25,6 +44,68 @@ def create_app() -> FastAPI:
def execution_mode() -> str:
return os.getenv("COMPUTE_EXECUTION_MODE", os.getenv("COMPUTE_MODE", "real")).lower()
def _int_env(name: str, default: int) -> int:
raw = os.getenv(name)
if raw is None or raw == "":
return default
return int(raw)
def _float_env(name: str, default: float) -> float:
raw = os.getenv(name)
if raw is None or raw == "":
return default
return float(raw)
def _path_inside(root: Path, candidate: Path) -> bool:
try:
candidate.resolve().relative_to(root.resolve())
return True
except ValueError:
return False
def _llama_factory_version() -> str:
for command in (["llamafactory-cli", "version"], ["llamafactory-cli", "--version"]):
try:
result = subprocess.run(command, capture_output=True, text=True, timeout=5)
except Exception:
continue
output = (result.stdout or result.stderr).strip()
if result.returncode == 0 and output:
return output.splitlines()[0][:120]
return ""
def _slice_log_content(
content: str,
tail_lines: int | None = None,
offset: int | None = None,
limit: int | None = None,
) -> dict[str, Any]:
lines = content.splitlines()
total = len(lines)
if offset is not None or limit is not None:
start = max(0, offset or 0)
end = start + limit if limit else total
selected = lines[start:end]
else:
tail = tail_lines or 200
start = max(0, total - tail)
selected = lines[start:]
next_offset = start + len(selected)
return {
"content": "\n".join(selected),
"total_lines": total,
"offset": start,
"limit": len(selected),
"has_more": next_offset < total,
"next_offset": next_offset if next_offset < total else None,
}
def _safe_float(value: Any, default: float = 0) -> float:
try:
return float(str(value).replace("[N/A]", "").strip() or default)
except (TypeError, ValueError):
return default
def job_status(job: dict[str, Any]) -> dict[str, Any]:
if execution_mode() != "simulator":
return job
@@ -74,9 +155,80 @@ def create_app() -> FastAPI:
)
return "\n".join(lines)
def real_gpu_resources() -> list[dict[str, Any]]:
query = (
"index,uuid,name,memory.total,memory.used,utilization.gpu,"
"temperature.gpu,power.draw,power.limit"
)
try:
result = subprocess.run(
["nvidia-smi", f"--query-gpu={query}", "--format=csv,noheader,nounits"],
check=True,
capture_output=True,
text=True,
timeout=5,
)
except Exception:
return fallback_gpu_resources()
items: list[dict[str, Any]] = []
for line in result.stdout.splitlines():
parts = [part.strip() for part in line.split(",")]
if len(parts) < 9:
continue
idx, uuid, name, mem_total, mem_used, util, temp, power, power_limit = parts[:9]
total_gb = round(_safe_float(mem_total) / 1024, 2)
used_gb = round(_safe_float(mem_used) / 1024, 2)
memory_percent = round(used_gb / total_gb * 100, 1) if total_gb else 0
gpu_percent = int(_safe_float(util))
items.append(
{
"id": int(idx),
"gpu_index": int(idx),
"uuid": uuid,
"name": name,
"status": "busy" if gpu_percent >= 5 or used_gb > 1 else "idle",
"gpu_percent": gpu_percent,
"memory_used_gb": used_gb,
"memory_total_gb": total_gb,
"memory_percent": memory_percent,
"temperature": int(_safe_float(temp)),
"power_w": round(_safe_float(power), 1),
"power_limit_w": round(_safe_float(power_limit), 1),
"processes": [],
}
)
return items
def fallback_gpu_resources() -> list[dict[str, Any]]:
count = _int_env("COMPUTE_GPU_COUNT", 0)
if count <= 0:
return []
name = os.getenv("COMPUTE_GPU_NAME", "Configured GPU")
memory_total = _float_env("COMPUTE_GPU_MEMORY_GB", 80.0)
power_limit = _float_env("COMPUTE_GPU_POWER_LIMIT_W", 300.0)
return [
{
"id": idx,
"gpu_index": idx,
"uuid": f"GPU-{host_id().upper()}-{idx}",
"name": name,
"status": "idle",
"gpu_percent": 0,
"memory_used_gb": 0,
"memory_total_gb": memory_total,
"memory_percent": 0,
"temperature": _int_env("COMPUTE_GPU_BASE_TEMPERATURE", 35),
"power_w": 0,
"power_limit_w": power_limit,
"processes": [],
}
for idx in range(count)
]
def gpu_resources() -> list[dict[str, Any]]:
if execution_mode() != "simulator":
return []
return real_gpu_resources()
active_jobs = [job_status(job) for job in jobs.values() if job["status"] in {"queued", "running"}]
gpus: list[dict[str, Any]] = []
for idx in range(4):
@@ -109,6 +261,101 @@ def create_app() -> FastAPI:
)
return gpus
def _check_path_item(item: dict[str, Any]) -> dict[str, Any]:
path = Path(str(item.get("path") or ""))
exists = path.exists()
expected_type = str(item.get("type") or "any")
ok = exists
if exists and expected_type == "dir":
ok = path.is_dir()
if exists and expected_type == "file":
ok = path.is_file()
return {
"name": item.get("name") or "",
"path": str(path),
"type": expected_type,
"required": bool(item.get("required", True)),
"exists": exists,
"is_dir": path.is_dir() if exists else False,
"is_file": path.is_file() if exists else False,
"ok": ok or not item.get("required", True),
}
def _job_preview(payload: dict[str, Any], check_paths: bool) -> dict[str, Any]:
warnings: list[str] = []
try:
command = build_command(payload, os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
except ValueError as exc:
return {
"valid": False,
"errors": [part.strip() for part in str(exc).split(";") if part.strip()],
"warnings": warnings,
"engine": str(payload.get("engine") or payload.get("training_engine") or "llama_factory"),
"command": [],
"command_text": "",
"work_dir": os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"),
"env": {},
"path_checks": [],
}
errors: list[str] = []
engine = str(payload.get("engine") or payload.get("training_engine") or "llama_factory")
path_checks: list[dict[str, Any]] = []
if check_paths and engine != "smoke":
path_checks = [
_check_path_item(
{
"name": "model_name_or_path",
"path": payload.get("model_name_or_path") or payload.get("base_model") or "",
"type": "any",
"required": True,
}
)
]
if payload.get("dataset_dir"):
path_checks.append(
_check_path_item(
{
"name": "dataset_dir",
"path": payload.get("dataset_dir"),
"type": "dir",
"required": True,
}
)
)
output_dir = Path(str(payload.get("output_dir") or "/data/yg-ft/outputs/training-job"))
path_checks.append(
_check_path_item(
{
"name": "output_parent",
"path": str(output_dir.parent),
"type": "dir",
"required": False,
}
)
)
errors.extend(
[f"{item['name']} path not available: {item['path']}" for item in path_checks if not item["ok"] and item["required"]]
)
if shutil.which(command.command[0]) is None:
errors.append(f"training command not found: {command.command[0]}")
if not Path(command.work_dir).exists():
errors.append(f"llama_factory_home not found: {command.work_dir}")
elif engine == "smoke":
warnings.append("smoke engine skips model and dataset path checks")
return {
"valid": not errors,
"errors": errors,
"warnings": warnings,
"engine": engine,
"command": command.command,
"command_text": " ".join(command.command),
"work_dir": command.work_dir,
"env": command.env,
"path_checks": path_checks,
}
@app.get(f"{route_prefix}/health")
async def health_check() -> dict[str, str]:
return {
@@ -116,41 +363,121 @@ def create_app() -> FastAPI:
"compute_host_id": os.getenv("COMPUTE_HOST_ID", "unknown"),
}
@app.get("/health")
async def health_check_root() -> dict[str, str]:
return await health_check()
@app.get(f"{route_prefix}/v1/compute/health")
async def compute_health_check() -> dict[str, str | bool]:
async def compute_health_check() -> dict[str, Any]:
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
dataset_root = Path(os.getenv("YG_FT_DATASET_ROOT", str(data_root / "datasets")))
output_root = Path(os.getenv("YG_FT_OUTPUT_ROOT", str(data_root / "outputs")))
llama_factory_home = Path(os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
return {
"status": "ok",
"api_version": "v1",
"compute_host_id": os.getenv("COMPUTE_HOST_ID", "unknown"),
"app_callback_enabled": os.getenv("ENABLE_APP_CALLBACK", "false").lower() == "true",
"data_root": str(data_root),
"data_root_exists": data_root.exists(),
"model_root": os.getenv("YG_FT_MODEL_ROOT", str(data_root / "models")),
"dataset_root": str(dataset_root),
"dataset_root_exists": dataset_root.exists(),
"output_root": str(output_root),
"output_root_exists": output_root.exists(),
"log_root": os.getenv("TRAINING_LOG_ROOT", "/opt/yg-ft/logs/training"),
"llama_factory_home": str(llama_factory_home),
"llama_factory_home_exists": llama_factory_home.exists(),
"llama_factory_version": os.getenv("LLAMA_FACTORY_VERSION", ""),
"execution_mode": execution_mode(),
"gpu_count": _int_env("COMPUTE_GPU_COUNT", 0),
"gpu_discovery_endpoint": f"{route_prefix}/compute/resources/gpus",
"capabilities": ["gpu_discovery", "llama_factory", "file_gateway", "job_polling"],
}
@app.get(f"{route_prefix}/v1/compute/jobs")
async def list_jobs_alias() -> dict[str, list[dict[str, Any]]]:
return {"items": [job_status(job) for job in jobs.values()]}
items = process_manager.list_jobs() if execution_mode() != "simulator" else [job_status(job) for job in jobs.values()]
return {"items": items}
@app.get(f"{route_prefix}/compute/resources/gpus")
async def list_gpus() -> dict[str, Any]:
return {"items": gpu_resources(), "compute_host_id": host_id()}
@app.get(f"{route_prefix}/v1/compute/resources/gpus")
async def list_gpus_v1() -> dict[str, Any]:
return {"items": gpu_resources(), "compute_host_id": host_id()}
@app.post(f"{route_prefix}/compute/jobs/preview")
async def preview_job(payload: dict[str, Any]) -> dict[str, Any]:
return _job_preview(payload, check_paths=False)
@app.post(f"{route_prefix}/compute/jobs/validate")
async def validate_job(payload: dict[str, Any]) -> dict[str, Any]:
return _job_preview(payload, check_paths=True)
@app.post(f"{route_prefix}/v1/compute/jobs/preview")
async def preview_job_v1(payload: dict[str, Any]) -> dict[str, Any]:
return await preview_job(payload)
@app.post(f"{route_prefix}/v1/compute/jobs/validate")
async def validate_job_v1(payload: dict[str, Any]) -> dict[str, Any]:
return await validate_job(payload)
@app.post(f"{route_prefix}/compute/files/check-paths")
async def check_paths(payload: dict[str, Any]) -> dict[str, Any]:
items = [_check_path_item(item) for item in payload.get("paths", []) if isinstance(item, dict)]
return {"valid": all(item["ok"] for item in items), "items": items}
@app.get(f"{route_prefix}/compute/files/list")
async def list_files(
root: str = Query(default="data"),
relative_path: str = Query(default=""),
directories_only: bool = Query(default=False),
) -> dict[str, Any]:
roots = {
"data": Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft")),
"models": Path(os.getenv("YG_FT_MODEL_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/models")),
"datasets": Path(os.getenv("YG_FT_DATASET_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/datasets")),
"outputs": Path(os.getenv("YG_FT_OUTPUT_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/outputs")),
}
base = roots.get(root)
if base is None:
raise HTTPException(status_code=400, detail="invalid root")
target = (base / relative_path.lstrip("/\\")).resolve()
if not _path_inside(base, target):
raise HTTPException(status_code=400, detail="path must stay inside selected root")
if not target.exists():
return {"root": root, "base_path": str(base), "relative_path": relative_path, "items": []}
items = []
for child in sorted(target.iterdir(), key=lambda path: (not path.is_dir(), path.name.lower())):
if directories_only and not child.is_dir():
continue
items.append(
{
"name": child.name,
"path": str(child),
"relative_path": str(child.relative_to(base)).replace("\\", "/"),
"type": "directory" if child.is_dir() else "file",
"byte_size": child.stat().st_size if child.is_file() else 0,
}
)
return {"root": root, "base_path": str(base), "relative_path": relative_path, "items": items}
@app.post(f"{route_prefix}/compute/jobs")
async def create_job(payload: dict[str, Any]) -> dict[str, Any]:
try:
command = build_command(payload, os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc))
if execution_mode() != "simulator":
raise HTTPException(
status_code=501,
detail="real compute executor is not implemented yet; set COMPUTE_EXECUTION_MODE=simulator only for isolated development",
)
job_id = str(payload.get("id") or f"job_{int(now() * 1000)}")
if execution_mode() != "simulator":
try:
return process_manager.create_job({**payload, "id": job_id}, command.command, command.work_dir)
except FileNotFoundError as exc:
raise HTTPException(status_code=500, detail=f"training command not found: {exc.filename}")
except ValueError as exc:
raise HTTPException(status_code=409, detail=str(exc))
job = {
"id": job_id,
"name": payload.get("name", job_id),
@@ -169,17 +496,29 @@ def create_app() -> FastAPI:
@app.get(f"{route_prefix}/compute/jobs")
async def list_jobs() -> dict[str, Any]:
return {"items": [job_status(job) for job in jobs.values()]}
items = process_manager.list_jobs() if execution_mode() != "simulator" else [job_status(job) for job in jobs.values()]
return {"items": items}
@app.get(f"{route_prefix}/compute/jobs/{{job_id}}")
async def get_job(job_id: str) -> dict[str, Any]:
job = jobs.get(job_id)
if execution_mode() != "simulator":
job = process_manager.get_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job_status(job)
@app.post(f"{route_prefix}/compute/jobs/{{job_id}}/stop")
async def stop_job(job_id: str) -> dict[str, Any]:
if execution_mode() != "simulator":
job = process_manager.stop_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
@@ -188,22 +527,101 @@ def create_app() -> FastAPI:
return job
@app.get(f"{route_prefix}/compute/jobs/{{job_id}}/logs")
async def job_logs(job_id: str) -> dict[str, Any]:
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
job = job_status(job)
metrics = [parse_log_line(line) for line in job["logs"].splitlines()]
return {"job_id": job_id, "content": job["logs"], "metrics": [m for m in metrics if m]}
async def job_logs(
job_id: str,
tail_lines: int | None = Query(default=200, ge=1, le=5000),
offset: int | None = Query(default=None, ge=0),
limit: int | None = Query(default=None, ge=1, le=5000),
) -> dict[str, Any]:
if execution_mode() != "simulator":
job = process_manager.get_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
content = process_manager.logs(job_id)
else:
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
job = job_status(job)
content = job["logs"]
window = _slice_log_content(content, tail_lines, offset, limit)
metrics = [parse_log_line(line) for line in window["content"].splitlines()]
return {"job_id": job_id, **window, "metrics": [m for m in metrics if m]}
@app.post(f"{route_prefix}/compute/files/upload")
async def upload_file(payload: dict[str, Any]) -> dict[str, Any]:
file_id = str(payload.get("id") or f"file_{int(now() * 1000)}")
return {"id": file_id, "status": "available", "local_path": f"/data/yg-ft/uploads/{file_id}"}
async def upload_file(
file: UploadFile | None = File(default=None),
target_relative_path: str | None = Form(default=None),
resource_type: str | None = Form(default=None),
resource_id: str | None = Form(default=None),
) -> dict[str, Any]:
file_id = f"file_{int(now() * 1000)}"
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
data_root.mkdir(parents=True, exist_ok=True)
filename = Path(file.filename if file else file_id).name
if target_relative_path:
target = (data_root / target_relative_path.lstrip("/\\")).resolve()
if not _path_inside(data_root, target):
raise HTTPException(status_code=400, detail="target path must stay inside YG_FT_DATA_ROOT")
else:
target = data_root / "uploads" / f"{file_id}_{filename}"
if file:
target.parent.mkdir(parents=True, exist_ok=True)
with target.open("wb") as output:
while chunk := await file.read(1024 * 1024):
output.write(chunk)
else:
target.parent.mkdir(parents=True, exist_ok=True)
target.write_text("", encoding="utf-8")
return {
"id": file_id,
"resource_type": resource_type,
"resource_id": resource_id,
"status": "available",
"local_path": str(target),
"byte_size": target.stat().st_size,
"checksum_sha256": hashlib.sha256(target.read_bytes()).hexdigest() if target.is_file() else "",
}
@app.post(f"{route_prefix}/compute/files/import-local")
async def import_local_file(payload: dict[str, Any]) -> dict[str, Any]:
source = Path(str(payload.get("source_path") or ""))
if not source.exists():
raise HTTPException(status_code=404, detail="source path not found")
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
data_root.mkdir(parents=True, exist_ok=True)
relative = str(payload.get("target_relative_path") or f"imports/{source.name}").lstrip("/\\")
target = (data_root / relative).resolve()
if not _path_inside(data_root, target):
raise HTTPException(status_code=400, detail="target path must stay inside YG_FT_DATA_ROOT")
target.parent.mkdir(parents=True, exist_ok=True)
if source.is_dir():
if target.exists():
shutil.rmtree(target)
shutil.copytree(source, target)
byte_size = sum(path.stat().st_size for path in target.rglob("*") if path.is_file())
checksum = ""
else:
shutil.copy2(source, target)
byte_size = target.stat().st_size
checksum = hashlib.sha256(target.read_bytes()).hexdigest()
return {
"id": str(payload.get("id") or f"file_{int(now() * 1000)}"),
"resource_type": payload.get("resource_type"),
"resource_id": payload.get("resource_id"),
"status": "available",
"local_path": str(target),
"byte_size": byte_size,
"checksum_sha256": checksum,
}
@app.get(f"{route_prefix}/compute/files/{{file_id}}/download")
async def download_file(file_id: str) -> dict[str, Any]:
return {"id": file_id, "status": "ready", "download_url": 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"
matches = list(upload_root.glob(f"{file_id}_*"))
if not matches:
raise HTTPException(status_code=404, detail="file not found")
return FileResponse(matches[0])
return app

View File

@@ -19,22 +19,57 @@ def validate_config(config: dict[str, Any]) -> list[str]:
errors.append("base_model or model_name_or_path is required")
if not config.get("dataset") and not config.get("dataset_dir"):
errors.append("dataset or dataset_dir is required")
learning_rate = float(config.get("learning_rate", 0.0002))
try:
learning_rate = float(config.get("learning_rate", 0.0002))
except (TypeError, ValueError):
learning_rate = 0
if learning_rate <= 0:
errors.append("learning_rate must be greater than zero")
epochs = int(config.get("n_epochs", config.get("num_train_epochs", 1)))
try:
epochs = int(config.get("n_epochs", config.get("num_train_epochs", 1)))
except (TypeError, ValueError):
epochs = 0
if epochs <= 0:
errors.append("n_epochs must be greater than zero")
return errors
def _optional_arg(config: dict[str, Any], command: list[str], option: str, *keys: str) -> None:
for key in keys:
value = config.get(key)
if value is not None and value != "":
command.extend([option, str(value)])
return
def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-Factory") -> LlamaFactoryCommand:
errors = validate_config(config)
if errors:
raise ValueError("; ".join(errors))
engine = str(config.get("engine") or config.get("training_engine") or "llama_factory")
if engine == "smoke":
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-smoke')}"
script = (
"import json, os, time; "
f"out={str(output_dir)!r}; "
"os.makedirs(out, exist_ok=True); "
"print('[INFO] smoke training started', flush=True); "
"\nfor step in range(1, 7):\n"
" loss=round(1.8/(step+1), 4)\n"
" lr=round(0.0002*(1-step/10), 8)\n"
" print({'loss': loss, 'grad_norm': round(0.4 + step*0.03, 4), 'learning_rate': lr, 'epoch': round(step/6, 4)}, flush=True)\n"
" time.sleep(0.4)\n"
"\nopen(os.path.join(out, 'adapter_config.json'), 'w', encoding='utf-8').write(json.dumps({'engine':'smoke','status':'completed'})); "
"print('***** train metrics *****', flush=True); "
"print('train_loss = 0.12', flush=True); "
"print('***** train metrics end *****', flush=True)"
)
return LlamaFactoryCommand(command=["python", "-u", "-c", script], work_dir="/app", env={})
model_path = config.get("base_model") or config.get("model_name_or_path")
dataset = config.get("dataset") or config.get("dataset_dir")
dataset = config.get("dataset") or config.get("dataset_name")
dataset_dir = config.get("dataset_dir")
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-job')}"
command = [
"llamafactory-cli",
@@ -46,7 +81,7 @@ def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-
"--model_name_or_path",
str(model_path),
"--dataset",
str(dataset),
str(dataset or "default"),
"--template",
str(config.get("template", "qwen")),
"--finetuning_type",
@@ -61,7 +96,22 @@ def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-
str(config.get("n_epochs", 3)),
"--save_steps",
str(config.get("save_steps", 50)),
"--logging_steps",
str(config.get("logging_steps", 10)),
"--overwrite_output_dir",
"true",
"--plot_loss",
"true",
]
if dataset_dir:
command.extend(["--dataset_dir", str(dataset_dir)])
_optional_arg(config, command, "--cutoff_len", "max_length", "cutoff_len")
_optional_arg(config, command, "--lr_scheduler_type", "lr_scheduler_type")
_optional_arg(config, command, "--warmup_ratio", "warmup_ratio")
_optional_arg(config, command, "--weight_decay", "weight_decay")
_optional_arg(config, command, "--lora_rank", "lora_rank", "rank")
_optional_arg(config, command, "--lora_alpha", "lora_alpha")
_optional_arg(config, command, "--lora_dropout", "lora_dropout")
quantization_bit = int(config.get("quantization_bit", 0) or 0)
if quantization_bit in {4, 8}:
command.extend(["--quantization_bit", str(quantization_bit)])

View File

@@ -1,5 +1,6 @@
fastapi>=0.111.0
uvicorn[standard]>=0.30.0
python-multipart>=0.0.9
pydantic>=2.7.0
python-dotenv>=1.0.1
httpx>=0.27.0