Skip to content
Open
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
4 changes: 3 additions & 1 deletion docs/source/async_distillation_trainer.md
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,8 @@ logprob for at each `beta` regime.

After every `weight_sync_steps` training steps, the updated student weights are transferred to its vLLM server via
NCCL. As with [`experimental.async_grpo.AsyncGRPOTrainer`], generation runs ahead of training, so samples may reflect a slightly stale
policy; `max_staleness` controls how many weight updates a sample can lag behind before being discarded.
policy; `max_staleness` controls how many weight updates a sample can lag behind before in-flight work is cancelled or
a queued sample is discarded.

## Quick start

Expand Down Expand Up @@ -229,6 +230,7 @@ A **rollout** is one prompt taken all the way through: generated by the student,
| `rollout/inflight` | rollouts in flight, i.e. generating or being scored |
| `rollout/vllm_retry_total` | retried vLLM requests, to either the student's server or a teacher's. A degraded server otherwise looks like unexplained slowness. It sits here rather than in `completions/` because it counts requests to a server, not generated text |
| `rollout/backpressure_s` | how long generation was blocked because the rollout queue was full. See [the rollout queue](#the-rollout-queue) |
| `rollout/stale_samples_total` | in-flight samples cancelled because the policy moved more than `max_staleness` versions past the one they started at |

### Samples arriving from the queue

Expand Down
5 changes: 4 additions & 1 deletion docs/source/async_grpo_trainer.md
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ The rollout worker runs in a separate process spawned from the trainer, so rewar

After every `weight_sync_steps` training steps, the updated weights are transferred to the vLLM server via NCCL so that subsequent generations reflect the latest policy.

Because generation and training run concurrently, the training samples may have been generated by a slightly older version of the model. The `max_staleness` parameter controls how many weight updates a sample can lag behind before being discarded.
Because generation and training run concurrently, the training samples may have been generated by a slightly older version of the model. The `max_staleness` parameter controls how many weight updates a sample can lag behind before being discarded. The worker applies the same limit to its in-flight rollouts: when the policy advances, it cancels the generations of any group that started too many versions ago, so vLLM does not finish work the trainer would drop.

The number of concurrent requests sent to the vLLM server is controlled by `max_inflight_tasks`. By default it is set automatically to `max_staleness × per_device_train_batch_size × gradient_accumulation_steps × num_processes` — the maximum number of samples the trainer can consume before they become stale. Generating more than this is wasteful since the excess samples will be discarded.

Expand Down Expand Up @@ -259,6 +259,9 @@ A **rollout** is **one full** conversation: a prompt generated to completion, in
| `rollout/score_s`, `rollout/score_wait_s`, `rollout/score_block_s` | scoring: time to score a group, group wait time to be scored, and how long generation was blocked because the scoring queue was full |
| `rollout/vllm_retry_total` | retried vLLM requests. A degraded server otherwise looks like unexplained slowness. It sits here rather than in `completions/` because it counts requests to the server, not generated text: a retried request produced no completion at all |
| `rollout/backpressure_s` | how long generation was blocked because the rollout queue was full. See [the rollout queue](#the-rollout-queue) |
| `rollout/failed_total` | rollouts that raised (a request that failed every retry, a completion that could not be parsed). The rollout is dropped and its group scored with the rest; only a group where every rollout failed takes the worker down |
| `rollout/dropped_groups_total` | groups dropped because a single rollout survived: a group-relative advantage needs at least two |
| `rollout/stale_groups_total` | in-flight groups cancelled because the policy moved more than `max_staleness` versions past the one they started at. The trainer would drop their samples anyway (see `sample/dropped_stale_total`), so finishing them only burns vLLM compute |

### Tools

Expand Down
65 changes: 54 additions & 11 deletions tests/experimental/test_async_distillation_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@

import pytest
import torch
from accelerate import PartialState
from datasets import Dataset, load_dataset
from transformers import AutoTokenizer
from transformers.testing_utils import torch_device
Expand Down Expand Up @@ -356,6 +357,24 @@ def _bare_loop(tokenizer, teacher_server_urls):
TWO_TEACHERS = {"math": "http://math:8002", "code": "http://code:8003"}


def _rollout_loop(dataset, **kwargs):
PartialState()
ctx = mp.get_context("spawn")
loop_kwargs = dict(
model_name="test",
dataset=dataset,
processing_class=MagicMock(),
rollout_buffer=ctx.Queue(),
metrics_queue=ctx.Queue(),
model_version_value=ctx.Value("i", 0),
heartbeat_value=ctx.Value("d", 0.0),
failed_event=ctx.Event(),
exception_info_queue=ctx.Queue(),
)
loop_kwargs.update(kwargs)
return _AsyncRolloutLoop(**loop_kwargs)


class TestWorkerMetrics:
"""The worker's payload has the same shape as the trainer's sink, so draining it in `log()` is an append."""

Expand Down Expand Up @@ -922,6 +941,39 @@ def test_stops_once_the_prompt_target_is_reached(self, trained, before_resume, s
assert control.should_training_stop is should_stop


class TestGenerateLoop(TrlTestCase):
def test_stale_in_flight_samples_are_cancelled_when_the_policy_advances(self):
loop = _rollout_loop(
Dataset.from_dict({"prompt": [f"q{i}" for i in range(8)]}), max_inflight_tasks=2, max_staleness=0
)
cancelled = []

async def run():
stop = asyncio.Event()

async def generate_and_score_one(prompt_id, row):
version = loop.model_version
if version == 0:
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
cancelled.append(prompt_id)
raise
stop.set()
return types.SimpleNamespace(completion_mask=[1], model_version=version, enqueued_at=None)

loop._generate_and_score_one = generate_and_score_one
task = asyncio.create_task(loop._generate_loop(stop))
await asyncio.sleep(0.1)
loop._model_version_value.value = 1
await asyncio.wait_for(task, 5)

asyncio.run(run())
assert len(cancelled) == 2
assert loop.rollout_buffer.get(timeout=5).model_version == 1
assert loop._metrics_queue.get(timeout=5)["rollout/stale_samples_total"] == 2


class TestRolloutStateCheckpoint(TrlTestCase):
"""Prompt-index checkpoint/resume logic — no GPU or vLLM required."""

Expand Down Expand Up @@ -974,17 +1026,8 @@ def record(*_args, **_kwargs):
assert written_before_super == [True]

def test_rollout_loop_skips_to_start_index(self):
ctx = mp.get_context("spawn")
loop = _AsyncRolloutLoop(
model_name="test",
dataset=Dataset.from_dict({"prompt": [f"row_{i}" for i in range(10)]}),
processing_class=MagicMock(),
rollout_buffer=ctx.Queue(),
model_version_value=ctx.Value("i", 0),
heartbeat_value=ctx.Value("d", 0.0),
failed_event=ctx.Event(),
exception_info_queue=ctx.Queue(),
metrics_queue=ctx.Queue(),
loop = _rollout_loop(
Dataset.from_dict({"prompt": [f"row_{i}" for i in range(10)]}),
dataset_start_index=3,
)
_prompt_id, row = next(loop._repeat_iterator())
Expand Down
139 changes: 138 additions & 1 deletion tests/experimental/test_async_grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -645,7 +645,7 @@ def test_rollout_loop_skips_to_start_index(self):
dataset = Dataset.from_dict({"prompt": [f"row_{i}" for i in range(10)]})
loop = self._make_rollout_loop(dataset, dataset_start_index=3)
it = loop._repeat_iterator()
_group_id, row = next(it)
_group_id, _index, row = next(it)
assert row["prompt"] == "row_3"

def test_inner_training_loop_sets_dataset_start_index_from_file(self):
Expand Down Expand Up @@ -1373,6 +1373,143 @@ def test_epoch_stop_is_fork_independent(self):
assert forked.state.global_step > no_fork.state.global_step


def _strict_reward(completions, answer, **kwargs):
return [1.0 for _ in zip(completions, answer, strict=True)]


_ROLLOUT = (
[{"role": "assistant", "content": "c"}],
[1, 2],
[TrainingSequence([1, 2], [0, 1], [0.0, -0.1], "r")],
0,
0,
None,
)


class TestGenerateLoop(TrlTestCase):
def _loop(self, num_generations, max_inflight_tasks, max_staleness=4):
PartialState()
with patch("trl.experimental.async_grpo.async_rollout_worker.add_response_schema", side_effect=lambda x: x):
return _AsyncRolloutLoop(
model_name="test",
dataset=Dataset.from_dict({"prompt": [f"q{i}" for i in range(8)]}),
reward_funcs=[dummy_reward_func],
processing_class=MagicMock(),
rollout_buffer=queue.Queue(),
metrics_queue=queue.Queue(),
model_version_value=mp.Value("i", 0),
heartbeat_value=mp.Value("d", 0.0),
failed_event=mp.Event(),
exception_info_queue=queue.Queue(),
num_generations=num_generations,
max_inflight_tasks=max_inflight_tasks,
max_staleness=max_staleness,
)

async def _groups(self, loop, generate_one, n, bump_version_to=None):
"""Run the generate loop until `n` groups reach the score queue."""
loop._generate_one = generate_one
stop = asyncio.Event()
task = asyncio.create_task(loop._generate_loop(stop))
if bump_version_to is not None:
await asyncio.sleep(0.1)
loop._model_version_value.value = bump_version_to
groups = [await asyncio.wait_for(loop._groups_to_score.get(), 5) for _ in range(n)]
stop.set()
await task
return groups

def test_failed_rollout_is_dropped_and_the_group_scored_with_the_rest(self):
loop = self._loop(num_generations=4, max_inflight_tasks=4)
calls = itertools.count()

async def generate_one(prompt, tool_dict, tools, group_id):
if next(calls) == 1:
raise RuntimeError("boom")
return _ROLLOUT

(group,) = asyncio.run(self._groups(loop, generate_one, 1))
assert group.group_id == 0
assert len(group.completions) == 3

def test_group_with_a_single_surviving_rollout_is_dropped(self):
loop = self._loop(num_generations=3, max_inflight_tasks=3)
calls = itertools.count()

async def generate_one(prompt, tool_dict, tools, group_id):
if next(calls) < 2:
raise RuntimeError("boom")
return _ROLLOUT

(group,) = asyncio.run(self._groups(loop, generate_one, 1))
assert group.group_id == 1

def test_group_where_every_rollout_fails_raises(self):
loop = self._loop(num_generations=2, max_inflight_tasks=2)

async def generate_one(prompt, tool_dict, tools, group_id):
raise RuntimeError("boom")

loop._generate_one = generate_one
with pytest.raises(RuntimeError, match="boom"):
asyncio.run(asyncio.wait_for(loop._generate_loop(asyncio.Event()), 5))

def test_stale_in_flight_groups_are_cancelled_when_the_policy_advances(self):
loop = self._loop(num_generations=2, max_inflight_tasks=4, max_staleness=1)
cancelled = []

async def generate_one(prompt, tool_dict, tools, group_id):
if loop.model_version == 0:
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
await asyncio.sleep(0)
cancelled.append(group_id)
raise
assert len(cancelled) == 4
return _ROLLOUT

groups = asyncio.run(self._groups(loop, generate_one, 2, bump_version_to=2))
assert len(cancelled) == 4
assert [g.group_id for g in groups] == [2, 3]
assert [g.model_version for g in groups] == [2, 2]

def test_partially_dispatched_stale_group_is_regenerated_as_a_smaller_group(self):
loop = self._loop(num_generations=4, max_inflight_tasks=2, max_staleness=0)

async def generate_one(prompt, tool_dict, tools, group_id):
if loop.model_version == 0:
await asyncio.Event().wait()
return _ROLLOUT

(group,) = asyncio.run(self._groups(loop, generate_one, 1, bump_version_to=1))
assert group.group_id == 0
assert len(group.completions) == 2
assert group.model_version == 1

def test_stale_group_with_one_undispatched_rollout_is_skipped(self):
loop = self._loop(num_generations=4, max_inflight_tasks=3, max_staleness=0)

async def generate_one(prompt, tool_dict, tools, group_id):
if loop.model_version == 0:
await asyncio.Event().wait()
if group_id == 0:
raise RuntimeError("the stale tail must not be dispatched")
return _ROLLOUT

(group,) = asyncio.run(self._groups(loop, generate_one, 1, bump_version_to=1))
assert group.group_id == 1
assert group.model_version == 1

def test_reward_kwargs_are_trimmed_to_the_surviving_rollouts(self):
PartialState()
group = _group([[_ROLLOUT[2][0]]] * 3, [[1, 2]] * 3)
group.reward_kwargs = {"answer": ["a"] * 4}
samples = asyncio.run(_bare_loop([_strict_reward])._score_group(group))
assert len(samples) == 3


@require_peft
class TestValidateLoraForVLLMSync(TrlTestCase):
model_id = "trl-internal-testing/tiny-Qwen2ForCausalLM-2.5"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -149,8 +149,8 @@ class AsyncDistillationConfig(_BaseConfig):
`-1` (auto), which sets it to `max_staleness * per_device_train_batch_size * gradient_accumulation_steps *
num_processes`.
max_staleness (`int`, *optional*, defaults to `4`):
Maximum number of weight update steps a rollout sample can lag behind the current model version before
being discarded.
Maximum number of weight update steps a rollout sample can lag behind the current model version before an
in-flight sample is cancelled or a queued sample is discarded.
queue_maxsize (`int`, *optional*, defaults to `1024`):
Maximum number of rollout samples to buffer in the rollout queue.
weight_sync_steps (`int`, *optional*, defaults to `1`):
Expand Down Expand Up @@ -358,7 +358,7 @@ class AsyncDistillationConfig(_BaseConfig):
default=4,
metadata={
"help": "Maximum number of weight update steps a rollout sample can lag behind the current model "
"version before being discarded."
"version before an in-flight sample is cancelled or a queued sample is discarded."
},
)
queue_maxsize: int = field(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1133,6 +1133,7 @@ def __init__(
teacher_top_k=self.args.teacher_top_k,
teacher_temperature=self.args.teacher_temperature,
max_tokens=self.args.max_completion_length,
max_staleness=self.args.max_staleness,
temperature=self.args.temperature,
top_p=self.args.top_p,
top_k=self.args.top_k,
Expand Down
30 changes: 27 additions & 3 deletions trl/experimental/async_distillation/async_rollout_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,7 @@ def __init__(
teacher_top_k: int = 8,
teacher_temperature: float = 1.0,
max_tokens: int = 32,
max_staleness: int = 4,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 0,
Expand Down Expand Up @@ -243,6 +244,7 @@ def __init__(
self.max_inflight_tasks = max_inflight_tasks
self.queue_maxsize = queue_maxsize
self.max_tokens = max_tokens
self.max_staleness = max_staleness
self.temperature = temperature
self.top_p = top_p
self.top_k = top_k
Expand Down Expand Up @@ -321,19 +323,41 @@ async def _resolve_teacher_model_names(self) -> None:
logger.info(f"teacher {teacher_id!r} at {url} serves {self.teacher_model_names[teacher_id]}")

async def _generate_loop(self, stop_event: asyncio.Event) -> None:
inflight_tasks: dict[asyncio.Task, int] = {}
# Keep the dispatch version beside the slot: a sample does not expose its model version until its task returns.
inflight_tasks: dict[asyncio.Task, tuple[int, int]] = {}
free_slots = set(range(self.max_inflight_tasks))
work_iter = self._repeat_iterator()
last_version = self.model_version

self._generation_start_time = time.monotonic()
try:
while True:
self._heartbeat_value.value = time.time()

version = self.model_version
if version != last_version:
last_version = version
stale_tasks = [
task
for task, (_slot, task_version) in inflight_tasks.items()
if version - task_version > self.max_staleness
]
for task in stale_tasks:
task.cancel()
if stale_tasks:
await asyncio.gather(*stale_tasks, return_exceptions=True)
for task in stale_tasks:
slot, _task_version = inflight_tasks.pop(task)
free_slots.add(slot)
if stale_tasks:
self._counters["rollout/stale_samples_total"] += len(stale_tasks)
logger.info(f"cancelled {len(stale_tasks)} stale rollout(s) at version {version}")

while free_slots and not stop_event.is_set():
prompt_id, row = next(work_iter)
slot = free_slots.pop()
task = asyncio.create_task(self._generate_and_score_one(prompt_id, row))
inflight_tasks[task] = slot
inflight_tasks[task] = (slot, self.model_version)

if not inflight_tasks:
if stop_event.is_set():
Expand All @@ -346,7 +370,7 @@ async def _generate_loop(self, stop_event: asyncio.Event) -> None:
continue

for task in done:
slot = inflight_tasks.pop(task)
slot, _task_version = inflight_tasks.pop(task)
free_slots.add(slot)
if task.exception() is not None:
raise task.exception()
Expand Down
Loading
Loading