2026-07-27 09:12:47 +08:00
|
|
|
|
"""
|
|
|
|
|
|
模型训练业务编排(移植自模型服务 projects/backend 的 training_service)
|
|
|
|
|
|
|
|
|
|
|
|
- preset 参数预设(quick / standard / high)
|
|
|
|
|
|
- train_type → stage 映射(sft/dpo/cpt/cot)
|
2026-07-31 16:10:34 +08:00
|
|
|
|
- 训练任务的启动 / 暂停 / 恢复 / 取消
|
|
|
|
|
|
- compute_gateway 状态轮询线程(当任务派发到算力时自动同步状态)
|
2026-07-27 09:12:47 +08:00
|
|
|
|
"""
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
2026-07-31 16:10:34 +08:00
|
|
|
|
import asyncio
|
2026-07-27 09:12:47 +08:00
|
|
|
|
import threading
|
2026-07-31 16:10:34 +08:00
|
|
|
|
import time
|
2026-07-27 09:12:47 +08:00
|
|
|
|
from typing import Any
|
|
|
|
|
|
|
2026-07-31 16:10:34 +08:00
|
|
|
|
from app.core.config import get_settings
|
2026-07-27 09:12:47 +08:00
|
|
|
|
from app.db.platform_store import get_platform_store
|
|
|
|
|
|
from app.modules.fine_tune import runner
|
|
|
|
|
|
|
|
|
|
|
|
PRESETS: dict[str, dict[str, Any]] = {
|
|
|
|
|
|
"quick": {"learning_rate": "5e-5", "n_epochs": 1, "batch_size": 4, "lora_rank": 8},
|
|
|
|
|
|
"standard": {"learning_rate": "2e-5", "n_epochs": 3, "batch_size": 8, "lora_rank": 16},
|
|
|
|
|
|
"high": {"learning_rate": "1e-5", "n_epochs": 5, "batch_size": 4, "lora_rank": 32},
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-07-31 16:10:34 +08:00
|
|
|
|
# ── compute sync 轮询线程 ──────────────────────────────────────────────
|
|
|
|
|
|
_sync_thread: threading.Thread | None = None
|
|
|
|
|
|
_sync_thread_stop = threading.Event()
|
2026-07-27 09:12:47 +08:00
|
|
|
|
|
2026-07-31 16:10:34 +08:00
|
|
|
|
|
|
|
|
|
|
def _compute_sync_loop() -> None:
|
|
|
|
|
|
"""后台线程:周期性轮询算力节点,同步训练任务状态/日志/指标。"""
|
|
|
|
|
|
interval = get_settings().compute_poll_interval_seconds or 3
|
|
|
|
|
|
while not _sync_thread_stop.is_set():
|
|
|
|
|
|
try:
|
|
|
|
|
|
store = get_platform_store()
|
|
|
|
|
|
running = store.running_compute_tasks()
|
|
|
|
|
|
if running:
|
|
|
|
|
|
asyncio.run(_poll_once())
|
|
|
|
|
|
except Exception: # noqa: BLE001 - keep polling loop alive
|
|
|
|
|
|
pass
|
|
|
|
|
|
time.sleep(interval)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def _poll_once() -> None:
|
|
|
|
|
|
from app.modules.compute_gateway.sync import poll_compute_jobs_once
|
|
|
|
|
|
await poll_compute_jobs_once()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def start_compute_sync_worker() -> None:
|
|
|
|
|
|
"""启动后台轮询线程(幂等,多次调用安全)。"""
|
|
|
|
|
|
global _sync_thread
|
|
|
|
|
|
if _sync_thread is not None and _sync_thread.is_alive():
|
|
|
|
|
|
return
|
|
|
|
|
|
_sync_thread_stop.clear()
|
|
|
|
|
|
_sync_thread = threading.Thread(target=_compute_sync_loop, daemon=True)
|
|
|
|
|
|
_sync_thread.start()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def stop_compute_sync_worker() -> None:
|
|
|
|
|
|
"""停止后台轮询线程。"""
|
|
|
|
|
|
_sync_thread_stop.set()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ── preset / config ────────────────────────────────────────────────────
|
2026-07-27 09:12:47 +08:00
|
|
|
|
def apply_presets(payload: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
|
"""根据 preset 字段补全缺失的超参;preset=custom 时不覆盖。"""
|
|
|
|
|
|
payload = dict(payload)
|
|
|
|
|
|
preset = payload.get("preset", "standard")
|
|
|
|
|
|
if preset in PRESETS and payload.get("preset") != "custom":
|
|
|
|
|
|
for key, value in PRESETS[preset].items():
|
|
|
|
|
|
payload.setdefault(key, value)
|
|
|
|
|
|
return payload
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_training_config(payload: dict[str, Any]) -> dict[str, Any]:
|
|
|
|
|
|
"""把前端创建/启动载荷标准化为执行器可消费的 config。"""
|
|
|
|
|
|
payload = apply_presets(dict(payload))
|
|
|
|
|
|
gpus = payload.get("gpus") or [0]
|
|
|
|
|
|
return {
|
|
|
|
|
|
"name": payload.get("name", ""),
|
|
|
|
|
|
"description": payload.get("description", ""),
|
|
|
|
|
|
"train_type": payload.get("train_type", "SFT"),
|
|
|
|
|
|
"train_method": payload.get("train_method", "lora"),
|
|
|
|
|
|
"template": payload.get("template", "qwen"),
|
|
|
|
|
|
"base_model": payload.get("base_model", "") or payload.get("base_model_id", ""),
|
|
|
|
|
|
"train_dataset_id": payload.get("train_dataset_id", ""),
|
|
|
|
|
|
"eval_dataset_id": payload.get("eval_dataset_id", ""),
|
|
|
|
|
|
"auto_merge": bool(payload.get("auto_merge", False)),
|
|
|
|
|
|
"output_model_name": payload.get("output_model_name", ""),
|
|
|
|
|
|
"gpus": gpus,
|
|
|
|
|
|
"num_gpus": payload.get("num_gpus", len(gpus)),
|
|
|
|
|
|
"batch_size": payload.get("batch_size", 2),
|
|
|
|
|
|
"learning_rate": payload.get("learning_rate", 0.0002),
|
|
|
|
|
|
"n_epochs": payload.get("n_epochs", 3),
|
|
|
|
|
|
"save_steps": payload.get("save_steps", 50),
|
|
|
|
|
|
"lr_scheduler_type": payload.get("lr_scheduler_type", "cosine"),
|
|
|
|
|
|
"max_length": payload.get("max_length", 2048),
|
|
|
|
|
|
"warmup_ratio": payload.get("warmup_ratio", 0.03),
|
|
|
|
|
|
"weight_decay": payload.get("weight_decay", 0.01),
|
|
|
|
|
|
"lora_rank": payload.get("lora_rank", 8),
|
|
|
|
|
|
"lora_alpha": payload.get("lora_alpha", 16),
|
|
|
|
|
|
"lora_dropout": payload.get("lora_dropout", 0.05),
|
|
|
|
|
|
"resume_from": payload.get("resume_from"),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def launch_training(task_id: str) -> None:
|
2026-07-31 16:10:34 +08:00
|
|
|
|
"""在后台线程启动真实训练(本机 subprocess fallback,当算力节点不可用时使用)。
|
|
|
|
|
|
|
|
|
|
|
|
架构原则:GPU 计算应派发到算力服务进程执行。
|
|
|
|
|
|
当 platform_store.start_task 检测到在线算力节点时,会走 _dispatch_to_compute 派发路径;
|
|
|
|
|
|
仅当无可用算力节点且非 simulator 模式时,降级到本机 runner(违反 §1.1,待移除)。
|
|
|
|
|
|
"""
|
2026-07-27 09:12:47 +08:00
|
|
|
|
threading.Thread(target=runner.run_training, args=(task_id,), daemon=True).start()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def pause(task_id: str) -> bool:
|
|
|
|
|
|
return runner.pause(task_id)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def resume(task_id: str) -> bool:
|
|
|
|
|
|
return runner.resume(task_id)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def cancel(task_id: str) -> bool:
|
|
|
|
|
|
return runner.cancel(task_id)
|