Files
YG_FT/backend/app/modules/fine_tune/service.py

80 lines
3.1 KiB
Python
Raw Normal View History

2026-07-27 09:12:47 +08:00
"""
模型训练业务编排移植自模型服务 projects/backend training_service
- preset 参数预设quick / standard / high
- train_type stage 映射sft/dpo/cpt/cot
- 训练任务的启动 / 暂停 / 恢复 / 取消委托 runner 真实执行
"""
from __future__ import annotations
import threading
from typing import Any
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},
}
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:
"""在后台线程启动真实训练。"""
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)