From f192062a52abab85f09587557230a87a440bdce7 Mon Sep 17 00:00:00 2001 From: dahai80 <121743945@qq.com> Date: Sat, 8 Aug 2026 12:16:29 +0800 Subject: [PATCH] =?UTF-8?q?feat(#424):=20=E6=8E=A5=E5=85=A5=20reward=20?= =?UTF-8?q?=E8=AE=AD=E7=BB=83=20HTTP=20=E8=B7=AF=E7=94=B1=E4=B8=8E=20serve?= =?UTF-8?q?r=20=E5=90=AF=E5=8A=A8=E8=A3=85=E9=85=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 补齐 PR #427 缺失的对外暴露层: - fine_tune_route.py: 新增 /admin/api/fine-tune/reward/jobs 的 create/list/get/cancel/delete 五个路由 + set_reward_context / _get_reward_service / _create_reward_job 共享辅助,校验 model_id 与 preference_pairs[{prompt,chosen,rejected}] schema。 - server.py: 启动期实例化 RewardService 单例,经 set_reward_context 注入 engine pool 与 event loop,镜像 DPO 装配块。 端到端验证(真实 Qwen3-0.6B-4bit, 4 iters): status=completed, loss 0.69→0.0003, reward_margin 7.98, acc_chosen 1.0, adapter 写盘 (reward_model=true)。 14 项 reward route 单测全绿。 Co-Authored-By: Claude Fable 5 --- fusion_mlx/admin/fine_tune_route.py | 132 ++++++++++++++++++++++++++++ fusion_mlx/server.py | 9 ++ 2 files changed, 141 insertions(+) diff --git a/fusion_mlx/admin/fine_tune_route.py b/fusion_mlx/admin/fine_tune_route.py index dc722f7..8f0feac 100644 --- a/fusion_mlx/admin/fine_tune_route.py +++ b/fusion_mlx/admin/fine_tune_route.py @@ -39,6 +39,7 @@ _engine_pool_ref = None _grpo_service = None _dpo_service = None +_reward_service = None _router = APIRouter() @@ -65,6 +66,13 @@ def set_dpo_context(pool, service=None): service.set_engine_pool(pool) +def set_reward_context(pool, service=None): + global _reward_service + _reward_service = service + if service is not None and pool is not None: + service.set_engine_pool(pool) + + def _get_grpo_service(): global _grpo_service if _grpo_service is None: @@ -87,6 +95,17 @@ def _get_dpo_service(): return _dpo_service +def _get_reward_service(): + global _reward_service + if _reward_service is None: + from fusion_mlx.training.reward_service import RewardService + + _reward_service = RewardService() + if _engine_pool_ref is not None: + _reward_service.set_engine_pool(_engine_pool_ref) + return _reward_service + + def _get_service(): global _fine_tune_service if _fine_tune_service is None: @@ -724,4 +743,117 @@ async def event_generator(): ) +# ============================================================================= +# Reward model training (#424) — /api/fine-tune/reward/jobs +# ============================================================================= + + +def _create_reward_job(request_body: dict): + # Body: {model_id, preference_pairs: [{prompt, chosen, rejected}], + # adapter_name?, config?}. + model_id = request_body.get("model_id", "") + pairs = request_body.get("preference_pairs", []) + adapter_name = request_body.get("adapter_name", "") + + if not model_id: + raise HTTPException(status_code=400, detail="model_id is required") + if not pairs or not isinstance(pairs, list): + raise HTTPException( + status_code=400, detail="preference_pairs (non-empty list) required" + ) + for idx, p in enumerate(pairs): + if not isinstance(p, dict) or not all( + k in p for k in ("prompt", "chosen", "rejected") + ): + raise HTTPException( + status_code=400, + detail=f"preference_pairs[{idx}] must have prompt/chosen/rejected", + ) + + from fusion_mlx.training.reward import RewardConfig + + config_body = dict(request_body.get("config", {})) + try: + config = RewardConfig(**config_body) + except Exception as e: + raise HTTPException(status_code=400, detail=f"Invalid config: {e}") + + pool = _get_engine_pool() + if pool is not None: + entry = pool.get_entry(model_id) + if entry is None: + raise HTTPException(status_code=404, detail=f"Model not found: {model_id}") + if entry.model_type not in ("llm", "vlm", None): + raise HTTPException( + status_code=400, + detail=f"Model {model_id} is not a text model (type: {entry.model_type})", + ) + + svc = _get_reward_service() + job = svc.create_job( + model_id=model_id, + preference_pairs=pairs, + config=config, + adapter_name=adapter_name, + ) + svc.start_processing() + return job.to_dict() + + +@_router.post("/api/fine-tune/reward/jobs") +async def create_reward_job( + request: Request, + is_admin: bool = Depends(require_admin), +): + body = await request.json() + return _create_reward_job(body) + + +@_router.get("/api/fine-tune/reward/jobs") +async def list_reward_jobs( + is_admin: bool = Depends(require_admin), +): + svc = _get_reward_service() + return [job.to_dict() for job in svc.list_jobs()] + + +@_router.get("/api/fine-tune/reward/jobs/{job_id}") +async def get_reward_job( + job_id: str, + is_admin: bool = Depends(require_admin), +): + svc = _get_reward_service() + job = svc.get_job(job_id) + if job is None: + raise HTTPException(status_code=404, detail=f"Job not found: {job_id}") + return job.to_dict() + + +@_router.post("/api/fine-tune/reward/jobs/{job_id}/cancel") +async def cancel_reward_job( + job_id: str, + is_admin: bool = Depends(require_admin), +): + svc = _get_reward_service() + if not svc.cancel_job(job_id): + raise HTTPException( + status_code=404, detail=f"Job not found or not cancellable: {job_id}" + ) + job = svc.get_job(job_id) + return job.to_dict() if job else {"status": "cancelled"} + + +@_router.delete("/api/fine-tune/reward/jobs/{job_id}") +async def delete_reward_job( + job_id: str, + is_admin: bool = Depends(require_admin), +): + svc = _get_reward_service() + if not svc.delete_job(job_id): + raise HTTPException( + status_code=404, detail=f"Job not found or currently running: {job_id}" + ) + return {"status": "deleted"} + + router = _router diff --git a/fusion_mlx/server.py b/fusion_mlx/server.py index c68e17b..ebcbfba 100644 --- a/fusion_mlx/server.py +++ b/fusion_mlx/server.py @@ -1226,6 +1226,15 @@ async def _startup(self): _dpo_svc.set_loop(asyncio.get_running_loop()) set_dpo_context(self.pool, _dpo_svc) + # Wire reward-model training service (#424) + from .admin.fine_tune_route import set_reward_context + from .training.reward_service import RewardService + + _reward_svc = RewardService() + _reward_svc.set_engine_pool(self.pool) + _reward_svc.set_loop(asyncio.get_running_loop()) + set_reward_context(self.pool, _reward_svc) + # Auto-add adapters dir to FUSION_LORA_ALLOWED_DIRS so trained # adapters can be served via EnginePool hot-swap without manual env config import os