Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions areal/api/cli_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -2262,6 +2262,14 @@ class InferenceEngineConfig:
default=1,
metadata={"help": "Batch size for consuming rollouts from the queue."},
)
min_valid_group_size: int = field(
Comment thread
EazyReal marked this conversation as resolved.
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={
Expand Down Expand Up @@ -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
Expand Down
31 changes: 29 additions & 2 deletions areal/infra/remote_inf_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
1 change: 1 addition & 0 deletions docs/en/cli_reference.md
Original file line number Diff line number Diff line change
Expand Up @@ -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. |
Expand Down
10 changes: 9 additions & 1 deletion docs/en/reference/rollout_workflow.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`:
Expand Down
1 change: 1 addition & 0 deletions docs/zh/cli_reference.md
Original file line number Diff line number Diff line change
Expand Up @@ -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. |
Expand Down
9 changes: 8 additions & 1 deletion docs/zh/reference/rollout_workflow.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` 中:
Expand Down
80 changes: 80 additions & 0 deletions tests/test_grouped_rollout_min_valid.py
Original file line number Diff line number Diff line change
@@ -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
Loading