mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
refactor(rollout): rename shared filter output (#3095)
This commit is contained in:
@@ -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
|
||||
...
|
||||
```
|
||||
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user