refactor(rollout): rename shared filter output (#3095)

This commit is contained in:
Jiajun Li
2026-09-06 22:56:06 -07:00
committed by GitHub
parent d2fc97ce58
commit 51853a165a
7 changed files with 28 additions and 16 deletions
+1 -1
View File
@@ -177,7 +177,7 @@ Per-group filter after scoring (DAPO-style). Stock:
```python
def filter_function(args, samples: list[Sample], **kwargs):
# return DynamicFilterOutput(keep=..., reason=...) or a bool
# return FilterOutput(keep=..., reason=...) or a bool
...
```
+2 -2
View File
@@ -161,8 +161,8 @@ Hook to normalize rewards differently from the default GRPO normalization.
Per-group filter; runs after scoring, before queueing for training.
```python
def filter_function(args, samples: list[Sample], **kwargs) -> DynamicFilterOutput:
return DynamicFilterOutput(keep=True, reason=None)
def filter_function(args, samples: list[Sample], **kwargs) -> FilterOutput:
return FilterOutput(keep=True, reason=None)
```
**Stock implementation:** `miles.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std`.
+7 -4
View File
@@ -3,20 +3,23 @@ from dataclasses import dataclass
@dataclass
class DynamicFilterOutput:
class FilterOutput:
keep: bool
reason: str | None = None
DynamicFilterOutput = FilterOutput
def call_dynamic_filter(fn, *args, **kwargs):
if fn is None:
return DynamicFilterOutput(keep=True)
return FilterOutput(keep=True)
output = fn(*args, **kwargs)
# compatibility for legacy version
if not isinstance(output, DynamicFilterOutput):
output = DynamicFilterOutput(keep=output)
if not isinstance(output, FilterOutput):
output = FilterOutput(keep=output)
return output
@@ -1,6 +1,6 @@
import torch
from miles.rollout.filter_hub.base_types import DynamicFilterOutput
from miles.rollout.filter_hub.base_types import FilterOutput
from miles.utils.types import Sample
__all__ = ["check_reward_nonzero_std", "check_no_aborted"]
@@ -18,7 +18,7 @@ def _flatten_samples(samples: list[Sample | list[Sample]]):
def check_reward_nonzero_std(args, samples: list[Sample | list[Sample]], **kwargs):
rewards = [sample.get_reward_value(args) for sample in _flatten_samples(samples)]
keep = torch.tensor(rewards, dtype=torch.float64).std() > 1e-8
return DynamicFilterOutput(
return FilterOutput(
keep=keep,
reason=None if keep else f"zero_std_{round(rewards[0], 1)}",
)
@@ -27,5 +27,5 @@ def check_reward_nonzero_std(args, samples: list[Sample | list[Sample]], **kwarg
def check_no_aborted(args, samples: list[Sample | list[Sample]], **kwargs):
"""Reject entire group if any sample was aborted (e.g. env timeout, Docker crash)."""
if any(s.status == Sample.Status.ABORTED for s in _flatten_samples(samples)):
return DynamicFilterOutput(keep=False, reason="group_has_aborted")
return DynamicFilterOutput(keep=True)
return FilterOutput(keep=False, reason="group_has_aborted")
return FilterOutput(keep=True)
@@ -7,7 +7,7 @@ from miles.rollout.base_types import (
RolloutFnOutput,
RolloutFnTrainInput,
)
from miles.rollout.filter_hub.base_types import DynamicFilterOutput
from miles.rollout.filter_hub.base_types import FilterOutput
from miles.rollout.inference_rollout.compatibility import call_rollout_function, load_rollout_function
from miles.utils.types import Sample, WeightVersionsPerCall
@@ -85,5 +85,5 @@ def load_and_call_train(args, data_source):
def filter_by_reward(args, samples, **kwargs):
reward = samples[0].reward if not isinstance(samples[0], list) else samples[0][0].reward
if reward == 1:
return DynamicFilterOutput(keep=True)
return DynamicFilterOutput(keep=False, reason="reward_zero")
return FilterOutput(keep=True)
return FilterOutput(keep=False, reason="reward_zero")
+9
View File
@@ -0,0 +1,9 @@
from tests.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=20, suite="stage-a-cpu", labels=[])
from miles.rollout.filter_hub.base_types import DynamicFilterOutput, FilterOutput
def test_dynamic_filter_output_is_a_compatibility_alias():
assert DynamicFilterOutput is FilterOutput
@@ -12,7 +12,7 @@ import pytest
import miles.rollout.fully_async_data_buffer as data_buffer
import miles.rollout.fully_async_rollout as fully_async
from miles.rollout.base_types import BaseRolloutFn, RolloutFnConstructorInput, RolloutFnEvalInput, RolloutFnTrainInput
from miles.rollout.filter_hub.base_types import DynamicFilterOutput
from miles.rollout.filter_hub.base_types import FilterOutput
from miles.utils.types import Sample, WeightVersionSpan, WeightVersionsPerCall
N_SAMPLES_PER_PROMPT = 2
@@ -347,7 +347,7 @@ async def test_nested_group_recycles_the_flat_prompt_group(monkeypatch):
def reject_group_1(args, group, **kwargs):
keep = group[0].group_index != 1
return DynamicFilterOutput(keep=keep, reason=None if keep else "rejected")
return FilterOutput(keep=keep, reason=None if keep else "rejected")
async def test_dynamic_filter_drops_group_without_recycling(monkeypatch):