Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 40 additions & 7 deletions verifiers/rubrics/rubric_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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()
}