Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
25 commits
Select commit Hold shift + click to select a range
bb72789
test: add FSDP2 distributed coverage for AsyncGRPOTrainer
behroozazarkhalili Jun 23, 2026
ffd5945
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Jun 24, 2026
d6d198d
style: reflow worker docstrings to max_len 119 (doc-builder)
behroozazarkhalili Jun 24, 2026
2fc1fe0
test: mark AsyncGRPO FSDP2 smoke as slow so it runs in the multi-GPU …
behroozazarkhalili Jul 3, 2026
b3e2d3e
Merge main into test/async-grpo-fsdp2-coverage; resolve test_async_gr…
behroozazarkhalili Aug 4, 2026
002095d
test: fix the AsyncGRPO FSDP2 worker stub against the current rollout…
behroozazarkhalili Aug 6, 2026
7488de2
Merge main into test/async-grpo-fsdp2-coverage; resolve the test_asyn…
behroozazarkhalili Aug 20, 2026
cfb6e30
test: give the FSDP2 stub worker the metrics_queue the trainer now dr…
behroozazarkhalili Aug 20, 2026
9d68df7
test: fix the two defects that kept the FSDP2 smoke from ever reporting
behroozazarkhalili Aug 20, 2026
be42989
Merge main into test/async-grpo-fsdp2-coverage; resolve the import bl…
behroozazarkhalili Aug 20, 2026
99c9054
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Aug 21, 2026
6c431cd
test: stop the FSDP2 smoke test starving on token-budgeted batching
behroozazarkhalili Aug 21, 2026
2433b1e
test: xfail the FSDP2 smoke on the FA head_size limit (#6837)
behroozazarkhalili Aug 21, 2026
f63584f
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Aug 21, 2026
72bf811
test(async-grpo): use a Flash Attention compatible model in the FSDP2…
behroozazarkhalili Aug 21, 2026
c75c659
test(async-grpo): write FSDP2 worker artifacts to the test's temp dir
behroozazarkhalili Aug 21, 2026
666b976
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Aug 27, 2026
79ab84f
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Sep 3, 2026
2a112af
test(async-grpo): kill the whole launcher group on timeout, require e…
behroozazarkhalili Sep 3, 2026
059e072
test(async-grpo): kill the whole launcher tree on timeout, not the la…
behroozazarkhalili Sep 3, 2026
a4cffd9
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Sep 3, 2026
83dff6b
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Sep 3, 2026
f4c3814
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Sep 3, 2026
0c43a43
test(async_grpo): make the FSDP2 worker report its launch shape
behroozazarkhalili Sep 4, 2026
d03ecc2
Merge remote-tracking branch 'origin/main' into test/async-grpo-fsdp2…
behroozazarkhalili Sep 4, 2026
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
220 changes: 220 additions & 0 deletions tests/experimental/_async_grpo_fsdp2_worker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,220 @@
# Copyright 2020-2026 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Companion worker launched under ``accelerate launch --config_file <fsdp2>`` by the FSDP2 case in
``test_async_grpo_trainer.py``.

It runs a couple of :class:`AsyncGRPOTrainer` steps on an FSDP2-sharded model, driven by an in-process stub rollout
worker (no vLLM server, no NCCL weight transfer), and checks that training actually progresses under FSDP2: the loss is
finite and the parameters change. It then prints one machine-parseable result line (``ASYNC_GRPO_FSDP2_RESULT {json}``)
that the pytest side asserts on.

This is a *functional* FSDP2 smoke, not a #6077 all-gather microbenchmark.

Self-contained on purpose (mirrors ``tests/experimental/_openreward_echo_env.py``): it imports only public TRL symbols
and carries its own stub, so it never imports pytest-internal classes across the subprocess boundary.
"""

from __future__ import annotations

import itertools
import json
import os
import queue

import numpy as np
import torch
from datasets import load_dataset
from transformers import AutoTokenizer

from trl.experimental.async_grpo import AsyncGRPOConfig, AsyncGRPOTrainer
from trl.experimental.async_grpo.async_rollout_worker import RolloutSample


# The trainer loads the model with Flash Attention, which requires a `head_size` multiple of 8. Hence the `small-*`
# model (`head_size=32`) below, rather than the usual `tiny-*` one (`head_size=2`).
MODEL_ID = "trl-internal-testing/small-Qwen2ForCausalLM-2.5"
RESULT_PREFIX = "ASYNC_GRPO_FSDP2_RESULT"
# Reported alongside the measured step count so the launcher can require the whole loop to have run.
_MAX_STEPS = 2


def dummy_reward_func(completions, **kwargs):
# Mirrors tests/experimental/test_async_grpo_trainer.py: the stub pre-computes rewards, so this is
# only here to satisfy the trainer's required `reward_funcs` argument.
return [float(hash(c[0]["content"]) % 100) / 100.0 for c in completions]


class _StubRolloutWorker:
"""Minimal in-process rollout worker — same shape as the one in test_async_grpo_trainer.py.

Reproduced here (rather than imported) because this module runs as ``__main__`` under ``accelerate launch``, not as
a pytest module, so importing the test class would be fragile. Keeping it self-contained matches the openreward
companion-script precedent.
"""

def __init__(self, tokenizer, dataset, num_generations: int = 3, samples_per_weight_sync: int = 10):
self.rollout_buffer = queue.Queue()
self.metrics_queue = queue.Queue() # drained by the trainer in `log()`; this stub measures nothing
self._samples_per_weight_sync = samples_per_weight_sync
self._model_version = 0
self._sample_iter = self._make_sample_iter(tokenizer, dataset, num_generations)

def _make_sample_iter(self, tokenizer, dataset, num_generations):
for group_id, row in enumerate(itertools.cycle(dataset)):
completions = [
[{"role": "assistant", "content": f"{row['completion'][0]['content']} {idx}"}]
for idx in range(num_generations)
]
prompt_completions = [row["prompt"] + completion for completion in completions]
prompt_ids = tokenizer.apply_chat_template(
row["prompt"], tokenize=True, add_generation_prompt=True, return_dict=False
)
prompt_completion_ids = tokenizer.apply_chat_template(
prompt_completions, tokenize=True, add_generation_prompt=False, return_dict=False
)
# Distinct rewards by construction. Hash-derived rewards can collide within a group under a randomized
# PYTHONHASHSEED, which makes `rewards.std()` zero and the advantages NaN, so a run would fail for a reason
# unrelated to FSDP2.
rewards = np.linspace(0.0, 1.0, num_generations)
advantages = (rewards - rewards.mean()) / rewards.std()
for idx in range(num_generations):
completion_ids = prompt_completion_ids[idx][len(prompt_ids) :]
yield RolloutSample(
prompt=row["prompt"],
completion=completions[idx],
input_ids=prompt_ids + completion_ids,
completion_mask=[0] * len(prompt_ids) + [1] * len(completion_ids),
old_log_probs=[0.0] * len(prompt_ids) + [-0.5] * len(completion_ids),
advantage=float(advantages[idx]),
model_version=self._model_version,
group_id=group_id, # every completion of one prompt belongs to the same group
metrics={"reward": float(rewards[idx]), "reward_std": float(rewards.std())},
)
Comment thread
cursor[bot] marked this conversation as resolved.

def _fill_queue(self):
for _ in range(self._samples_per_weight_sync):
self.rollout_buffer.put(next(self._sample_iter))

def start(self):
self._fill_queue()

def update_model_version(self, version):
self._model_version = version
self._fill_queue()

def stop(self):
pass

def check_health(self, stale_after_s):
pass


class _NoOpWeightTransfer:
"""No-op `WeightTransferProtocol`, so the smoke exercises only the FSDP2 parameter lifecycle.

Without it `AsyncGRPOTrainer` builds the default `WeightTransferClient`, which needs a live vLLM server to stream
weights into over NCCL.
"""

def init_weight_transfer(self) -> None: ...

def pause(self) -> None: ...

def send_weights(self, iterator) -> None:
# Drain the iterator so the trainer's weight-gathering path still runs end to end.
for _ in iterator:
pass

def resume(self) -> None: ...

def destroy(self) -> None: ...


def main() -> None:
dataset = load_dataset("trl-internal-testing/zen", "conversational_prompt_completion", split="train")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)

# Same minimal, memory-frugal config as the existing single-process test_train, with 2 steps so we
# exercise the optimizer loop more than once under FSDP2.
args = AsyncGRPOConfig(
# The launcher runs this worker with `cwd` at the repo root, so a relative `output_dir` would drop trainer
# artifacts into the working tree. The test owns a temporary directory and passes it in.
output_dir=os.environ["ASYNC_GRPO_FSDP2_OUTPUT_DIR"],
learning_rate=0.1,
per_device_train_batch_size=3,
num_generations=3,
max_completion_length=8,
# 0 selects the count-based FixedCountBatcher and, being non-None, still skips the vLLM
# max_model_len lookup. A positive budget starves here: the stub's samples are short enough
# that 2 rank-rows never fill, so TokenBudgetBatcher would never emit a micro-batch.
token_budget=0,
max_steps=_MAX_STEPS,
vllm_server_timeout=5.0,
report_to="none",
)
Comment thread
cursor[bot] marked this conversation as resolved.
trainer = AsyncGRPOTrainer(
model=MODEL_ID,
reward_funcs=dummy_reward_func,
args=args,
train_dataset=dataset,
# 24 = 4 x microbatch_size (per_device_train_batch_size 3 x 2 ranks); max_steps=2 needs 2,
# so the initial fill covers the whole run without relying on a weight-sync refill.
rollout_worker=_StubRolloutWorker(tokenizer, dataset, num_generations=3, samples_per_weight_sync=24),
weight_transfer=_NoOpWeightTransfer(),
)
Comment thread
cursor[bot] marked this conversation as resolved.

# Snapshot params before training so we can confirm FSDP2 training actually updated them.
before = {n: p.detach().clone() for n, p in trainer.model.named_parameters()}

trainer.train()

# Did any parameter change? Materialize DTensors (full_tensor) and move both operands to CPU before
# comparing: the `before` snapshot is captured at construction (pre-FSDP-wrap, plain tensor) while the
# post-train param is an FSDP2 DTensor on CUDA, so a direct torch.equal would raise a device mismatch.
def _materialize(t):
t = t.full_tensor() if isinstance(t, torch.distributed.tensor.DTensor) else t
return t.detach().cpu()

# Compare every parameter on every rank before deciding: `_materialize` calls the collective
# `full_tensor()`, so breaking early would leave the ranks issuing different numbers of collectives and
# rank 0 would hang instead of reporting `params_changed: false`. A list comprehension is deliberate,
# since `any()` over a generator short-circuits the same way `break` does. Only rank 0's verdict is
# asserted: with `fsdp_cpu_ram_efficient_loading`, the pre-wrap snapshot on other ranks holds
# placeholders rather than the loaded weights.
diffs = [not torch.equal(_materialize(before[n]), _materialize(p)) for n, p in trainer.model.named_parameters()]
changed = any(diffs)

last = trainer.state.log_history[-1] if trainer.state.log_history else {}
train_loss = last.get("train_loss")
# The pytest side asserts on the launch shape too: a replicated single-process run would also change the
# parameters and report a finite loss, so world size, the distributed type and sharded parameters are reported.
accelerator = trainer.accelerator
result = {
"steps": trainer.state.global_step,
"max_steps": _MAX_STEPS,
"params_changed": changed,
"train_loss_finite": train_loss is not None and bool(np.isfinite(train_loss)),
"num_processes": accelerator.num_processes,
"distributed_type": accelerator.distributed_type.value,
"fsdp_version": accelerator.state.fsdp_plugin.fsdp_version if accelerator.state.fsdp_plugin else None,
"sharded_params": sum(isinstance(p, torch.distributed.tensor.DTensor) for p in trainer.model.parameters()),
}
# Only rank 0 prints the asserted line, so the pytest side parses exactly one result.
if trainer.accelerator.is_main_process:
print(f"{RESULT_PREFIX} {json.dumps(result)}", flush=True) # noqa: T201 - result channel for the launcher


if __name__ == "__main__":
main()
29 changes: 29 additions & 0 deletions tests/experimental/data/accelerate_configs/fsdp2_reshard.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
# 2-process FSDP2 config for the async-GRPO FSDP2 functional test (test_train_fsdp2).
#
# `fsdp_reshard_after_forward: true` is set explicitly so the test exercises the resharding parameter
# lifecycle rather than relying on FSDP's default. Mirrors `examples/accelerate_configs/fsdp2.yaml`
# with `num_processes: 2` for a 2-GPU node.
compute_environment: LOCAL_MACHINE
debug: false
distributed_type: FSDP
downcast_bf16: 'no'
enable_cpu_affinity: false
fsdp_config:
fsdp_activation_checkpointing: false
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_cpu_ram_efficient_loading: true
fsdp_offload_params: false
fsdp_reshard_after_forward: true
fsdp_state_dict_type: FULL_STATE_DICT
fsdp_version: 2
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 2
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
95 changes: 94 additions & 1 deletion tests/experimental/test_async_grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,10 @@
import multiprocessing as mp
import os
import queue
import signal
import subprocess
from collections import defaultdict
from pathlib import Path
from unittest.mock import MagicMock, patch

import numpy as np
Expand Down Expand Up @@ -54,7 +57,14 @@
)
from trl.trainer.base_trainer import _BaseTrainer

from ..testing_utils import TrlTestCase, is_ampere_or_newer
from ..testing_utils import TrlTestCase, is_ampere_or_newer, require_torch_multi_accelerator


ROOT = Path(__file__).resolve().parents[2]
_HERE = Path(__file__).parent
_FSDP2_WORKER = _HERE / "_async_grpo_fsdp2_worker.py"
_FSDP2_CONFIG = _HERE / "data" / "accelerate_configs" / "fsdp2_reshard.yaml"
_FSDP2_RESULT_PREFIX = "ASYNC_GRPO_FSDP2_RESULT"


# The trainer loads the model with Flash Attention, which requires a `head_size` multiple of 8. Hence the `small-*`
Expand All @@ -65,6 +75,22 @@ def dummy_reward_func(completions, **kwargs):
return [float(hash(c[0]["content"]) % 100) / 100.0 for c in completions]


def _descendant_pids(pid: int) -> list[int]:
"""Every process below `pid` in the /proc tree, so a timeout can kill ranks that torch elastic detached into their
own sessions."""
found, stack = [], [pid]
while stack:
parent = stack.pop()
try:
children = Path(f"/proc/{parent}/task/{parent}/children").read_text().split()
except OSError:
children = []
for child in map(int, children):
found.append(child)
stack.append(child)
return found


class _StubRolloutWorker:
"""Minimal rollout worker stub for testing the trainer in isolation."""

Expand Down Expand Up @@ -303,6 +329,73 @@ def reset(self, **kwargs): ...
environment_factory={"a": EnvA, "b": EnvB},
)

@pytest.mark.slow
@require_torch_multi_accelerator
def test_train_fsdp2(self):
# Functional smoke: AsyncGRPOTrainer trains under a 2-process FSDP2 group, confirming the optimizer
# updates the FSDP2-sharded parameters. Uses an in-process stub rollout worker (no vLLM server /
# NCCL weight transfer), so the only distributed surface is the FSDP2 parameter lifecycle.
#
# `@pytest.mark.slow` marks the cost, but no lane collects this test today: `slow_tests` is
# `pytest -m "slow" tests/`, and `norecursedirs` excludes `tests/experimental` from that recursive
# collection. `test_experimental` passes an explicit path so it does collect the test, but its runner
# is single-GPU, where `@require_torch_multi_accelerator` skips it.
#
# Pin the repo root onto PYTHONPATH for the child: `accelerate launch` re-execs each rank via
# torch.distributed.elastic, which sets sys.path[0] to the launched script's directory
# (tests/experimental/), not cwd. Without this, a non-editable `trl` already in site-packages
# would shadow the working tree and the test would exercise the wrong code.
env = os.environ.copy()
env["PYTHONPATH"] = os.pathsep.join([str(ROOT), env.get("PYTHONPATH", "")]).rstrip(os.pathsep)
# `cwd` below is the repo root, so hand the worker somewhere else to write trainer artifacts.
env["ASYNC_GRPO_FSDP2_OUTPUT_DIR"] = str(self.tmp_dir)
# Bound the child: the trainer's rollout consumer blocks indefinitely on an empty queue (it only
# calls `check_health`, which this stub implements as a no-op), so a starved run would hang the
# pytest process rather than fail it. The timeout turns that into a readable failure.
# `accelerate launch` starts each rank through torch elastic, which puts every worker in its own session
# (`start_new_session=True`), so neither `subprocess.run(timeout=...)` nor a kill of the launcher's process
# group reaches the ranks: they would keep the GPUs and hold the stdout pipe open, and the read after the
# timeout would never return. Collect the descendants from /proc while the launcher is still their parent
# and kill the whole tree.
proc = subprocess.Popen(
["accelerate", "launch", "--config_file", str(_FSDP2_CONFIG), str(_FSDP2_WORKER)],
env=env,
cwd=ROOT,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
try:
stdout, stderr = proc.communicate(timeout=900)
except subprocess.TimeoutExpired:
for pid in [proc.pid, *_descendant_pids(proc.pid)]:
try:
os.kill(pid, signal.SIGKILL)
except ProcessLookupError:
pass
stdout, stderr = proc.communicate()
pytest.fail(f"FSDP2 worker timed out after 900s:\nSTDOUT:\n{stdout}\nSTDERR:\n{stderr}")
Comment thread
cursor[bot] marked this conversation as resolved.
result = subprocess.CompletedProcess(proc.args, proc.returncode, stdout, stderr)
assert result.returncode == 0, f"FSDP2 worker failed:\nSTDOUT:\n{result.stdout}\nSTDERR:\n{result.stderr}"

result_lines = [ln for ln in result.stdout.splitlines() if ln.startswith(_FSDP2_RESULT_PREFIX)]
assert len(result_lines) == 1, f"expected exactly one result line, got {result_lines}\n{result.stdout}"
measured = json.loads(result_lines[0][len(_FSDP2_RESULT_PREFIX) :].strip())

# Training actually ran under FSDP2, produced a finite loss, and updated the parameters.
# The worker configures more than one step so the optimizer loop runs repeatedly under FSDP2; accepting fewer
# would let an early stop after step 1 pass.
assert measured["max_steps"] > 1, f"worker must configure more than one step: {measured}"
assert measured["steps"] == measured["max_steps"], f"not every configured step ran: {measured}"
assert measured["train_loss_finite"], f"train loss not finite under FSDP2: {measured}"
assert measured["params_changed"], f"parameters did not change under FSDP2: {measured}"
Comment thread
behroozazarkhalili marked this conversation as resolved.
# A replicated single-process run would pass the checks above too, so pin the launch shape the config asks
# for: two ranks, FSDP version 2, and parameters that are DTensors after wrapping.
assert measured["num_processes"] == 2, f"expected a two-process launch: {measured}"
assert measured["distributed_type"] == "FSDP", f"not launched under FSDP: {measured}"
assert measured["fsdp_version"] == 2, f"not FSDP version 2: {measured}"
assert measured["sharded_params"] > 0, f"no parameter was sharded as a DTensor: {measured}"


def _vision_parameter_names(model) -> set[str]:
"""Names of the parameters belonging to a vision-language model's vision tower.
Expand Down
Loading