Report the raw reward of every generated group, filter or no filter (#3073)

This commit is contained in:
fzyzcjy
2026-09-26 20:44:18 +08:00
committed by GitHub
parent 6419b7258a
commit fbb1450cab
11 changed files with 291 additions and 25 deletions
@@ -809,6 +809,7 @@ class VerifiersRolloutFn(BaseRolloutFn):
continue
await self._apply_miles_rewards(group)
all_groups.append(group)
metrics.on_group_before_dynamic_filter(self.args, _flatten_samples(group))
filter_output = apply_preput_filters(
self.args,
self.dynamic_filter,
+26 -1
View File
@@ -1,3 +1,4 @@
import argparse
from collections import defaultdict
from collections.abc import Iterator
from dataclasses import dataclass
@@ -38,6 +39,19 @@ def call_dynamic_filter(fn, args, samples: list[Sample | list[Sample]], **kwargs
class MetricGatherer:
def __init__(self):
self._dynamic_filter_drop_reason_count = defaultdict(lambda: 0)
self._unfiltered_reward_sum = 0.0
self._unfiltered_reward_count = 0
def on_group_before_dynamic_filter(self, args: argparse.Namespace, group: list) -> None:
for sample in _iter_group_samples(group):
if sample.reward is None:
continue
if not args.reward_key and isinstance(sample.reward, dict):
continue
if (value := sample.get_reward_value(args)) is None:
continue
self._unfiltered_reward_sum += float(value)
self._unfiltered_reward_count += 1
def on_dynamic_filter_drop(self, reason: str | None):
if not reason:
@@ -45,7 +59,18 @@ class MetricGatherer:
self._dynamic_filter_drop_reason_count[reason] += 1
def collect(self):
return {
metrics = {
f"rollout/dynamic_filter/drop_{reason}": count
for reason, count in self._dynamic_filter_drop_reason_count.items()
}
if self._unfiltered_reward_count:
metrics["rollout/raw_reward_unfiltered"] = self._unfiltered_reward_sum / self._unfiltered_reward_count
return metrics
def _iter_group_samples(group: list) -> Iterator[Sample]:
for item in group:
if isinstance(item, list):
yield from item
else:
yield item
+1
View File
@@ -174,6 +174,7 @@ class DefaultDataBuffer(DataBuffer):
self._metric_gatherer.on_dynamic_filter_drop(reason=output.reason)
return False
self._metric_gatherer.on_group_before_dynamic_filter(self._args, input.group)
output = call_dynamic_filter(self._dynamic_filter, self._args, input.group)
if not output.keep:
self._metric_gatherer.on_dynamic_filter_drop(reason=output.reason)
@@ -144,6 +144,7 @@ async def generate_rollout_async(
assert len(group) == args.n_samples_per_prompt
all_data.append(group)
metric_gatherer.on_group_before_dynamic_filter(args, group)
filter_output = apply_preput_filters(args, dynamic_filter, group)
if not filter_output.keep:
metric_gatherer.on_dynamic_filter_drop(reason=filter_output.reason)
+1
View File
@@ -500,6 +500,7 @@ async def generate_rollout_async(
assert len(group) == args.n_samples_per_prompt
all_data.append(group)
metric_gatherer.on_group_before_dynamic_filter(args, group)
filter_output = apply_preput_filters(args, dynamic_filter, group)
if not filter_output.keep:
metric_gatherer.on_dynamic_filter_drop(reason=filter_output.reason)
+3 -4
View File
@@ -137,12 +137,11 @@ def rollout_env(tmp_path, request) -> RolloutEnv:
data_path = str(tmp_path / "data.jsonl")
_write_jsonl(data_path, data_rows)
router_port = find_available_port(20000)
args = _build_args(data_path=data_path, router_port=router_port, extra_argv=config.extra_argv)
SingletonMeta.clear_all_instances()
with with_mock_server(model_name=args.hf_checkpoint, latency=config.latency) as mock_server:
with with_mock_server(latency=config.latency) as mock_server:
router_port = find_available_port(20000)
args = _build_args(data_path=data_path, router_port=router_port, extra_argv=config.extra_argv)
with _with_miles_router(args) as router_server:
r = requests.post(
f"{router_server.url}/add_worker",
@@ -0,0 +1,70 @@
from argparse import Namespace
from miles.rollout.filter_hub.base_types import MetricGatherer
from miles.utils.types import Sample
class TestMetricGathererUnfilteredRawReward:
def test_the_mean_covers_kept_and_dropped_groups_alike(self):
"""The gatherer sees every group before the filter, so the mean must not depend on the keep decision."""
gatherer = MetricGatherer()
gatherer.on_group_before_dynamic_filter(_args(), [_sample(reward=1.0), _sample(reward=1.0)])
gatherer.on_group_before_dynamic_filter(_args(), [_sample(reward=0.0), _sample(reward=0.0)])
gatherer.on_dynamic_filter_drop(reason="zero_std_0")
assert gatherer.collect()["rollout/raw_reward_unfiltered"] == 0.5
def test_a_nested_group_contributes_every_inner_sample(self):
"""Multi-turn groups arrive as lists of lists; counting outer items would weight turns unevenly."""
gatherer = MetricGatherer()
gatherer.on_group_before_dynamic_filter(
_args(), [[_sample(reward=1.0)], [_sample(reward=0.0), _sample(reward=0.5)]]
)
assert gatherer.collect()["rollout/raw_reward_unfiltered"] == 0.5
def test_a_dict_reward_is_read_through_the_reward_key(self):
"""Custom generate functions store dict rewards; the metric must follow --reward-key like raw_reward does."""
gatherer = MetricGatherer()
gatherer.on_group_before_dynamic_filter(
_args(reward_key="reward_value"), [_sample(reward={"reward_value": 0.25, "outcome": "x"})]
)
assert gatherer.collect()["rollout/raw_reward_unfiltered"] == 0.25
def test_an_unkeyed_structured_reward_does_not_report_a_raw_reward(self):
"""Structured custom-RM payloads have no unambiguous scalar value to average."""
gatherer = MetricGatherer()
gatherer.on_group_before_dynamic_filter(
_args(), [_sample(reward={"teacher": {"reward": 0.25}, "student": {"reward": 0.5}})]
)
assert "rollout/raw_reward_unfiltered" not in gatherer.collect()
def test_no_offered_group_reports_no_metric(self):
"""A window without generation must yield no point rather than a fabricated zero."""
gatherer = MetricGatherer()
assert "rollout/raw_reward_unfiltered" not in gatherer.collect()
def _args(reward_key: str | None = None) -> Namespace:
return Namespace(reward_key=reward_key)
def _sample(reward: float | dict | None) -> Sample:
return Sample(group_index=0, index=0, prompt="p", label="l", reward=reward)
class TestMetricGathererUnscoredSamples:
def test_an_unscored_sample_is_left_out_of_the_mean_instead_of_crashing_it(self):
"""A late-aborted group-RM group arrives COMPLETED with reward None; it carries no score to average."""
gatherer = MetricGatherer()
gatherer.on_group_before_dynamic_filter(_args(), [_sample(reward=1.0), _sample(reward=None)])
assert gatherer.collect()["rollout/raw_reward_unfiltered"] == 1.0
@@ -44,3 +44,27 @@ def test_filter_effect(rollout_env, use_filter, expect_all_correct):
assert rewards == {1}, "Filter should keep only correct samples"
else:
assert 0 in rewards, "Without filter, incorrect samples should be present"
@pytest.mark.parametrize(
"rollout_env",
[
pytest.param(
integration_env_config(
["--rollout-batch-size", "3", "--dynamic-sampling-filter-path", "test:filter_by_reward"],
data_rows=MIXED_DATA_ROWS,
),
id="with_filter",
)
],
indirect=["rollout_env"],
)
def test_unfiltered_raw_reward_still_sees_the_dropped_groups(rollout_env):
"""The accepted-only raw_reward is pinned by the filter; this mean must count what was filtered away."""
env = rollout_env
with function_registry.temporary("test:filter_by_reward", filter_by_reward):
out = load_and_call_train(env.args, env.data_source)
assert {group[0].reward for group in out.samples} == {1}
assert 0 < out.metrics["rollout/raw_reward_unfiltered"] < 1
@@ -0,0 +1,126 @@
from argparse import Namespace
from types import SimpleNamespace
import pytest
from miles.rollout import fully_async_data_buffer
from miles.rollout.fully_async_data_buffer import (
DataBuffer,
DataBufferConstructorInput,
DataBufferInput,
DefaultMultiDataBuffer,
)
class _RecordingBuffer(DataBuffer):
"""A custom DataBuffer of the kind --custom-async-data-buffer-path-per-model names."""
def __init__(self, input: DataBufferConstructorInput) -> None:
self.input = input
self.metrics_asked_about: list[str | None] = []
async def put(self, input: DataBufferInput) -> None:
raise NotImplementedError
async def get(self, **context) -> DataBufferInput:
raise NotImplementedError
def get_metrics(self, trainer_model_id: str | None = None) -> dict[str, float]:
self.metrics_asked_about.append(trainer_model_id)
return {"asked": float(len(self.metrics_asked_about))}
def _multi_buffer(monkeypatch: pytest.MonkeyPatch, *, model_ids: list[str]) -> DefaultMultiDataBuffer:
monkeypatch.setattr(
fully_async_data_buffer, "resolve_megatron_config", lambda args: SimpleNamespace(model_ids=model_ids)
)
monkeypatch.setattr(fully_async_data_buffer, "load_function", lambda path: _RecordingBuffer)
args = Namespace(custom_async_data_buffer_path_per_model=[f"{one}=recording.Buffer" for one in model_ids])
return DefaultMultiDataBuffer(DataBufferConstructorInput(args=args, unused_handler_fn=lambda samples: None))
def _composed(multi: DefaultMultiDataBuffer, model_id: str) -> _RecordingBuffer:
return multi._inners[model_id]
class TestTheMetricsOfOnePolicy:
def test_the_policy_the_drain_asked_about_reaches_the_buffer_it_composes(self, monkeypatch):
"""A buffer that selects by policy saw None and could attribute its metrics to the wrong one."""
multi = _multi_buffer(monkeypatch, model_ids=["solver", "verifier"])
multi.get_metrics("solver")
assert _composed(multi, "solver").metrics_asked_about == ["solver"]
def test_a_policy_is_never_asked_about_another_one(self, monkeypatch):
"""Each policy keeps its own window counters, and a drain resets the ones it reads."""
multi = _multi_buffer(monkeypatch, model_ids=["solver", "verifier"])
multi.get_metrics("solver")
multi.get_metrics("verifier")
assert _composed(multi, "solver").metrics_asked_about == ["solver"]
assert _composed(multi, "verifier").metrics_asked_about == ["verifier"]
def test_the_metrics_returned_are_the_ones_that_policy_buffer_reported(self, monkeypatch):
"""Forwarding the policy may not cost the caller the numbers it came for."""
multi = _multi_buffer(monkeypatch, model_ids=["solver", "verifier"])
assert multi.get_metrics("solver") == {"asked": 1.0}
assert multi.get_metrics("solver") == {"asked": 2.0}
def test_a_policy_this_run_does_not_train_is_refused(self, monkeypatch):
"""The composed buffers are one per policy of the run, so any other name selects nothing."""
multi = _multi_buffer(monkeypatch, model_ids=["solver"])
with pytest.raises(AssertionError, match="trains no policy of this run"):
multi.get_metrics("stranger")
from tests.fast.fixtures.megatron_config_fixtures import encode_megatron_config
from miles.utils.types import Sample
def _make_args() -> Namespace:
return Namespace(
async_data_buffer_capacity_factor=1.0,
custom_async_data_buffer_path_per_model=None,
dynamic_sampling_filter_path=None,
max_weight_staleness=None,
megatron_config=encode_megatron_config("solver", "verifier"),
reward_key=None,
rollout_batch_size=1,
)
def _make_sample(*, index: int, reward: float, trainer_model_id: str) -> Sample:
sample = Sample(
index=index,
prompt="prompt",
response="response",
reward=reward,
status=Sample.Status.COMPLETED,
)
sample.trainer_model_id = trainer_model_id
return sample
def _ignore_group(group: list[Sample]) -> None:
pass
class TestPerPolicyMetrics:
async def test_collecting_one_policy_metrics_does_not_reset_another_policy_window(self) -> None:
"""Collecting one policy preserves another policy's resettable raw-reward window."""
buffer = DefaultMultiDataBuffer(DataBufferConstructorInput(args=_make_args(), unused_handler_fn=_ignore_group))
group = [
_make_sample(index=1, reward=1.0, trainer_model_id="solver"),
_make_sample(index=2, reward=3.0, trainer_model_id="verifier"),
]
await buffer.put(DataBufferInput(prompt_group=group, group=group))
assert buffer.get_metrics("solver")["rollout/raw_reward_unfiltered"] == 1.0
assert buffer.get_metrics("verifier")["rollout/raw_reward_unfiltered"] == 3.0
assert "rollout/raw_reward_unfiltered" not in buffer.get_metrics("verifier")
+38 -20
View File
@@ -54,6 +54,7 @@ def make_group(
group_index: int,
status: Sample.Status = Sample.Status.COMPLETED,
weight_versions: list[str] | None = None,
reward: float = 1,
) -> list[Sample]:
versions = [
WeightVersionsPerCall(spans=[WeightVersionSpan(version=version, abs_start=0, abs_end=1)])
@@ -67,7 +68,7 @@ def make_group(
response="ok",
response_length=1,
label="ok",
reward=1,
reward=reward,
status=status,
weight_versions=list(versions),
)
@@ -93,6 +94,7 @@ def make_args(**overrides) -> Namespace:
rollout_sample_filter_path=None,
sglang_router_ip="127.0.0.1",
sglang_router_port=30000,
sglang_router_request_timeout_secs=14400,
eval_num_gpus=0,
)
defaults.update(overrides)
@@ -384,25 +386,6 @@ async def test_worker_cancellation_still_propagates(monkeypatch: pytest.MonkeyPa
await asyncio.sleep(0)
async def test_worker_bounds_in_flight_groups(monkeypatch):
release = asyncio.Event()
async def blocking_generate(state, group, sampling_params, evaluation=False, sample_done_callback=None):
await release.wait()
return group
data_source = FakeDataSource()
fn = make_fn(monkeypatch, make_args(rollout_batch_size=2), data_source, generate=blocking_generate)
drain = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=0)))
await asyncio.sleep(0.05)
assert data_source.num_get_calls == 2 # in-flight bound, not more
release.set()
output = await drain
assert len(output.samples) == 2
async def test_async_max_concurrent_samples_caps_in_flight_groups(monkeypatch):
release = asyncio.Event()
@@ -542,6 +525,22 @@ async def put_group(buffer, group):
await buffer.put(data_buffer.DataBufferInput(prompt_group=group, group=group))
async def test_buffer_reports_unfiltered_raw_reward_across_kept_and_dropped():
"""The accepted-only raw_reward is conditioned by the filter, so this mean must still see dropped groups."""
args = make_args(rollout_batch_size=1, dynamic_sampling_filter_path=f"{__name__}.reject_group_1")
buffer = data_buffer.DefaultDataBuffer(
data_buffer.DataBufferConstructorInput(args=args, unused_handler_fn=lambda group: None)
)
await put_group(buffer, make_group(1, reward=0))
await put_group(buffer, make_group(2, reward=1))
metrics = buffer.get_metrics()
assert metrics["rollout/raw_reward_unfiltered"] == 0.5
assert metrics["rollout/dynamic_filter/drop_rejected"] == 1
assert "rollout/raw_reward_unfiltered" not in buffer.get_metrics()
async def test_buffer_blocks_producer_when_full():
buffer, _ = make_buffer(max_groups=2)
await put_group(buffer, make_group(1))
@@ -1158,3 +1157,22 @@ class TestRolloutFnContract:
fn = make_fn(monkeypatch, make_args(rollout_batch_size=1), data_source)
assert fn.constructor_input.data_source is data_source
async def test_worker_bounds_in_flight_groups(monkeypatch):
release = asyncio.Event()
async def blocking_generate(state, group, sampling_params, evaluation=False, sample_done_callback=None):
await release.wait()
return group
data_source = FakeDataSource()
fn = make_fn(monkeypatch, make_args(rollout_batch_size=2), data_source, generate=blocking_generate)
drain = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=0)))
await asyncio.sleep(0.05)
assert data_source.num_get_calls == 2 # in-flight bound, not more
release.set()
output = await drain
assert len(output.samples) == 2