mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Report the raw reward of every generated group, filter or no filter (#3073)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user