Skip to content
Merged
Show file tree
Hide file tree
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
18 changes: 17 additions & 1 deletion src/prime_rl/orchestrator/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,17 @@ def setup(self) -> Stat:
def generation(self) -> Stat:
return Stat([r.timing.generation.duration for r in self.rollouts])

@property
def generation_model(self) -> Stat:
"""The share of the generation phase spent inside model calls (inference)."""
return Stat([r.timing.generation.model.duration for r in self.rollouts])

@property
def generation_harness(self) -> Stat:
"""The share of the generation phase spent outside model calls (harness, tools,
user simulation)."""
return Stat([r.timing.generation.harness.duration for r in self.rollouts])

@property
def finalize(self) -> Stat:
return Stat([r.timing.finalize.duration for r in self.rollouts])
Expand All @@ -113,7 +124,12 @@ def total(self) -> Stat:
return Stat([sum(getattr(r.timing, p).duration for p in self.PHASES) for r in self.rollouts])

def stats(self) -> dict[str, Stat]:
return {**{phase: getattr(self, phase) for phase in self.PHASES}, "total": self.total}
return {
**{phase: getattr(self, phase) for phase in self.PHASES},
"generation/model": self.generation_model,
"generation/harness": self.generation_harness,
"total": self.total,
}


class CustomMetrics(StatGroup):
Expand Down
12 changes: 10 additions & 2 deletions tests/unit/orchestrator/test_advantage.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,11 +55,12 @@ def _take(n: int) -> list[int]:
)
parent = len(nodes) - 1

# Trace token counts are usage-based, so carry provider usage on the final sampled turn:
# Trace token counts are usage-based, so carry provider usage on the final turn's call:
# every model-generated token as completion, the leading prompt + tool observations as the
# fed-in context (num_input_tokens = num_total_tokens - num_output_tokens).
output_tokens = sum(sampled_lengths)
input_tokens = 1 + sum(obs_lengths)
calls: list[vf.ModelCall] = []

for i, n_sampled in enumerate(sampled_lengths):
ids = _take(n_sampled)
Expand All @@ -72,10 +73,16 @@ def _take(n: int) -> list[int]:
logprobs=[-0.1] * n_sampled,
sampled=True,
parent=parent,
usage=vf.Usage(prompt_tokens=input_tokens, completion_tokens=output_tokens) if is_last else None,
)
)
parent = len(nodes) - 1
if is_last:
calls.append(
vf.ModelCall(
node=parent,
usage=vf.Usage(prompt_tokens=input_tokens, completion_tokens=output_tokens),
)
)
if i < len(obs_lengths):
obs_ids = _take(obs_lengths[i])
nodes.append(
Expand All @@ -93,6 +100,7 @@ def _take(n: int) -> list[int]:
rollout = Rollout[vf.TaskData](
task=vf.TraceTask(type="Task", data=vf.TaskData(idx=0, prompt=None)),
nodes=nodes,
calls=calls,
rewards={"reward": reward},
metrics=metrics or {},
)
Expand Down
15 changes: 13 additions & 2 deletions tests/unit/orchestrator/test_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,8 @@ def mk(
filter_results: dict | None = None,
setup: float = 0.0,
generation: float = 0.0,
generation_model: float = 0.0,
generation_harness: float = 0.0,
finalize: float = 0.0,
scoring: float = 0.0,
):
Expand All @@ -54,7 +56,11 @@ def mk(
filter_results=filter_results or {},
timing=SimpleNamespace(
setup=SimpleNamespace(duration=setup),
generation=SimpleNamespace(duration=generation),
generation=SimpleNamespace(
duration=generation,
model=SimpleNamespace(duration=generation_model),
harness=SimpleNamespace(duration=generation_harness),
),
finalize=SimpleNamespace(duration=finalize),
scoring=SimpleNamespace(duration=scoring),
),
Expand Down Expand Up @@ -145,11 +151,16 @@ def test_nested_metrics_and_rewards():


def test_nested_timing():
m = TrainRollouts([mk(setup=1.0, generation=2.0, finalize=0.5, scoring=0.5)]).metrics
m = TrainRollouts(
[mk(setup=1.0, generation=2.0, generation_model=1.5, generation_harness=0.5, finalize=0.5, scoring=0.5)]
).metrics
assert m.timing.setup.mean() == 1.0 and m.timing.total.mean() == 4.0 # total sums all four phases
assert m.timing.generation_model.mean() == 1.5 and m.timing.generation_harness.mean() == 0.5
out = m.to_wandb(prefix="train/agg", subset="all")
assert out["train/agg/all/timing/setup/mean"] == 1.0
assert out["train/agg/all/timing/total/mean"] == 4.0
assert out["train/agg/all/timing/generation/model/mean"] == 1.5
assert out["train/agg/all/timing/generation/harness/mean"] == 0.5


def test_train_only_metrics_absent_from_eval():
Expand Down