From e828e149b65c286f0982173a266fccfcd56ac2e1 Mon Sep 17 00:00:00 2001 From: aminediro Date: Sat, 30 May 2026 20:43:05 +0000 Subject: [PATCH] async_grpo: implement vLLM native 4-phase weight-sync (start/finish_weight_update), require vllm>=0.22.0 --- .../async_grpo/async_grpo_trainer.py | 1 - .../async_grpo/weight_transfer.py | 24 ++++++++++--------- 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/trl/experimental/async_grpo/async_grpo_trainer.py b/trl/experimental/async_grpo/async_grpo_trainer.py index 78a008a7ff8..3672cd0eb6c 100644 --- a/trl/experimental/async_grpo/async_grpo_trainer.py +++ b/trl/experimental/async_grpo/async_grpo_trainer.py @@ -442,7 +442,6 @@ def __init__( "dtype_names": weight_dtype_names, "shapes": weight_shapes, "packed": True, - "is_checkpoint_format": True, }, ) self.rollout_worker = AsyncRolloutWorker( diff --git a/trl/experimental/async_grpo/weight_transfer.py b/trl/experimental/async_grpo/weight_transfer.py index d0af83c2d2a..16e0b0646c5 100644 --- a/trl/experimental/async_grpo/weight_transfer.py +++ b/trl/experimental/async_grpo/weight_transfer.py @@ -21,7 +21,7 @@ from trl.import_utils import is_vllm_available -if is_vllm_available(min_version="0.17.1"): +if is_vllm_available(min_version="0.22.0"): from vllm.distributed.weight_transfer.nccl_engine import NCCLTrainerSendWeightsArgs, NCCLWeightTransferEngine from vllm.utils.network_utils import get_ip, get_open_port @@ -37,9 +37,9 @@ def __init__( server_timeout: float = 240.0, init_weight_transfer_timeout: int = 1800, ): - if not is_vllm_available(min_version="0.17.1"): + if not is_vllm_available(min_version="0.22.0"): raise ImportError( - "vLLM >= 0.17.1 is required to use WeightTransferClient. Install it with: pip install 'vllm>=0.17.1'" + "vLLM >= 0.22.0 is required to use WeightTransferClient. Install it with: pip install 'vllm>=0.22.0'" ) self.vllm_server_url = vllm_server_url.rstrip("/") self.server_timeout = server_timeout @@ -103,25 +103,27 @@ def send_weights(self, iterator) -> None: if self.model_update_group is None: return t0 = time.time() + # Prepare the workers for the reload; must complete before any weights are sent. + requests.post( + f"{self.vllm_server_url}/start_weight_update", + json={"is_checkpoint_format": True}, + timeout=1800, + ) + # The /update_weights POST drives the workers' blocking NCCL recv, so it runs on a thread + # concurrently with the trainer-side broadcast. t_update = threading.Thread( target=requests.post, args=(f"{self.vllm_server_url}/update_weights",), kwargs={"json": {"update_info": self._weight_update_info}, "timeout": 1800}, ) t_update.start() - logger.debug(f"[weight_sync] /update_weights POST sent ({time.time() - t0:.1f}s)") - t_nccl = time.time() NCCLWeightTransferEngine.trainer_send_weights( iterator=iterator, trainer_args=NCCLTrainerSendWeightsArgs(group=self.model_update_group, packed=True), ) - logger.debug(f"[weight_sync] NCCL transfer took {time.time() - t_nccl:.1f}s") - t_join = time.time() t_update.join() - logger.debug( - f"[weight_sync] /update_weights join took {time.time() - t_join:.1f}s " - f"(total send_weights: {time.time() - t0:.1f}s)" - ) + requests.post(f"{self.vllm_server_url}/finish_weight_update", timeout=1800) + logger.debug(f"[weight_sync] send_weights took {time.time() - t0:.1f}s") def pause(self) -> None: t0 = time.time()