Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
96 changes: 96 additions & 0 deletions fusion_mlx/admin/fine_tune_route.py
Original file line number Diff line number Diff line change
Expand Up @@ -856,4 +856,100 @@ async def delete_reward_job(
return {"status": "deleted"}


@_router.post("/api/fine-tune/reward/score")
async def score_reward_endpoint(
request: Request,
is_admin: bool = Depends(require_admin),
):
# Score (prompt, completions) under a trained reward adapter (#431).
# Loads standalone (separate from inference pool), attaches the trained
# value head, scores each completion, evicts. Closes the Phase1 (#424)
# reward-model -> Phase2 (#363 GRPO) loop: this URL is passed as GRPO's
# config.reward_endpoint, which POSTs {prompt, completions} -> {rewards}.
# GRPO's callback protocol is fixed at {prompt, completions} (no
# model_id/adapter_name), so those MUST be supplied as query params in
# the reward_endpoint URL, e.g. .../reward/score?key=T&model_id=X&adapter_name=Y.
body = await request.json()

model_id = body.get("model_id", "") or request.query_params.get("model_id", "")
adapter_name = body.get("adapter_name", "") or request.query_params.get(
"adapter_name", ""
)
prompt = body.get("prompt", "")
completions = body.get("completions", [])

if not model_id:
raise HTTPException(status_code=400, detail="model_id is required")
if not prompt:
raise HTTPException(status_code=400, detail="prompt is required")
if not completions or not isinstance(completions, list):
raise HTTPException(
status_code=400, detail="completions (non-empty list) required"
)
if not adapter_name:
raise HTTPException(
status_code=400,
detail="adapter_name is required (a trained reward adapter)",
)

svc = _get_service()
model_path = svc._resolve_model_path(model_id)
if model_path is None:
raise HTTPException(status_code=404, detail=f"Model not found: {model_id}")

from fusion_mlx.training.service import ADAPTER_BASE_DIR

adapter_path = str(ADAPTER_BASE_DIR / model_id / adapter_name)
import os

if not os.path.isdir(adapter_path):
raise HTTPException(
status_code=404,
detail=f"Adapter not found: {model_id}/{adapter_name}",
)

from fusion_mlx.training.reward import RewardScoreResult
from fusion_mlx.training.reward import score_text as reward_score_text

logger.info(
"reward score endpoint: model=%s adapter=%s n_completions=%d",
model_path,
adapter_path,
len(completions),
)

def _run():
import mlx_lm.utils as mlx_utils

model, tokenizer = mlx_utils.load(model_path, adapter_path=adapter_path)
try:
rewards = reward_score_text(
model, tokenizer, model_path, prompt, completions, adapter_path
)
return RewardScoreResult(
rewards=rewards, model_id=model_id, adapter_name=adapter_name
)
finally:
del model
del tokenizer
import gc

import mlx.core as mx

gc.collect()
mx.clear_cache()
logger.info("reward score endpoint: model evicted")

try:
result = await asyncio.to_thread(_run)
except ValueError as e:
logger.exception("reward score endpoint: bad adapter")
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
logger.exception("reward score endpoint failed")
raise HTTPException(status_code=500, detail=f"Scoring failed: {e}")

return result.to_dict()


router = _router
98 changes: 98 additions & 0 deletions fusion_mlx/training/reward.py
Original file line number Diff line number Diff line change
Expand Up @@ -220,3 +220,101 @@ def save_adapter(self, adapter_path):
weights = dict(tree_flatten(self.model.trainable_parameters()))
mx.save_safetensors(adapter_path, weights)
logger.info("REWARD: saved reward adapter to %s", adapter_path)


@dataclass
class RewardScoreResult:
rewards: list
model_id: str
adapter_name: str

def to_dict(self):
return {
"rewards": self.rewards,
"model_id": self.model_id,
"adapter_name": self.adapter_name,
}


def score_text(model, tokenizer, model_path, prompt, completions, adapter_path=None):
# Inference-time scalar reward scoring for (prompt, completion) pairs,
# using a trained reward adapter (LoRA backbone + value head, #424).
# Mirror of RewardTrainer._score but non-differentiable: forward the
# concatenated sequence, take the last-token hidden, project via the
# value head to a scalar. Loads value_head weights from the adapter's
# safetensors (the standard mlx_utils.load adapter_path path applies
# LoRA but does not restore the custom value_head submodule).
import os

from safetensors import safe_open

logger.info(
"reward score_text: model=%s adapter=%s prompt_len=%d n_completions=%d",
model_path,
adapter_path,
len(prompt),
len(completions),
)

head = None
if adapter_path:
weights_file = os.path.join(adapter_path, "adapters.safetensors")
vh_keys = {}
if os.path.isfile(weights_file):
with safe_open(weights_file, "mlx") as f:
for k in list(f.keys()):
if k.startswith("value_head."):
vh_keys[k] = f.get_tensor(k)
if "value_head.proj.weight" in vh_keys:
hidden = int(vh_keys["value_head.proj.weight"].shape[1])
head = _ValueHead(hidden)
head.load_weights(
[
("proj.weight", vh_keys["value_head.proj.weight"]),
("proj.bias", vh_keys["value_head.proj.bias"]),
]
)
logger.info("reward score_text: value head loaded hidden=%d", hidden)
else:
logger.warning(
"reward score_text: adapter has no value_head weights (%s), "
"scoring with untrained head",
weights_file,
)

if getattr(model, "value_head", None) is None:
if head is not None:
model.value_head = head
else:
raise ValueError(
"reward score_text: model has no value_head and adapter "
"provided no value_head weights; not a reward model"
)

def _score_one(prompt_ids, completion_ids):
full = mx.concatenate([prompt_ids, completion_ids])
trunk = getattr(model, "transformer", None) or getattr(model, "model", None)
if trunk is not None:
hidden = trunk(full[None, :])
if isinstance(hidden, tuple):
hidden = hidden[0]
hidden = hidden[0]
else:
logits = _model_forward_logits(model, full)
if isinstance(logits, tuple):
logits = logits[0]
n_comp = int(completion_ids.shape[0])
hidden = mx.mean(logits[0, -n_comp:, :].astype(mx.float32), axis=0)
hidden = mx.expand_dims(hidden, 0)
s = model.value_head(hidden)
mx.eval(s)
return float(s)

rewards = []
prompt_ids = mx.array(tokenizer.encode(prompt))
for completion in completions:
completion_ids = mx.array(tokenizer.encode(completion))
rewards.append(_score_one(prompt_ids, completion_ids))

logger.info("reward score_text: rewards=%s", rewards)
return rewards
Loading
Loading