feat(prm): PRM reward pure kernels (split_steps/aggregate) (v0.71.30)

This commit is contained in:
Alpamys 2026-07-05 16:49:51 +05:00
parent 3c21c5e7c8
commit 79e6c119de
2 changed files with 433 additions and 0 deletions

View File

@ -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",
]

103
tests/test_v07130.py Normal file
View File

@ -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