diff --git a/src/soup_cli/utils/prm_reward.py b/src/soup_cli/utils/prm_reward.py new file mode 100644 index 0000000..73b7e48 --- /dev/null +++ b/src/soup_cli/utils/prm_reward.py @@ -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/`` 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", +] diff --git a/tests/test_v07130.py b/tests/test_v07130.py new file mode 100644 index 0000000..b3de5b7 --- /dev/null +++ b/tests/test_v07130.py @@ -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