From 8777cb60520aaa0a7756113b52046c7a0b7c89c6 Mon Sep 17 00:00:00 2001 From: EazyReal <8047065+EazyReal@users.noreply.github.com> Date: Tue, 21 Jul 2026 05:47:34 +0000 Subject: [PATCH] feat(rollout): add minimum valid group size Allow grouped rollouts to reject under-filled groups while keeping partial groups as the default. Document the configuration boundary and cover wrapper plumbing. Signed-off-by: EazyReal <8047065+EazyReal@users.noreply.github.com> --- areal/api/cli_args.py | 12 ++++ areal/infra/remote_inf_engine.py | 31 +++++++++- docs/en/cli_reference.md | 1 + docs/en/reference/rollout_workflow.md | 10 +++- docs/zh/cli_reference.md | 1 + docs/zh/reference/rollout_workflow.md | 9 ++- tests/test_grouped_rollout_min_valid.py | 80 +++++++++++++++++++++++++ 7 files changed, 140 insertions(+), 4 deletions(-) create mode 100644 tests/test_grouped_rollout_min_valid.py diff --git a/areal/api/cli_args.py b/areal/api/cli_args.py index 50d05aa781..32cbbd5a9b 100644 --- a/areal/api/cli_args.py +++ b/areal/api/cli_args.py @@ -2262,6 +2262,14 @@ class InferenceEngineConfig: default=1, metadata={"help": "Batch size for consuming rollouts from the queue."}, ) + min_valid_group_size: int = field( + default=1, + metadata={ + "help": "Minimum valid trajectories required to keep a grouped rollout. " + "Default 1 keeps partial groups; set to gconfig.n_samples to require " + "full groups. Applies when group_size > 1 and must not exceed it." + }, + ) max_head_offpolicyness: int = field( default=0, metadata={ @@ -2420,6 +2428,10 @@ def __post_init__(self): ) if not self.admin_api_key or not self.admin_api_key.strip(): raise ValueError("admin_api_key must not be empty or whitespace-only") + if self.min_valid_group_size < 1: + raise ValueError( + f"min_valid_group_size must be >= 1, got {self.min_valid_group_size}" + ) if ( self._version == "v2" and self.agent is not None diff --git a/areal/infra/remote_inf_engine.py b/areal/infra/remote_inf_engine.py index fb91af095b..50e2d3c129 100644 --- a/areal/infra/remote_inf_engine.py +++ b/areal/infra/remote_inf_engine.py @@ -74,12 +74,19 @@ def __init__( workflow: RolloutWorkflow, group_size: int, logger: Logger, + min_valid_group_size: int = 1, ): if group_size < 1: raise ValueError(f"group_size must be >= 1, got {group_size}") + if not 1 <= min_valid_group_size <= group_size: + raise ValueError( + f"min_valid_group_size must be in [1, group_size={group_size}], " + f"got {min_valid_group_size}" + ) self.workflow = workflow self.group_size = group_size self.logger = logger + self.min_valid_group_size = min_valid_group_size async def arun_episode( self, engine: InferenceEngine, data: dict[str, Any] @@ -96,6 +103,16 @@ async def arun_episode( if not valid_results: return None + if len(valid_results) < self.min_valid_group_size: + self.logger.warning( + "GroupedRolloutWorkflow: only %s/%s trajectories valid " + "(min_valid_group_size=%s), dropping group", + len(valid_results), + len(results), + self.min_valid_group_size, + ) + return None + # Some results None -> warn and continue with valid ones if len(valid_results) < len(results): self.logger.warning( @@ -621,7 +638,12 @@ def _resolve_workflow( raise ValueError("proxy_addr is required for online mode") resolved = self._wrap_openai_agent(None, proxy_addr=proxy_addr) if group_size > 1: - resolved = GroupedRolloutWorkflow(resolved, group_size, self.logger) + resolved = GroupedRolloutWorkflow( + resolved, + group_size, + self.logger, + min_valid_group_size=self.config.min_valid_group_size, + ) return resolved # 1. Already a RolloutWorkflow instance @@ -712,7 +734,12 @@ def _resolve_workflow( # Wrap with GroupedRolloutWorkflow if group_size > 1 if group_size > 1: - resolved = GroupedRolloutWorkflow(resolved, group_size, self.logger) + resolved = GroupedRolloutWorkflow( + resolved, + group_size, + self.logger, + min_valid_group_size=self.config.min_valid_group_size, + ) return resolved diff --git a/docs/en/cli_reference.md b/docs/en/cli_reference.md index 5e4b9ca5d6..004ee8029f 100644 --- a/docs/en/cli_reference.md +++ b/docs/en/cli_reference.md @@ -547,6 +547,7 @@ Configuration for inference servers, including offpolicyness control. | `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | | `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | | `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | +| `min_valid_group_size` | integer | `1` | Minimum valid trajectories required to keep a grouped rollout. Default 1 keeps partial groups; set to gconfig.n_samples to require full groups. Applies when group_size > 1 and must not exceed it. | | `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | | `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | | `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | diff --git a/docs/en/reference/rollout_workflow.md b/docs/en/reference/rollout_workflow.md index 9b0720d595..38c734523d 100644 --- a/docs/en/reference/rollout_workflow.md +++ b/docs/en/reference/rollout_workflow.md @@ -193,10 +193,18 @@ engine.submit( Or via CLI: ```yaml +gconfig: + n_samples: 4 rollout: - group_size: 4 + min_valid_group_size: 4 ``` +`min_valid_group_size` is the minimum number of non-`None` results required to keep a +group. Its default of `1` keeps a partial group whenever at least one trajectory is +valid. Set it to `gconfig.n_samples` to require complete groups. If fewer results are +valid, `GroupedRolloutWorkflow` returns `None` and drops the whole group. The value must +be in `[1, gconfig.n_samples]` when grouped rollout is enabled. + ### How It Works When `group_size > 1`, the workflow is wrapped in `GroupedRolloutWorkflow`: diff --git a/docs/zh/cli_reference.md b/docs/zh/cli_reference.md index 53a04daa75..313cf5c570 100644 --- a/docs/zh/cli_reference.md +++ b/docs/zh/cli_reference.md @@ -545,6 +545,7 @@ Configuration for inference servers, including offpolicyness control. | `max_concurrent_rollouts` | integer \| None | `None` | Maximum number of concurrent rollouts to the inference engine. Defaults to consumer_batch_size. | | `queue_size` | integer \| None | `None` | Input/Output queue size for async rollout. | | `consumer_batch_size` | integer | `1` | Batch size for consuming rollouts from the queue. | +| `min_valid_group_size` | integer | `1` | Minimum valid trajectories required to keep a grouped rollout. Default 1 keeps partial groups; set to gconfig.n_samples to require full groups. Applies when group_size > 1 and must not exceed it. | | `max_head_offpolicyness` | integer | `0` | Maximum off-policyness for the head. If the current version is more than this many versions behind, the request will not be accepted. | | `enable_rollout_tracing` | boolean | `False` | Whether to output verbose tracing messages for each generation request. | | `check_trajectory_format` | boolean | `False` | Whether to check the format of produced trajectories of a customized workflow. Useful when debugging the workflow in isolation. Should be False during RL training. | diff --git a/docs/zh/reference/rollout_workflow.md b/docs/zh/reference/rollout_workflow.md index da3a791189..0f546165d5 100644 --- a/docs/zh/reference/rollout_workflow.md +++ b/docs/zh/reference/rollout_workflow.md @@ -184,10 +184,17 @@ engine.submit( 或通过 CLI: ```yaml +gconfig: + n_samples: 4 rollout: - group_size: 4 + min_valid_group_size: 4 ``` +`min_valid_group_size` 表示一个分组至少需要多少个非 `None` 结果才会被保留。默认值为 +`1`,即只要分组中至少有一个有效结果,就会保留这个部分分组。将它设为 `gconfig.n_samples` 可以要求只保留完整分组。如果有效结果数低于该阈值, +`GroupedRolloutWorkflow` 会返回 `None`,从而丢弃整个分组。启用分组 Rollout 时, 该值必须在 +`[1, gconfig.n_samples]` 范围内。 + ### 工作原理 当 `group_size > 1` 时,工作流被包装在 `GroupedRolloutWorkflow` 中: diff --git a/tests/test_grouped_rollout_min_valid.py b/tests/test_grouped_rollout_min_valid.py new file mode 100644 index 0000000000..8f2e57a76a --- /dev/null +++ b/tests/test_grouped_rollout_min_valid.py @@ -0,0 +1,80 @@ +import asyncio +import logging + +import pytest +import torch + +from areal.api import RolloutWorkflow +from areal.api.cli_args import InferenceEngineConfig +from areal.infra.remote_inf_engine import GroupedRolloutWorkflow, RemoteInfEngine + + +class _Workflow(RolloutWorkflow): + def __init__(self, none_count: int): + self._none_count = none_count + self._calls = 0 + + async def arun_episode(self, engine, data): + self._calls += 1 + if self._calls <= self._none_count: + return None + return { + "input_ids": torch.tensor([[1, 2, 3]]), + "attention_mask": torch.tensor([[1, 1, 1]]), + } + + +def _run_group(none_count: int, min_valid_group_size: int): + workflow = GroupedRolloutWorkflow( + _Workflow(none_count), + group_size=4, + logger=logging.getLogger("test"), + min_valid_group_size=min_valid_group_size, + ) + return asyncio.run(workflow.arun_episode(engine=None, data={})) + + +@pytest.mark.parametrize( + ("min_valid_group_size", "none_count", "expected_rows"), + [ + (1, 4, None), + (1, 1, 3), + (2, 2, 2), + (2, 3, None), + (4, 1, None), + ], +) +def test_min_valid_group_size_filters_underfilled_groups( + min_valid_group_size, none_count, expected_rows +): + out = _run_group(none_count, min_valid_group_size) + + if expected_rows is None: + assert out is None + else: + assert out is not None + assert out["input_ids"].shape[0] == expected_rows + + +@pytest.mark.parametrize("min_valid_group_size", [0, 5]) +def test_min_valid_group_size_outside_group_bounds_raises(min_valid_group_size): + with pytest.raises(ValueError, match="min_valid_group_size must be in"): + GroupedRolloutWorkflow( + _Workflow(0), + group_size=4, + logger=logging.getLogger("test"), + min_valid_group_size=min_valid_group_size, + ) + + +def test_configured_min_valid_group_size_reaches_grouped_workflow(): + engine = object.__new__(RemoteInfEngine) + engine.config = InferenceEngineConfig(backend="sglang:d1", min_valid_group_size=3) + engine.logger = logging.getLogger("test") + + resolved = engine._resolve_workflow( + _Workflow(none_count=0), workflow_kwargs=None, group_size=4 + ) + + assert isinstance(resolved, GroupedRolloutWorkflow) + assert resolved.min_valid_group_size == 3