From d8b79b1204ad513dac3edf39edb00b5d2487d9e3 Mon Sep 17 00:00:00 2001 From: Ekaanksh Patil Date: Wed, 8 Jul 2026 21:34:15 +0530 Subject: [PATCH] Fix/prm reward (#301) * chore: auto-add Punk Records notifier * feat(prm): batch forward pass in PRMScorer.__call__ * fix(prm): resolve ruff warnings in PRMScorer.__call__ * Delete .github/workflows/notify.yml * Optimize PRMScorer batched reward inference * style(test): rename B -> batch in PRM fake models (ruff N806) --------- Co-authored-by: Alpamys --- src/soup_cli/utils/prm_reward.py | 77 ++++++++++++++++++++++++++++++-- tests/test_v07130.py | 26 ++++++++--- 2 files changed, 94 insertions(+), 9 deletions(-) diff --git a/src/soup_cli/utils/prm_reward.py b/src/soup_cli/utils/prm_reward.py index f873af4..7372229 100644 --- a/src/soup_cli/utils/prm_reward.py +++ b/src/soup_cli/utils/prm_reward.py @@ -246,17 +246,86 @@ class PRMScorer: return aggregate_step_scores(per_step, self.aggregate) def __call__(self, completions: Any, **kwargs: Any) -> list[float]: + import torch + 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 + if not completion_list: + return [] + config = getattr(self._model,"config",None) + max_pos = getattr(config,"max_position_embeddings",_MAX_INPUT_TOKENS) + cap = _MAX_INPUT_TOKENS + if isinstance(max_pos,int)and not isinstance(max_pos,bool): + cap = min(max_pos,_MAX_INPUT_TOKENS) + all_input_ids: list[list[int]]=[] + all_step_positions: list[list[int]]=[] + skip_indices: set[int]= set() + + for i,completion in enumerate(completion_list): + prompt = prompts[i] if isinstance(prompts,(list, tuple)) and icap: + input_ids=input_ids[:cap] + step_positions = [p for p in step_positions if p dict[str, Any]: """Load ``reward_head.{weight,bias}`` tensors from a Soup-trained PRM dir. diff --git a/tests/test_v07130.py b/tests/test_v07130.py index d1f8e38..4e2acf9 100644 --- a/tests/test_v07130.py +++ b/tests/test_v07130.py @@ -269,9 +269,10 @@ def _make_fake_scorer(aggregate="min"): class _FakeModel: reward_head = head - def __call__(self, input_ids, output_hidden_states=False): + def __call__(self, input_ids, output_hidden_states=False, **kwargs): seq_len = input_ids.shape[1] - hs = torch.arange(seq_len).float().reshape(1, seq_len, 1).repeat(1, 1, hidden) + batch = input_ids.shape[0] + hs = torch.arange(seq_len).float().reshape(1, seq_len, 1).expand(batch, seq_len, hidden) return SimpleNamespace(hidden_states=[hs]) scorer = PRMScorer("./prm", aggregate=aggregate, device="cpu") @@ -358,6 +359,21 @@ class TestPRMScorer: assert len(out) == 1 assert isinstance(out[0], float) + def test_batched_parity_with_per_completion(self): + s = _make_fake_scorer("min") + completions = ["a\nb", "c\nd\ne"] + batched = s(completions) + for i, c in enumerate(completions): + single = s([c]) + assert abs(batched[i] - single[0]) < 1e-5, f"mismatch at {i}" + + def test_mixed_length_no_cross_row_contamination(self): + s = _make_fake_scorer("min") + short = s(["a\nb"]) + long_ = s(["c\nd\ne\nf\ng"]) + both = s(["a\nb", "c\nd\ne\nf\ng"]) + assert abs(both[0] - short[0]) < 1e-5 + assert abs(both[1] - long_[0]) < 1e-5 def _make_capped_scorer(max_pos, aggregate="min"): """PRMScorer whose fake model advertises a tiny max_position_embeddings.""" @@ -387,9 +403,10 @@ def _make_capped_scorer(max_pos, aggregate="min"): reward_head = head config = SimpleNamespace(max_position_embeddings=max_pos) - def __call__(self, input_ids, output_hidden_states=False): + def __call__(self, input_ids, output_hidden_states=False, **kwargs): seq_len = input_ids.shape[1] - hs = torch.arange(seq_len).float().reshape(1, seq_len, 1).repeat(1, 1, hidden) + batch = input_ids.shape[0] + hs = torch.arange(seq_len).float().reshape(1, seq_len, 1).expand(batch, seq_len, hidden) return SimpleNamespace(hidden_states=[hs]) scorer = PRMScorer("./prm", aggregate=aggregate, device="cpu") @@ -397,7 +414,6 @@ def _make_capped_scorer(max_pos, aggregate="min"): scorer._tokenizer = _FakeTok() return scorer - class TestPRMScorerInputCap: def test_truncates_to_max_position_embeddings(self): # steps "a b"(2 tok)->boundary 1, "c d"(2 tok)->boundary 3; cap=3 keeps