""" 模型训练业务编排(移植自模型服务 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)