diff --git a/verifiers/rubrics/rubric_group.py b/verifiers/rubrics/rubric_group.py index 85807f9023..b361b1dd63 100644 --- a/verifiers/rubrics/rubric_group.py +++ b/verifiers/rubrics/rubric_group.py @@ -114,14 +114,33 @@ async def teardown(self): async def score_group(self, states: list[State]): """ Evaluate all reward functions in-place for a group of rollouts. + + Mirrors ``score_rollout`` for metrics (full replace) and ``Rubric.score_group`` + for advantages: child rubrics may write per-rubric advantages/trajectory fields, + so the group recomputes state advantages from aggregated rewards and only fills + trajectory fields that are still ``None`` (preserving intentional step values). """ - aggregated_rewards = [0.0] * len(states) + num_states = len(states) + if num_states == 0: + self.logger.warning("No states to score") + return + + aggregated_rewards = [0.0] * num_states aggregated_metrics: dict[str, list[float]] = {} original_rewards = [state.get("reward", 0.0) for state in states] original_metrics = [ state.get("metrics", {}).copy() if state.get("metrics") else {} for state in states ] + original_advantages = [state.get("advantage") for state in states] + # Snapshot traj fields so child if-None fills can be rolled back between rubrics. + original_traj_fields = [ + [ + (step.get("reward"), step.get("advantage")) + for step in state.get("trajectory") or [] + ] + for state in states + ] for rubric in self.rubrics: await rubric.score_group(states) for i, state in enumerate(states): @@ -132,14 +151,28 @@ async def score_group(self, states: list[State]): aggregated_rewards[i] += rubric_reward for key, value in rubric_metrics.items(): if key not in aggregated_metrics: - aggregated_metrics[key] = [0.0] * len(states) + aggregated_metrics[key] = [0.0] * num_states aggregated_metrics[key][i] += value + # Restore so the next rubric sees the same pre-group inputs. state["reward"] = original_rewards[i] state["metrics"] = original_metrics[i].copy() + state["advantage"] = original_advantages[i] + for step, (reward, advantage) in zip( + state.get("trajectory") or [], original_traj_fields[i] + ): + step["reward"] = reward + step["advantage"] = advantage + + avg_reward = sum(aggregated_rewards) / num_states for i, state in enumerate(states): state["reward"] = aggregated_rewards[i] - if aggregated_metrics: - if "metrics" not in state: - state["metrics"] = {} - for key, values in aggregated_metrics.items(): - state["metrics"][key] = values[i] + state["advantage"] = aggregated_rewards[i] - avg_reward + # Same as Rubric.score_group: fill only unset traj fields. + for step in state.get("trajectory") or []: + if step.get("advantage") is None: + step["advantage"] = state["advantage"] + if step.get("reward") is None: + step["reward"] = state["reward"] + state["metrics"] = { + key: values[i] for key, values in aggregated_metrics.items() + }