mirror of https://github.com/razor-ai/soup.git
feat(prm): PRM reward pure kernels (split_steps/aggregate) (v0.71.30)
This commit is contained in:
parent
3c21c5e7c8
commit
79e6c119de
|
|
@ -0,0 +1,330 @@
|
|||
"""PRM-as-per-step-reward for GRPO — v0.71.30 (#PRM-guided GRPO).
|
||||
|
||||
Use a trained Soup PRM (the v0.53.11 ``PRMTrainerWrapper``) as the reward
|
||||
function inside GRPO: split each generated completion into reasoning steps,
|
||||
score every step with the PRM's scalar reward head, and fold the per-step
|
||||
scores into a single scalar reward (``min`` / ``prod`` / ``last``) that GRPO
|
||||
optimises. This is the o1-era process-supervision training signal.
|
||||
|
||||
Design:
|
||||
- The pure kernels (:func:`split_steps`, :func:`aggregate_step_scores`) carry
|
||||
NO torch dependency and are unit-testable on the light core.
|
||||
- :class:`PRMScorer` is a stateful ``(completions, **kwargs) -> list[float]``
|
||||
callable (torch lazy-loaded inside methods) that GRPO uses as its
|
||||
``reward_fn``. It rides the existing shaping + ``wrap_reward_funcs`` seam in
|
||||
``trainer/grpo.py`` unchanged, so the v0.71.26 reward-hack mitigation
|
||||
controller observes the PRM reward for free.
|
||||
|
||||
Honesty: proof-of-mechanism only — a tiny PRM signal is noisy; this is NOT a
|
||||
production reward-model claim (see #286). The step split is a newline
|
||||
heuristic (v1).
|
||||
|
||||
Security:
|
||||
- ``prm_reward`` local paths are containment-checked (realpath + commonpath
|
||||
under cwd) in :func:`build_prm_reward_fn`; the reward-head weights load via
|
||||
``safetensors.safe_open`` (no pickle).
|
||||
- Bounded step count / per-step chars; aggregate mode is a closed allowlist.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
# Aggregation modes for folding per-step scores into a scalar reward.
|
||||
AGGREGATE_MODES: tuple[str, ...] = ("min", "prod", "last")
|
||||
|
||||
# Bounds — a pathological completion cannot blow up the forward pass.
|
||||
_MAX_STEPS = 64
|
||||
_MAX_STEP_CHARS = 2_000
|
||||
|
||||
|
||||
def split_steps(text: Any) -> list[str]:
|
||||
"""Split a completion into reasoning steps (newline heuristic, v1).
|
||||
|
||||
Drops empty / whitespace-only lines, strips each line, truncates each step
|
||||
to ``_MAX_STEP_CHARS`` and the whole list to ``_MAX_STEPS``. Returns ``[]``
|
||||
for non-string input.
|
||||
"""
|
||||
if not isinstance(text, str):
|
||||
return []
|
||||
steps: list[str] = []
|
||||
for line in text.splitlines():
|
||||
stripped = line.strip()
|
||||
if not stripped:
|
||||
continue
|
||||
steps.append(stripped[:_MAX_STEP_CHARS])
|
||||
if len(steps) >= _MAX_STEPS:
|
||||
break
|
||||
return steps
|
||||
|
||||
|
||||
def _finite(value: Any) -> float:
|
||||
"""Coerce ``value`` to a finite float; non-finite / non-numeric → 0.0."""
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
return 0.0
|
||||
fv = float(value)
|
||||
return fv if math.isfinite(fv) else 0.0
|
||||
|
||||
|
||||
def aggregate_step_scores(scores: list[float], mode: Any) -> float:
|
||||
"""Fold per-step scalar scores into a single reward.
|
||||
|
||||
``mode`` is one of ``AGGREGATE_MODES``. Empty ``scores`` → 0.0. Non-finite
|
||||
per-step values are coerced to 0.0 before folding (a NaN step reward must
|
||||
not poison the whole reward).
|
||||
"""
|
||||
if isinstance(mode, bool) or not isinstance(mode, str) or mode not in AGGREGATE_MODES:
|
||||
raise ValueError(
|
||||
f"prm_aggregate must be one of {AGGREGATE_MODES}; got {mode!r}"
|
||||
)
|
||||
if not scores:
|
||||
return 0.0
|
||||
clean = [_finite(s) for s in scores]
|
||||
if mode == "min":
|
||||
return float(min(clean))
|
||||
if mode == "last":
|
||||
return float(clean[-1])
|
||||
# prod
|
||||
return float(math.prod(clean))
|
||||
|
||||
|
||||
class PRMScorer:
|
||||
"""Stateful GRPO reward: score each completion's steps with a Soup PRM.
|
||||
|
||||
Torch / transformers / safetensors are lazy-imported inside methods so the
|
||||
module stays importable on the light core. ``__name__`` is set to
|
||||
``"prm_reward"`` so TRL's ``rewards/<func_name>`` logging key is stable and
|
||||
the mitigation buffer records under a readable name.
|
||||
"""
|
||||
|
||||
__name__ = "prm_reward"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
prm_path: str,
|
||||
aggregate: str = "min",
|
||||
device: str = "cpu",
|
||||
trust_remote_code: bool = False,
|
||||
) -> None:
|
||||
if isinstance(aggregate, bool) or aggregate not in AGGREGATE_MODES:
|
||||
raise ValueError(
|
||||
f"aggregate must be one of {AGGREGATE_MODES}; got {aggregate!r}"
|
||||
)
|
||||
self.prm_path = prm_path
|
||||
self.aggregate = aggregate
|
||||
self.device = device
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
|
||||
def _ensure_loaded(self) -> None:
|
||||
"""Lazy-load the base CausalLM + re-attach and load the reward head."""
|
||||
if self._model is not None:
|
||||
return
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
self.prm_path, trust_remote_code=self.trust_remote_code
|
||||
)
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
dtype = torch.bfloat16 if self.device == "cuda" else torch.float32
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.prm_path,
|
||||
trust_remote_code=self.trust_remote_code,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
head_state = load_reward_head_weights(self.prm_path)
|
||||
hidden_size = model.config.hidden_size
|
||||
reward_head = nn.Linear(hidden_size, 1, bias=True)
|
||||
# Cast to the head weight dtype so load_state_dict matches.
|
||||
reward_head.load_state_dict(
|
||||
{
|
||||
"weight": head_state["weight"].to(torch.float32),
|
||||
"bias": head_state["bias"].to(torch.float32),
|
||||
}
|
||||
)
|
||||
reward_head = reward_head.to(dtype)
|
||||
model.reward_head = reward_head
|
||||
model.eval()
|
||||
model.requires_grad_(False)
|
||||
model.to(self.device)
|
||||
self._model = model
|
||||
self._tokenizer = tokenizer
|
||||
|
||||
def _render_prompt(self, prompt: Any) -> str:
|
||||
"""Best-effort render of a GRPO prompt (str or message list) to text."""
|
||||
if isinstance(prompt, str):
|
||||
return prompt
|
||||
if isinstance(prompt, (list, tuple)):
|
||||
parts: list[str] = []
|
||||
for msg in prompt:
|
||||
if isinstance(msg, dict):
|
||||
parts.append(str(msg.get("content", "")))
|
||||
else:
|
||||
parts.append(str(msg))
|
||||
return "\n".join(p for p in parts if p)
|
||||
return ""
|
||||
|
||||
def _completion_text(self, completion: Any) -> str:
|
||||
if isinstance(completion, str):
|
||||
return completion
|
||||
if isinstance(completion, dict):
|
||||
return str(completion.get("content", ""))
|
||||
if isinstance(completion, (list, tuple)):
|
||||
parts = [
|
||||
str(m.get("content", "")) if isinstance(m, dict) else str(m)
|
||||
for m in completion
|
||||
]
|
||||
return "".join(parts)
|
||||
return str(completion)
|
||||
|
||||
def _score_one(self, prompt_text: str, steps: list[str]) -> float:
|
||||
import torch
|
||||
|
||||
if not steps:
|
||||
return 0.0
|
||||
tokenizer = self._tokenizer
|
||||
# Prompt context (trained distribution) then step boundaries.
|
||||
prefix_ids = (
|
||||
tokenizer(prompt_text, add_special_tokens=False)["input_ids"]
|
||||
if prompt_text
|
||||
else []
|
||||
)
|
||||
input_ids = list(prefix_ids)
|
||||
step_positions: list[int] = []
|
||||
for step in steps:
|
||||
step_ids = tokenizer(step, add_special_tokens=False)["input_ids"]
|
||||
if not step_ids:
|
||||
continue
|
||||
input_ids.extend(step_ids)
|
||||
step_positions.append(len(input_ids) - 1)
|
||||
if not step_positions:
|
||||
return 0.0
|
||||
ids = torch.tensor([input_ids], dtype=torch.long, device=self.device)
|
||||
with torch.no_grad():
|
||||
outputs = self._model(input_ids=ids, output_hidden_states=True)
|
||||
last_hidden = outputs.hidden_states[-1][0] # [T, H]
|
||||
pos = torch.tensor(step_positions, dtype=torch.long, device=self.device)
|
||||
step_hidden = last_hidden.index_select(0, pos) # [S, H]
|
||||
scores = self._model.reward_head(step_hidden).squeeze(-1) # [S]
|
||||
per_step = [float(s) for s in scores.detach().float().cpu().tolist()]
|
||||
return aggregate_step_scores(per_step, self.aggregate)
|
||||
|
||||
def __call__(self, completions: Any, **kwargs: Any) -> list[float]:
|
||||
self._ensure_loaded()
|
||||
prompts = kwargs.get("prompts")
|
||||
rewards: list[float] = []
|
||||
completion_list = list(completions) if completions is not None else []
|
||||
for i, completion in enumerate(completion_list):
|
||||
prompt = prompts[i] if isinstance(prompts, (list, tuple)) and i < len(prompts) else None
|
||||
prompt_text = self._render_prompt(prompt)
|
||||
steps = split_steps(self._completion_text(completion))
|
||||
rewards.append(self._score_one(prompt_text, steps))
|
||||
return rewards
|
||||
|
||||
|
||||
def load_reward_head_weights(prm_path: str) -> dict:
|
||||
"""Load ``reward_head.{weight,bias}`` tensors from a Soup-trained PRM dir.
|
||||
|
||||
Scans every ``*.safetensors`` shard in ``prm_path`` via
|
||||
``safetensors.safe_open`` and collects the ``reward_head.*`` tensors.
|
||||
Raises ``ValueError`` (friendly "not a Soup-trained PRM") when absent —
|
||||
``AutoModelForCausalLM.from_pretrained`` silently drops these keys, so a
|
||||
base checkpoint without a head must be rejected loudly.
|
||||
"""
|
||||
import os
|
||||
|
||||
from safetensors import safe_open
|
||||
|
||||
if not os.path.isdir(prm_path):
|
||||
raise ValueError(
|
||||
f"prm_reward path is not a directory: {prm_path!r}. Expected a "
|
||||
"Soup-trained PRM produced by `soup train` with task='prm'."
|
||||
)
|
||||
collected: dict = {}
|
||||
for entry in sorted(os.listdir(prm_path)):
|
||||
if not entry.endswith(".safetensors"):
|
||||
continue
|
||||
shard = os.path.join(prm_path, entry)
|
||||
with safe_open(shard, framework="pt") as handle:
|
||||
for key in handle.keys(): # noqa: SIM118 — safe_open handle API
|
||||
if key.startswith("reward_head."):
|
||||
collected[key[len("reward_head."):]] = handle.get_tensor(key)
|
||||
if "weight" not in collected or "bias" not in collected:
|
||||
raise ValueError(
|
||||
f"No reward_head weights found in {prm_path!r} — this is not a "
|
||||
"Soup-trained PRM. Train one with `soup train` (task='prm') first."
|
||||
)
|
||||
return collected
|
||||
|
||||
|
||||
def build_prm_reward_fn(
|
||||
tcfg: Any,
|
||||
device: str,
|
||||
trust_remote_code: bool,
|
||||
) -> PRMScorer:
|
||||
"""Build the :class:`PRMScorer` for GRPO from a ``TrainingConfig``.
|
||||
|
||||
Validates local-path containment (realpath + commonpath under cwd) and
|
||||
surfaces a ``trust_remote_code`` probe/warning, then returns the scorer.
|
||||
A non-existent local path is treated as a Hugging Face repo id (loaded via
|
||||
``from_pretrained``); only *existing local paths* are containment-checked.
|
||||
"""
|
||||
import os
|
||||
|
||||
from rich.console import Console
|
||||
|
||||
console = Console()
|
||||
prm_path = tcfg.prm_reward
|
||||
if prm_path is None:
|
||||
raise ValueError("build_prm_reward_fn called with prm_reward=None")
|
||||
|
||||
# Containment: only enforce for a path that exists on disk. A bare repo id
|
||||
# (no local existence) is handled by from_pretrained's own network path.
|
||||
if os.path.exists(prm_path):
|
||||
real = os.path.realpath(prm_path)
|
||||
cwd = os.path.realpath(os.getcwd())
|
||||
if os.path.commonpath([real, cwd]) != cwd:
|
||||
raise ValueError(
|
||||
"prm_reward path must stay under the current working "
|
||||
f"directory; got {prm_path!r}"
|
||||
)
|
||||
prm_path = real
|
||||
|
||||
resolved_trust = _resolve_trust(prm_path, trust_remote_code, console)
|
||||
return PRMScorer(
|
||||
prm_path=prm_path,
|
||||
aggregate=tcfg.prm_aggregate,
|
||||
device=device,
|
||||
trust_remote_code=resolved_trust,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_trust(base: str, requested: bool, console: Any) -> bool:
|
||||
"""Trust-remote-code probe + warn, mirroring the trainer convention."""
|
||||
from soup_cli.utils.trust_remote import (
|
||||
model_requires_trust_remote_code,
|
||||
resolve_trust_remote_code,
|
||||
)
|
||||
|
||||
requires = model_requires_trust_remote_code(base) or False
|
||||
return resolve_trust_remote_code(
|
||||
base,
|
||||
requested=requested,
|
||||
console=console,
|
||||
requires_remote_code=requires,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AGGREGATE_MODES",
|
||||
"PRMScorer",
|
||||
"aggregate_step_scores",
|
||||
"build_prm_reward_fn",
|
||||
"load_reward_head_weights",
|
||||
"split_steps",
|
||||
]
|
||||
|
|
@ -0,0 +1,103 @@
|
|||
"""v0.71.30 — PRM-guided GRPO + bundled rollout envs.
|
||||
|
||||
Tests the pure PRM reward kernels (split_steps / aggregate_step_scores), the
|
||||
schema fields + cross-validators, the torch-lazy PRMScorer (safetensors head
|
||||
load + scoring), the GRPO wiring, the bundled rollout envs, and the recipes.
|
||||
"""
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
import pytest
|
||||
|
||||
from soup_cli.utils.prm_reward import (
|
||||
AGGREGATE_MODES,
|
||||
aggregate_step_scores,
|
||||
split_steps,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 1 — pure kernels
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestSplitSteps:
|
||||
def test_splits_on_newlines(self):
|
||||
assert split_steps("a\nb\nc") == ["a", "b", "c"]
|
||||
|
||||
def test_drops_empty_and_whitespace(self):
|
||||
assert split_steps("a\n\n \nb\n") == ["a", "b"]
|
||||
|
||||
def test_strips_each_step(self):
|
||||
assert split_steps(" a \n\tb\t") == ["a", "b"]
|
||||
|
||||
def test_non_string_returns_empty(self):
|
||||
assert split_steps(None) == []
|
||||
assert split_steps(123) == []
|
||||
|
||||
def test_empty_returns_empty(self):
|
||||
assert split_steps("") == []
|
||||
assert split_steps(" \n ") == []
|
||||
|
||||
def test_caps_step_count(self):
|
||||
from soup_cli.utils.prm_reward import _MAX_STEPS
|
||||
|
||||
text = "\n".join(str(i) for i in range(_MAX_STEPS + 50))
|
||||
assert len(split_steps(text)) == _MAX_STEPS
|
||||
|
||||
def test_caps_step_chars(self):
|
||||
from soup_cli.utils.prm_reward import _MAX_STEP_CHARS
|
||||
|
||||
long = "x" * (_MAX_STEP_CHARS + 100)
|
||||
out = split_steps(long)
|
||||
assert len(out) == 1
|
||||
assert len(out[0]) == _MAX_STEP_CHARS
|
||||
|
||||
|
||||
class TestAggregate:
|
||||
def test_min(self):
|
||||
assert aggregate_step_scores([0.9, 0.2, 0.7], "min") == pytest.approx(0.2)
|
||||
|
||||
def test_last(self):
|
||||
assert aggregate_step_scores([0.9, 0.2, 0.7], "last") == pytest.approx(0.7)
|
||||
|
||||
def test_prod(self):
|
||||
assert aggregate_step_scores([0.5, 0.5, 0.5], "prod") == pytest.approx(0.125)
|
||||
|
||||
def test_empty_returns_zero(self):
|
||||
assert aggregate_step_scores([], "min") == 0.0
|
||||
assert aggregate_step_scores([], "prod") == 0.0
|
||||
assert aggregate_step_scores([], "last") == 0.0
|
||||
|
||||
def test_single(self):
|
||||
assert aggregate_step_scores([0.42], "min") == pytest.approx(0.42)
|
||||
|
||||
def test_non_finite_is_safe(self):
|
||||
# NaN / inf must not propagate — coerced to 0.0
|
||||
out = aggregate_step_scores([float("nan"), 0.5], "min")
|
||||
assert out == 0.0
|
||||
|
||||
def test_bad_mode_rejected(self):
|
||||
with pytest.raises(ValueError, match="min|prod|last"):
|
||||
aggregate_step_scores([0.5], "mean")
|
||||
|
||||
def test_bool_mode_rejected(self):
|
||||
with pytest.raises(ValueError):
|
||||
aggregate_step_scores([0.5], True)
|
||||
|
||||
def test_aggregate_modes_constant(self):
|
||||
assert set(AGGREGATE_MODES) == {"min", "prod", "last"}
|
||||
|
||||
|
||||
class TestNoTopLevelTorch:
|
||||
def test_prm_reward_has_no_top_level_torch(self):
|
||||
import soup_cli.utils.prm_reward as mod
|
||||
|
||||
source = inspect.getsource(mod)
|
||||
tree = ast.parse(source)
|
||||
heavy = {"torch", "transformers", "peft", "safetensors"}
|
||||
for node in tree.body:
|
||||
if isinstance(node, ast.Import):
|
||||
for alias in node.names:
|
||||
assert alias.name.split(".")[0] not in heavy, alias.name
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
root = (node.module or "").split(".")[0]
|
||||
assert root not in heavy, node.module
|
||||
Loading…
Reference in New Issue