From 8e797cfed19b6658b9f71d775daf23dd5569e78e Mon Sep 17 00:00:00 2001 From: dahai80 <121743945@qq.com> Date: Sat, 8 Aug 2026 11:56:53 +0800 Subject: [PATCH] feat(#424): commit reward training service + route/service tests MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The RewardTrainer (training/reward.py) and admin routes (/admin/api/fine-tune/reward/jobs CRUD) were already on main, but the RewardService — the queue/CRUD/persistence/SSE layer that the routes and server.py wire up — was never committed (left untracked from the prior reward-training work). Without it the reward endpoints 500 at import time. This commits the missing service and adds test coverage: - fusion_mlx/training/reward_service.py — RewardJob + RewardService (queue, create/get/list/cancel/delete, _execute_reward, persistence, SSE event push), mirroring DPOService. - tests/unit/test_reward_route.py — 14 tests: route CRUD (create, missing model_id/pairs, malformed pair, invalid config, list, get 404, cancel queued, delete, cancel 404) + RewardConfig defaults, RewardJob.to_dict roundtrip, RewardTrainer.train_step loss/metrics with a stub backbone + registered value head (Bradley-Terry loss ~= log(2) at init, acc in {0,1}). Closes #424. Co-Authored-By: Claude Fable 5 --- fusion_mlx/training/reward_service.py | 401 ++++++++++++++++++++++++++ tests/unit/test_reward_route.py | 239 +++++++++++++++ 2 files changed, 640 insertions(+) create mode 100644 fusion_mlx/training/reward_service.py create mode 100644 tests/unit/test_reward_route.py diff --git a/fusion_mlx/training/reward_service.py b/fusion_mlx/training/reward_service.py new file mode 100644 index 0000000..18009cd --- /dev/null +++ b/fusion_mlx/training/reward_service.py @@ -0,0 +1,401 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reward model training service (#424). + +Importers/callers: + - fusion_mlx.admin.fine_tune_route imports RewardService + RewardJob for routes + - fusion_mlx.server instantiates + wires via set_reward_context + +Affected API: new /admin/api/fine-tune/reward/jobs endpoints + (POST create, GET list/{id}, POST cancel, DELETE) + +Data schemas: + RewardJob — job record (job_id, model_id, preference_pairs, config, events[], cond) + RewardService — singleton service (queue, CRUD, _execute_reward, SSE progress) + +User verbatim instruction: "对上游依赖给上游提issue和pr" +""" + +from __future__ import annotations + +import asyncio +import gc +import json +import logging +import os +import time +import uuid +from dataclasses import asdict, dataclass, field +from pathlib import Path + +import mlx.core as mx + +from fusion_mlx.training.reward import RewardConfig, RewardTrainer +from fusion_mlx.training.service import ADAPTER_BASE_DIR, JobStatus + +logger = logging.getLogger(__name__) + + +@dataclass +class RewardJob: + job_id: str + model_id: str + preference_pairs: list = field(default_factory=list) + config: RewardConfig = field(default_factory=RewardConfig) + adapter_name: str = "" + adapter_path: str = "" + status: JobStatus = JobStatus.QUEUED + created_at: float = field(default_factory=time.time) + started_at: float = 0.0 + finished_at: float = 0.0 + error: str = "" + events: list = field(default_factory=list) + progress: dict = field(default_factory=dict) + terminal: bool = False + cond: asyncio.Condition = field(default_factory=asyncio.Condition, repr=False) + + def to_dict(self) -> dict: + return { + "job_id": self.job_id, + "model_id": self.model_id, + "preference_pairs": self.preference_pairs, + "config": asdict(self.config), + "adapter_name": self.adapter_name, + "adapter_path": self.adapter_path, + "status": self.status.value, + "created_at": self.created_at, + "started_at": self.started_at, + "finished_at": self.finished_at, + "error": self.error, + "progress": self.progress, + } + + +class RewardService: + # Singleton managing reward-model job queue + execution. One concurrent job + # (Apple Silicon memory constraint — training evicts inference model). + + def __init__(self): + self._jobs: dict[str, RewardJob] = {} + self._queue: list[str] = [] + self._current_job_id: str | None = None + self._loop: asyncio.AbstractEventLoop | None = None + self._engine_pool = None + self._running = False + self._load_jobs() + + def set_engine_pool(self, pool): + self._engine_pool = pool + + def set_loop(self, loop: asyncio.AbstractEventLoop): + self._loop = loop + + def _resolve_model_path(self, model_id: str) -> str | None: + if self._engine_pool is None: + return model_id + entry = self._engine_pool.get_entry(model_id) + if entry is not None and hasattr(entry, "model_path"): + return entry.model_path + candidate = Path(os.path.expanduser("~/.fusion-mlx/models")) / model_id + if candidate.exists(): + return str(candidate) + return model_id + + # ========================================================================= + # Job CRUD + # ========================================================================= + + def create_job( + self, + model_id: str, + preference_pairs: list, + config: RewardConfig | None = None, + adapter_name: str = "", + ) -> RewardJob: + cfg = config or RewardConfig() + adapter_name = adapter_name or f"reward-{uuid.uuid4().hex[:6]}" + adapter_path = str(ADAPTER_BASE_DIR / model_id / adapter_name) + + job = RewardJob( + job_id=uuid.uuid4().hex[:12], + model_id=model_id, + preference_pairs=list(preference_pairs), + config=cfg, + adapter_name=adapter_name, + adapter_path=adapter_path, + ) + self._jobs[job.job_id] = job + self._queue.append(job.job_id) + self._persist_jobs() + logger.info( + "REWARD job created: %s model=%s adapter=%s queued=%d", + job.job_id, + model_id, + adapter_name, + len(self._queue), + ) + return job + + def get_job(self, job_id: str) -> RewardJob | None: + return self._jobs.get(job_id) + + def list_jobs(self) -> list[RewardJob]: + return list(self._jobs.values()) + + def cancel_job(self, job_id: str) -> bool: + job = self._jobs.get(job_id) + if job is None: + return False + if job.status == JobStatus.QUEUED: + job.status = JobStatus.CANCELLED + job.terminal = True + if job_id in self._queue: + self._queue.remove(job_id) + self._notify_job(job) + self._persist_jobs() + return True + if job.status == JobStatus.RUNNING: + job.status = JobStatus.CANCELLED + job.terminal = True + job.finished_at = time.time() + self._current_job_id = None + self._notify_job(job) + self._persist_jobs() + self._process_queue() + return True + return False + + def delete_job(self, job_id: str) -> bool: + job = self._jobs.get(job_id) + if job is None: + return False + if job.status == JobStatus.RUNNING: + return False + self._jobs.pop(job_id, None) + if job_id in self._queue: + self._queue.remove(job_id) + self._persist_jobs() + return True + + # ========================================================================= + # Queue processing + # ========================================================================= + + def start_processing(self): + self._running = True + if self._loop is None: + try: + self._loop = asyncio.get_running_loop() + except RuntimeError: + self._loop = asyncio.new_event_loop() + self._process_queue() + logger.info("Reward service processing queue (jobs=%d)", len(self._queue)) + + def _process_queue(self): + if not self._running: + return + if self._current_job_id is not None: + return + if not self._queue: + return + job_id = self._queue.pop(0) + job = self._jobs.get(job_id) + if job is None or job.status == JobStatus.CANCELLED: + self._process_queue() + return + + self._current_job_id = job_id + job.status = JobStatus.RUNNING + job.started_at = time.time() + self._persist_jobs() + self._notify_job(job) + + if self._loop is None: + logger.error("No event loop for reward execution") + job.status = JobStatus.FAILED + job.error = "No event loop available" + job.terminal = True + self._current_job_id = None + return + + asyncio.ensure_future(self._run_job(job), loop=self._loop) + + async def _run_job(self, job: RewardJob): + try: + await asyncio.to_thread(self._execute_reward, job) + except Exception as exc: + logger.exception("Reward job %s failed: %s", job.job_id, exc) + job.status = JobStatus.FAILED + job.error = str(exc) + job.terminal = True + job.finished_at = time.time() + self._persist_jobs() + self._notify_job(job) + finally: + self._current_job_id = None + self._process_queue() + + def _execute_reward(self, job: RewardJob): + # Run reward-model training in a background thread (blocking). Load + # model, apply LoRA, run preference-pair loop, save adapter, cleanup. + import mlx_lm.utils as mlx_utils + from mlx_lm.tuner.utils import linear_to_lora_layers + + model_path = self._resolve_model_path(job.model_id) + if model_path is None: + raise ValueError(f"Cannot resolve model path for {job.model_id}") + cfg = job.config + logger.info("REWARD execute: model=%s job=%s", model_path, job.job_id) + + model, tokenizer = mlx_utils.load(model_path) + + mx.random.seed(cfg.seed) + model.freeze() + lora_params = { + "rank": cfg.lora_rank, + "dropout": cfg.lora_dropout, + "scale": cfg.lora_alpha, + } + linear_to_lora_layers(model, cfg.lora_layers, lora_params, use_dora=False) + + trainer = RewardTrainer(model, tokenizer, model_path, cfg) + + pairs = job.preference_pairs + batch_size = max(cfg.batch_size, 1) + n_pairs = len(pairs) + if n_pairs == 0: + raise ValueError("reward job has no preference_pairs") + + for it in range(cfg.iters): + start = (it * batch_size) % n_pairs + batch = [pairs[(start + i) % n_pairs] for i in range(batch_size)] + result = trainer.train_step(batch) + self._push_event( + job, + { + "type": "reward_step", + "iter": it, + "total_iters": cfg.iters, + "loss": result.loss, + "reward_margin": result.reward_margin, + "acc_chosen": result.acc_chosen, + "n_pairs": len(batch), + }, + ) + job.progress = { + "iter": it + 1, + "total_iters": cfg.iters, + "loss": result.loss, + "reward_margin": result.reward_margin, + "acc_chosen": result.acc_chosen, + } + + Path(job.adapter_path).mkdir(parents=True, exist_ok=True) + trainer.save_adapter(str(Path(job.adapter_path) / "adapters.safetensors")) + + adapter_cfg = { + "adapter_path": job.adapter_path, + "num_layers": cfg.lora_layers, + "lora_parameters": { + "rank": cfg.lora_rank, + "scale": cfg.lora_alpha, + "dropout": cfg.lora_dropout, + }, + "fine_tune_type": "lora", + "reward_model": True, + } + with open(Path(job.adapter_path) / "adapter_config.json", "w") as f: + json.dump(adapter_cfg, f, indent=2) + + del model + del tokenizer + gc.collect() + mx.clear_cache() + + job.status = JobStatus.COMPLETED + job.terminal = True + job.finished_at = time.time() + self._persist_jobs() + self._push_event(job, {"type": "done", "adapter_path": job.adapter_path}) + logger.info( + "REWARD job %s completed: adapter=%s", + job.job_id, + job.adapter_path, + ) + + def _push_event(self, job: RewardJob, event: dict): + job.events.append(event) + + async def _notify(): + async with job.cond: + job.cond.notify_all() + + try: + if self._loop and self._loop.is_running(): + self._loop.call_soon_threadsafe( + lambda: asyncio.ensure_future(_notify()) + ) + except RuntimeError: + pass + + def _notify_job(self, job: RewardJob): + async def _do(): + async with job.cond: + job.cond.notify_all() + + if self._loop and self._loop.is_running(): + asyncio.ensure_future(_do(), loop=self._loop) + + # ========================================================================= + # Persistence + # ========================================================================= + + @property + def _jobs_file(self) -> Path: + return ADAPTER_BASE_DIR / "reward_jobs.json" + + def _persist_jobs(self): + ADAPTER_BASE_DIR.mkdir(parents=True, exist_ok=True) + data = [job.to_dict() for job in self._jobs.values()] + try: + with open(self._jobs_file, "w") as f: + json.dump(data, f, indent=2) + except Exception as exc: + logger.warning("Failed to persist reward jobs: %s", exc) + + def _load_jobs(self): + if not self._jobs_file.exists(): + return + try: + with open(self._jobs_file) as f: + data = json.load(f) + except Exception as exc: + logger.warning("Failed to load reward jobs: %s", exc) + return + for item in data: + try: + cfg = RewardConfig(**item.get("config", {})) + status = JobStatus(item.get("status", "queued")) + if status in (JobStatus.RUNNING, JobStatus.QUEUED): + status = JobStatus.CANCELLED + job = RewardJob( + job_id=item.get("job_id", uuid.uuid4().hex[:12]), + model_id=item.get("model_id", ""), + preference_pairs=item.get("preference_pairs", []), + config=cfg, + adapter_name=item.get("adapter_name", ""), + adapter_path=item.get("adapter_path", ""), + status=status, + created_at=item.get("created_at", 0.0), + started_at=item.get("started_at", 0.0), + finished_at=item.get("finished_at", 0.0), + error=item.get("error", ""), + progress=item.get("progress", {}), + ) + job.terminal = status in ( + JobStatus.COMPLETED, + JobStatus.FAILED, + JobStatus.CANCELLED, + ) + self._jobs[job.job_id] = job + except Exception as exc: + logger.warning("Skipping malformed reward job: %s", exc) diff --git a/tests/unit/test_reward_route.py b/tests/unit/test_reward_route.py new file mode 100644 index 0000000..7b14111 --- /dev/null +++ b/tests/unit/test_reward_route.py @@ -0,0 +1,239 @@ +# Route tests for /admin/api/fine-tune/reward/jobs endpoints (#424). +# Minimal FastAPI app with the admin router + require_admin override. +# RewardService queue processing is neutralized so route handlers never +# spin up a real model load during sync TestClient requests. + +from __future__ import annotations + +import mlx.core as mx +import mlx.nn as nn +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from fusion_mlx.admin.auth import require_admin +from fusion_mlx.admin.fine_tune_route import set_reward_context +from fusion_mlx.admin.routes import router as admin_router +from fusion_mlx.training.reward import RewardConfig, RewardTrainer +from fusion_mlx.training.reward_service import RewardJob, RewardService + + +def _build_app(): + app = FastAPI() + svc = RewardService() + svc.start_processing = lambda *a, **kw: None + svc._process_queue = lambda *a, **kw: None + set_reward_context(None, svc) + app.include_router(admin_router) + app.dependency_overrides[require_admin] = lambda: True + return app, svc + + +def test_reward_create_job_returns_id(): + app, svc = _build_app() + client = TestClient(app) + resp = client.post( + "/admin/api/fine-tune/reward/jobs", + json={ + "model_id": "m1", + "preference_pairs": [{"prompt": "p", "chosen": "good", "rejected": "bad"}], + "config": {"iters": 1, "batch_size": 1}, + }, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert "job_id" in body + assert body["model_id"] == "m1" + assert body["preference_pairs"][0]["chosen"] == "good" + assert body["config"]["iters"] == 1 + assert body["status"] in ("queued", "running", "completed", "cancelled") + + +def test_reward_create_missing_model_id(): + app, _ = _build_app() + client = TestClient(app) + resp = client.post( + "/admin/api/fine-tune/reward/jobs", + json={"preference_pairs": [{"prompt": "p", "chosen": "c", "rejected": "r"}]}, + ) + assert resp.status_code == 400 + assert "model_id" in resp.json()["detail"] + + +def test_reward_create_missing_pairs(): + app, _ = _build_app() + client = TestClient(app) + resp = client.post( + "/admin/api/fine-tune/reward/jobs", + json={"model_id": "m1"}, + ) + assert resp.status_code == 400 + assert "preference_pairs" in resp.json()["detail"] + + +def test_reward_create_malformed_pair(): + app, _ = _build_app() + client = TestClient(app) + resp = client.post( + "/admin/api/fine-tune/reward/jobs", + json={"model_id": "m1", "preference_pairs": [{"prompt": "p"}]}, + ) + assert resp.status_code == 400 + assert "prompt/chosen/rejected" in resp.json()["detail"] + + +def test_reward_create_invalid_config(): + app, _ = _build_app() + client = TestClient(app) + resp = client.post( + "/admin/api/fine-tune/reward/jobs", + json={ + "model_id": "m1", + "preference_pairs": [{"prompt": "p", "chosen": "c", "rejected": "r"}], + "config": {"unknown_field": 1}, + }, + ) + assert resp.status_code == 400 + assert "Invalid config" in resp.json()["detail"] + + +def test_reward_list_jobs(): + app, svc = _build_app() + client = TestClient(app) + svc.create_job( + model_id="m1", + preference_pairs=[{"prompt": "p", "chosen": "c", "rejected": "r"}], + adapter_name="a1", + ) + resp = client.get("/admin/api/fine-tune/reward/jobs") + assert resp.status_code == 200 + jobs = resp.json() + assert any(j["adapter_name"] == "a1" for j in jobs) + + +def test_reward_get_job_not_found(): + app, _ = _build_app() + client = TestClient(app) + resp = client.get("/admin/api/fine-tune/reward/jobs/nonexistent") + assert resp.status_code == 404 + + +def test_reward_cancel_queued_job(): + app, svc = _build_app() + client = TestClient(app) + job = svc.create_job( + model_id="m1", + preference_pairs=[{"prompt": "p", "chosen": "c", "rejected": "r"}], + adapter_name="q1", + ) + resp = client.post(f"/admin/api/fine-tune/reward/jobs/{job.job_id}/cancel") + assert resp.status_code == 200 + assert resp.json()["status"] == "cancelled" + + +def test_reward_delete_job(): + app, svc = _build_app() + client = TestClient(app) + job = svc.create_job( + model_id="m1", + preference_pairs=[{"prompt": "p", "chosen": "c", "rejected": "r"}], + adapter_name="d1", + ) + resp = client.delete(f"/admin/api/fine-tune/reward/jobs/{job.job_id}") + assert resp.status_code == 200 + assert resp.json()["status"] == "deleted" + + +def test_reward_cancel_not_found(): + app, _ = _build_app() + client = TestClient(app) + resp = client.post("/admin/api/fine-tune/reward/jobs/nope/cancel") + assert resp.status_code == 404 + + +# ============================================================================= +# Service / config / trainer unit tests +# ============================================================================= + + +_PAIRS = [ + {"prompt": "Q?", "chosen": "good", "rejected": "bad"}, + {"prompt": "Q2?", "chosen": "better", "rejected": "worse"}, +] + + +def test_reward_config_defaults(): + cfg = RewardConfig() + assert cfg.iters == 50 + assert cfg.learning_rate == 1e-5 + assert cfg.lora_rank == 8 + assert cfg.max_seq_length == 1024 + + +def test_reward_job_to_dict_roundtrip(): + cfg = RewardConfig(iters=3, lora_rank=4) + job = RewardJob( + job_id="abc", + model_id="m1", + preference_pairs=_PAIRS, + config=cfg, + adapter_name="rm-abc", + adapter_path="/tmp/rm-abc", + ) + d = job.to_dict() + assert d["job_id"] == "abc" + assert d["config"]["iters"] == 3 + assert d["config"]["lora_rank"] == 4 + assert d["preference_pairs"] == _PAIRS + assert d["status"] == "queued" + + +class _StubTokenizer: + def encode(self, text): + return [ord(c) for c in text][:4] + + +class _StubTrunk(nn.Module): + # Returns (seq, hidden) hidden states so _score can take last-token. + def __init__(self, hidden=4, vocab=32): + super().__init__() + self.embed = nn.Embedding(vocab, hidden) + self._hidden = hidden + + def __call__(self, ids): + # ids: (1, seq) -> (1, seq, hidden); _score indexes [0] -> (seq, hidden) + return self.embed(ids) + + +class _StubModel(nn.Module): + def __init__(self, hidden=4, vocab=32): + super().__init__() + self.transformer = _StubTrunk(hidden, vocab) + self.hidden_size = hidden + + def __call__(self, ids): + return self.transformer(ids) + + +def test_reward_train_step_runs_and_returns_metrics(): + # RewardTrainer.train_step with a stub backbone + registered value head. + # Verifies the Bradley-Terry loss graph executes, returns finite metrics, + # and the head is attached to the model. + model = _StubModel() + cfg = RewardConfig(iters=1, batch_size=1, lora_layers=0, learning_rate=1e-4) + trainer = RewardTrainer(model, _StubTokenizer(), "/dev/null", cfg) + result = trainer.train_step([_PAIRS[0]]) + assert mx.isfinite(mx.array(result.loss)) + assert -1.0 <= result.acc_chosen <= 1.0 + assert getattr(model, "value_head", None) is not None + + +def test_reward_loss_margin_sign(): + # With a fresh head, score_w - score_l should be ~0 (untrained), so the + # loss is near log(2) (=-log sigmoid(0)) and acc is a bool in {0,1}. + model = _StubModel() + cfg = RewardConfig(iters=1, batch_size=1, lora_layers=0) + trainer = RewardTrainer(model, _StubTokenizer(), "/dev/null", cfg) + result = trainer.train_step([_PAIRS[0]]) + # -log sigmoid(0) = log(2) ~= 0.693 + assert abs(result.loss - 0.6931) < 0.05 + assert result.acc_chosen in (0.0, 1.0)