Skip to content

Commit 6493bd9

Browse files
authored
fix: fix rollout metrics fail in unit test (#3173)
Signed-off-by: Yuki Huang <yukih@nvidia.com>
1 parent 1f4d989 commit 6493bd9

6 files changed

Lines changed: 176 additions & 71 deletions

File tree

nemo_rl/experience/metric_utils.py

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,53 @@
1+
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# http://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
"""Shared aggregation helpers for rollout metrics."""
16+
17+
import math
18+
import statistics
19+
from collections.abc import Sequence
20+
21+
from wandb import Histogram
22+
23+
24+
def calculate_single_metric(
25+
values: Sequence[float | int], batch_size: int, key_name: str
26+
) -> dict:
27+
"""Compute summary statistics for a metric as slash-prefixed keys.
28+
29+
Args:
30+
values: Per-sample metric values to aggregate.
31+
batch_size: Denominator for the mean (sum(values) / batch_size, not len(values)); stddev still uses len(values).
32+
key_name: Prefix for the returned metric keys (e.g. "total_reward").
33+
34+
Returns:
35+
Dict mapping "{key_name}/{stat}" to its value for stat in mean, max, min, median, stddev (nan for a single value), and histogram (a wandb.Histogram).
36+
"""
37+
return {
38+
f"{key_name}/mean": sum(values) / batch_size,
39+
f"{key_name}/max": max(values),
40+
f"{key_name}/min": min(values),
41+
f"{key_name}/median": statistics.median(values),
42+
f"{key_name}/stddev": statistics.stdev(values) if len(values) > 1 else math.nan,
43+
f"{key_name}/histogram": Histogram(values),
44+
}
45+
46+
47+
def pct(values: Sequence[float | int], p: float) -> float:
48+
"""Percentile helper for buffer starvation diagnostics."""
49+
if not values:
50+
return 0.0
51+
sorted_v = sorted(values)
52+
idx = min(int(len(sorted_v) * p / 100), len(sorted_v) - 1)
53+
return float(sorted_v[idx])

nemo_rl/experience/rollout_manager.py

Lines changed: 43 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -25,11 +25,8 @@
2525
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
2626
from nemo_rl.environments.interfaces import EnvironmentInterface
2727
from nemo_rl.experience.interfaces import Completion, PromptGroupRecord
28-
from nemo_rl.experience.rollouts import (
29-
_calculate_single_metric,
30-
_tensorize_by_key,
31-
calculate_rewards,
32-
)
28+
from nemo_rl.experience.metric_utils import calculate_single_metric, pct
29+
from nemo_rl.experience.rollouts import _tensorize_by_key, calculate_rewards
3330
from nemo_rl.models.generation.interfaces import (
3431
GenerationConfig,
3532
GenerationDatumSpec,
@@ -337,17 +334,29 @@ def _aggregate_rollout_metrics(
337334
terminated = [m["terminated"] for m in all_sample_metrics]
338335
max_turns_reached = [m["max_turns_reached"] for m in all_sample_metrics]
339336

337+
# max_gen_tokens_per_turn: Diagnostic for long single generations
338+
max_gen_tokens_per_turn = [
339+
max(m["turn_gen_tokens"]) if m["turn_gen_tokens"] else 0
340+
for m in all_sample_metrics
341+
]
342+
340343
# Aggregate metrics across all samples.
341344
n = len(all_sample_metrics)
342345
rollout_metrics: dict[str, Any] = {
343-
**_calculate_single_metric(total_reward, n, "total_reward"),
346+
**calculate_single_metric(total_reward, n, "total_reward"),
344347
# turn metrics
345348
"total_turns": sum(turn_count),
346-
**_calculate_single_metric(turn_count, n, "turns_per_sample"),
349+
**calculate_single_metric(turn_count, n, "turns_per_sample"),
350+
"turns_per_sample/p95": pct(turn_count, 95),
351+
"turns_per_sample/p99": pct(turn_count, 99),
347352
# token metrics
348-
**_calculate_single_metric(total_tokens, n, "total_tokens_per_sample"),
349-
**_calculate_single_metric(assistant_tokens, n, "gen_tokens_per_sample"),
350-
**_calculate_single_metric(env_tokens, n, "env_tokens_per_sample"),
353+
**calculate_single_metric(total_tokens, n, "total_tokens_per_sample"),
354+
**calculate_single_metric(assistant_tokens, n, "gen_tokens_per_sample"),
355+
**calculate_single_metric(env_tokens, n, "env_tokens_per_sample"),
356+
# max_gen_tokens_per_turn: Diagnostic for long single generations
357+
"max_gen_tokens_per_turn/max": max(max_gen_tokens_per_turn),
358+
"max_gen_tokens_per_turn/mean": sum(max_gen_tokens_per_turn) / n,
359+
"max_gen_tokens_per_turn/p95": pct(max_gen_tokens_per_turn, 95),
351360
# truncated metrics
352361
"truncation_rate": sum(truncated) / n,
353362
"natural_termination_rate": sum(terminated) / n,
@@ -362,7 +371,7 @@ def _aggregate_rollout_metrics(
362371
rollout_metrics["per_worker_token_counts"] = per_worker_token_counts
363372

364373
# Per-turn token histograms (flat across all turns, distinct from the
365-
# per-sample histograms emitted via _calculate_single_metric above).
374+
# per-sample histograms emitted via calculate_single_metric above).
366375
rollout_metrics["histogram/gen_tokens_length"] = [
367376
t for m in all_sample_metrics for t in m["turn_gen_tokens"]
368377
]
@@ -542,18 +551,36 @@ def _compute_rollout_metrics(
542551
sum(len(m["token_ids"]) for m in c.message_log if m["role"] == "assistant")
543552
for c in completions
544553
]
554+
# max_gen_tokens_per_turn: Diagnostic for long single generations
555+
max_gen_tokens_per_turn = [
556+
max(
557+
(
558+
len(m["token_ids"])
559+
for m in c.message_log
560+
if m["role"] == "assistant"
561+
),
562+
default=0,
563+
)
564+
for c in completions
565+
]
545566
# truncated metrics
546567
truncated = [c.truncated for c in completions]
547568

548569
# Aggregate metrics across all samples.
549570
n = len(completions)
550571
rollout_metrics: dict[str, Any] = {
551-
**_calculate_single_metric(total_reward, n, "total_reward"),
572+
**calculate_single_metric(total_reward, n, "total_reward"),
552573
# turn metrics
553-
**_calculate_single_metric(turn_count, n, "turns_per_sample"),
574+
**calculate_single_metric(turn_count, n, "turns_per_sample"),
575+
"turns_per_sample/p95": pct(turn_count, 95),
576+
"turns_per_sample/p99": pct(turn_count, 99),
554577
# token metrics
555-
**_calculate_single_metric(total_tokens, n, "total_tokens_per_sample"),
556-
**_calculate_single_metric(assistant_tokens, n, "gen_tokens_per_sample"),
578+
**calculate_single_metric(total_tokens, n, "total_tokens_per_sample"),
579+
**calculate_single_metric(assistant_tokens, n, "gen_tokens_per_sample"),
580+
**calculate_single_metric(
581+
max_gen_tokens_per_turn, n, "max_gen_tokens_per_turn"
582+
),
583+
"max_gen_tokens_per_turn/p95": pct(max_gen_tokens_per_turn, 95),
557584
# truncated metrics
558585
"natural_termination_rate": sum(not t for t in truncated) / n,
559586
"truncation_rate": sum(truncated) / n,
@@ -569,7 +596,7 @@ def _compute_rollout_metrics(
569596
]
570597
if values:
571598
rollout_metrics.update(
572-
_calculate_single_metric(values, n, f"{agent_name}/{key}")
599+
calculate_single_metric(values, n, f"{agent_name}/{key}")
573600
)
574601
rollout_metrics[f"{agent_name}/full_result"] = Table(
575602
data=[[json.dumps(r, separators=(",", ":"))] for r in agent_extras],

nemo_rl/experience/rollouts.py

Lines changed: 14 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -18,19 +18,17 @@
1818
import asyncio
1919
import copy
2020
import json
21-
import math
2221
import statistics
2322
import warnings
2423
from collections import defaultdict
25-
from collections.abc import Sequence
2624
from dataclasses import dataclass
2725
from typing import Any, Optional
2826

2927
import ray
3028
import torch
3129
from pydantic import BaseModel
3230
from transformers import PreTrainedTokenizerBase
33-
from wandb import Histogram, Table
31+
from wandb import Table
3432

3533
from nemo_rl.algorithms.utils import get_gdpo_reward_component_keys
3634
from nemo_rl.data.interfaces import (
@@ -48,6 +46,7 @@
4846
EnvironmentReturn,
4947
)
5048
from nemo_rl.environments.nemo_gym import DEFAULT_THINKING_TAGS
49+
from nemo_rl.experience.metric_utils import calculate_single_metric, pct
5150
from nemo_rl.models.generation.interfaces import (
5251
GenerationConfig,
5352
GenerationDatumSpec,
@@ -1263,14 +1262,6 @@ async def run_single_sample_with_error_handling(i, sample_state):
12631262
if key not in final_batch:
12641263
final_batch[key] = input_batch[key]
12651264

1266-
# Helper for percentile (buffer starvation diagnostics)
1267-
def _pct(values: Sequence[float | int], p: float) -> float:
1268-
if not values:
1269-
return 0.0
1270-
sorted_v = sorted(values)
1271-
idx = min(int(len(sorted_v) * p / 100), len(sorted_v) - 1)
1272-
return float(sorted_v[idx])
1273-
12741265
turn_counts = [m["turn_count"] for m in all_sample_metrics]
12751266
max_gen_tokens_per_turn_values = [
12761267
m["max_gen_tokens_per_turn"] for m in all_sample_metrics
@@ -1282,8 +1273,8 @@ def _pct(values: Sequence[float | int], p: float) -> float:
12821273
"total_turns": sum(turn_counts),
12831274
"avg_turns_per_sample": sum(turn_counts) / batch_size,
12841275
"max_turns_per_sample": max(turn_counts),
1285-
"turns_per_sample/p95": _pct(turn_counts, 95),
1286-
"turns_per_sample/p99": _pct(turn_counts, 99),
1276+
"turns_per_sample/p95": pct(turn_counts, 95),
1277+
"turns_per_sample/p99": pct(turn_counts, 99),
12871278
"natural_termination_rate": sum(m["terminated"] for m in all_sample_metrics)
12881279
/ batch_size,
12891280
"truncation_rate": sum(m["truncated"] for m in all_sample_metrics)
@@ -1312,7 +1303,7 @@ def _pct(values: Sequence[float | int], p: float) -> float:
13121303
"max_gen_tokens_per_turn/max": max(max_gen_tokens_per_turn_values),
13131304
"max_gen_tokens_per_turn/mean": sum(max_gen_tokens_per_turn_values)
13141305
/ batch_size,
1315-
"max_gen_tokens_per_turn/p95": _pct(max_gen_tokens_per_turn_values, 95),
1306+
"max_gen_tokens_per_turn/p95": pct(max_gen_tokens_per_turn_values, 95),
13161307
# Reward metrics
13171308
"mean_total_reward": sum(m["total_reward"] for m in all_sample_metrics)
13181309
/ batch_size,
@@ -1359,19 +1350,6 @@ class AsyncNemoGymRolloutResult:
13591350
rollout_metrics: dict[str, Any]
13601351

13611352

1362-
def _calculate_single_metric(
1363-
values: Sequence[float | int], batch_size: int, key_name: str
1364-
) -> dict:
1365-
return {
1366-
f"{key_name}/mean": sum(values) / batch_size,
1367-
f"{key_name}/max": max(values),
1368-
f"{key_name}/min": min(values),
1369-
f"{key_name}/median": statistics.median(values),
1370-
f"{key_name}/stddev": statistics.stdev(values) if len(values) > 1 else math.nan,
1371-
f"{key_name}/histogram": Histogram(values),
1372-
}
1373-
1374-
13751353
def get_nemo_gym_thinking_tags(env_config: dict[str, Any]) -> list[str]:
13761354
"""Return thinking tags used by the Gym-side detector."""
13771355
nemo_gym_config = env_config.get("nemo_gym")
@@ -1918,39 +1896,32 @@ def run_async_nemo_gym_rollout(
19181896
m["max_gen_tokens_per_turn"] for m in all_sample_metrics
19191897
]
19201898

1921-
def _pct(values: Sequence[float | int], p: float) -> float:
1922-
if not values:
1923-
return 0.0
1924-
sorted_v = sorted(values)
1925-
idx = min(int(len(sorted_v) * p / 100), len(sorted_v) - 1)
1926-
return float(sorted_v[idx])
1927-
19281899
rollout_metrics = {
19291900
**rollout_loop_timing_metrics,
1930-
**_calculate_single_metric(
1901+
**calculate_single_metric(
19311902
turn_counts,
19321903
batch_size,
19331904
"turns_per_sample",
19341905
),
1935-
"turns_per_sample/p95": _pct(turn_counts, 95),
1936-
"turns_per_sample/p99": _pct(turn_counts, 99),
1937-
**_calculate_single_metric(
1906+
"turns_per_sample/p95": pct(turn_counts, 95),
1907+
"turns_per_sample/p99": pct(turn_counts, 99),
1908+
**calculate_single_metric(
19381909
[m["total_tokens"] for m in all_sample_metrics],
19391910
batch_size,
19401911
"total_tokens_per_sample",
19411912
),
1942-
**_calculate_single_metric(
1913+
**calculate_single_metric(
19431914
[m["assistant_tokens"] for m in all_sample_metrics],
19441915
batch_size,
19451916
"gen_tokens_per_sample",
19461917
),
1947-
**_calculate_single_metric(
1918+
**calculate_single_metric(
19481919
max_gen_tokens_per_turn_values,
19491920
batch_size,
19501921
"max_gen_tokens_per_turn",
19511922
),
1952-
"max_gen_tokens_per_turn/p95": _pct(max_gen_tokens_per_turn_values, 95),
1953-
**_calculate_single_metric(
1923+
"max_gen_tokens_per_turn/p95": pct(max_gen_tokens_per_turn_values, 95),
1924+
**calculate_single_metric(
19541925
[m["total_reward"] for m in all_sample_metrics],
19551926
batch_size,
19561927
"total_reward",
@@ -1989,7 +1960,7 @@ def _pct(values: Sequence[float | int], p: float) -> float:
19891960
]
19901961
if values:
19911962
per_agent_metrics.update(
1992-
_calculate_single_metric(
1963+
calculate_single_metric(
19931964
values, len(agent_results), f"{agent_name}/{key}"
19941965
)
19951966
)

pyrefly.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,7 @@ project-includes = [
133133
"nemo_rl/evals/answer_parsing.py",
134134
"nemo_rl/experience/__init__.py",
135135
"nemo_rl/experience/interfaces.py",
136+
"nemo_rl/experience/metric_utils.py",
136137
"nemo_rl/experience/rollout_manager.py",
137138
"nemo_rl/experience/rollouts.py",
138139
"nemo_rl/modelopt/__init__.py",

tests/unit/excluded_unit_tests.sh

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -209,10 +209,18 @@ EXCLUDED_UNIT_TESTS=(
209209
--deselect=tests/unit/models/dtensor/test_parallelize.py::test_parallelize_plan_keys
210210

211211
###########################################################################
212-
# EXPERIENCE — rollout tests all require vLLM setup (~67-91s setup each)
213-
###########################################################################
214-
215-
--ignore=unit/experience/test_rollouts.py
212+
# EXPERIENCE — rollout tests need vLLM setup (~67-91s each).
213+
# Keep pure Python tests, plus the 3 matches_original tests that guard correctness.
214+
###########################################################################
215+
216+
--deselect=tests/unit/experience/test_rollouts.py::test_run_multi_step_calculator_vllm_sync
217+
--deselect=tests/unit/experience/test_rollouts.py::test_run_multi_step_calculator_vllm_async
218+
--deselect=tests/unit/experience/test_rollouts.py::test_max_seqlen_respected_sync
219+
--deselect=tests/unit/experience/test_rollouts.py::test_max_seqlen_respected_async
220+
--deselect=tests/unit/experience/test_rollouts.py::test_run_sliding_puzzle_vllm
221+
--deselect=tests/unit/experience/test_rollouts.py::test_async_rollout_manager
222+
--deselect=tests/unit/experience/test_rollouts.py::test_async_rollout_manager_truncation
223+
--deselect=tests/unit/experience/test_rollouts.py::test_async_nemo_gym_rollout_manager
216224

217225
###########################################################################
218226
# ENVIRONMENTS

0 commit comments

Comments
 (0)