From 6760d815af318db23f9eedaf48db93e1123115db Mon Sep 17 00:00:00 2001 From: Guyue Huang Date: Thu, 30 Apr 2026 11:59:08 -0700 Subject: [PATCH 1/2] Discard weight when finish generation in the main loop Signed-off-by: Guyue Huang --- nemo_rl/algorithms/grpo.py | 6 +++++- nemo_rl/models/generation/vllm/vllm_generation.py | 2 ++ nemo_rl/models/generation/vllm/vllm_worker.py | 4 ++-- nemo_rl/models/generation/vllm/vllm_worker_async.py | 4 ++-- 4 files changed, 11 insertions(+), 5 deletions(-) diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index d26fb0bcae4..d5578874472 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -1586,7 +1586,11 @@ def grpo_train( max_rollout_turns=master_config.grpo["max_rollout_turns"], greedy=False, ) - policy_generation.finish_generation() + policy_generation.finish_generation( + discard_weights=colocated_inference + ) + if colocated_inference: + POLICY_GENERATION_STALE = True # Collect generation logger metrics for performance reporting after each generation step # inflight batch sizes and num pending samples are collected from each worker if policy_generation is not None: diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index 1bd20f5cbba..063cee005a6 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -732,10 +732,12 @@ def finish_generation(self, *args: Any, **kwargs: Any) -> bool: if self.cfg["vllm_cfg"]["async_engine"] else "reset_prefix_cache" ) + kwargs = {} # Use run_all_workers_single_data for methods that don't need data futures = self.worker_group.run_all_workers_single_data( method_name, run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + **kwargs, ) # Wait for all futures to complete results = ray.get(futures) diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index a91d1fac335..7255fcb4a32 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -986,7 +986,7 @@ def reset_prefix_cache(self): gc.collect() torch.cuda.empty_cache() - def sleep(self): + def sleep(self, discard_weights: bool = False): """Put the vLLM engine to sleep.""" assert self.llm is not None, ( "Attempting to sleep with either an uninitialized vLLM or non-model-owner" @@ -1009,7 +1009,7 @@ def sleep(self): self.llm.renderer, "clear_mm_cache" ): self.llm.renderer.clear_mm_cache() - self.llm.sleep(level=1) + self.llm.sleep(level=2 if discard_weights else 1) gc.collect() torch.cuda.empty_cache() diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index d6089ffb068..7fcd6cd0db6 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -1129,7 +1129,7 @@ async def reset_prefix_cache_async(self): gc.collect() torch.cuda.empty_cache() - async def sleep_async(self): + async def sleep_async(self, discard_weights: bool = False): """Async version of sleep.""" assert self.llm is not None, ( "Attempting to sleep with either an uninitialized vLLM or non-model-owner" @@ -1148,7 +1148,7 @@ async def sleep_async(self): # the receiver and sends data=None, causing an assertion error. if hasattr(self.llm, "reset_mm_cache"): await self.llm.reset_mm_cache() - await self.llm.sleep(level=1) + await self.llm.sleep(level=2 if discard_weights else 1) gc.collect() torch.cuda.empty_cache() From 8d80498c07aaf0f8af359794bcf0b437d1d4c6dc Mon Sep 17 00:00:00 2001 From: Guyue Huang Date: Tue, 19 May 2026 09:38:26 -0700 Subject: [PATCH 2/2] Add UT Signed-off-by: Guyue Huang --- tests/unit/algorithms/test_grpo.py | 1 + .../models/generation/test_vllm_generation.py | 83 ++++++++++++++++++- 2 files changed, 83 insertions(+), 1 deletion(-) diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index a6193f9f854..4b361ccac91 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -1398,6 +1398,7 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): ) assert policy_generation.clear_logger_metrics.called + policy_generation.finish_generation.assert_called_once_with(discard_weights=True) assert policy_generation.get_logger_metrics.called assert any( "generation_logger_metrics" in call.args[0] diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index 09793914e21..c8d1a6c1561 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -16,7 +16,7 @@ import os from copy import deepcopy from pathlib import Path -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest import ray @@ -33,7 +33,9 @@ GenerationDatumSpec, ) from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration +from nemo_rl.models.generation.vllm.vllm_worker import VllmGenerationWorkerImpl from nemo_rl.models.generation.vllm.vllm_worker_async import ( + VllmAsyncGenerationWorkerImpl, _replace_prefix_tokens, ) from nemo_rl.models.policy import LoRAConfig, PolicyConfig @@ -144,6 +146,85 @@ } +@pytest.mark.parametrize( + "colocated,async_engine,expected_method,expected_kwargs", + [ + (True, False, "sleep", {"discard_weights": True}), + (True, True, "sleep_async", {"discard_weights": True}), + (False, False, "reset_prefix_cache", {}), + (False, True, "reset_prefix_cache_async", {}), + ], +) +def test_vllm_finish_generation_routes_discard_weights( + monkeypatch, colocated, async_engine, expected_method, expected_kwargs +): + vllm_generation = VllmGeneration.__new__(VllmGeneration) + vllm_generation.cfg = { + "colocated": {"enabled": colocated}, + "vllm_cfg": {"async_engine": async_engine}, + } + vllm_generation.worker_group = MagicMock() + vllm_generation.worker_group.run_all_workers_single_data.return_value = [ + "worker_future" + ] + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_generation.ray.get", + lambda futures: [True], + ) + + assert vllm_generation.finish_generation(discard_weights=True) + + vllm_generation.worker_group.run_all_workers_single_data.assert_called_once_with( + expected_method, + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + **expected_kwargs, + ) + + +def test_vllm_worker_sleep_uses_discard_weight_level(monkeypatch): + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + worker.cfg = {"vllm_cfg": {"async_engine": False}} + worker.llm = MagicMock() + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker.gc.collect", lambda: None + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker.torch.cuda.empty_cache", + lambda: None, + ) + + worker.sleep(discard_weights=False) + worker.llm.sleep.assert_called_once_with(level=1) + + worker.llm.sleep.reset_mock() + worker.sleep(discard_weights=True) + worker.llm.sleep.assert_called_once_with(level=2) + + +@pytest.mark.asyncio +async def test_vllm_async_worker_sleep_uses_discard_weight_level(monkeypatch): + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + worker.cfg = {"vllm_cfg": {"async_engine": True}} + worker.llm = MagicMock() + worker.llm.reset_prefix_cache = AsyncMock() + worker.llm.reset_mm_cache = AsyncMock() + worker.llm.sleep = AsyncMock() + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker_async.gc.collect", lambda: None + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker_async.torch.cuda.empty_cache", + lambda: None, + ) + + await worker.sleep_async(discard_weights=False) + worker.llm.sleep.assert_awaited_once_with(level=1) + + worker.llm.sleep.reset_mock() + await worker.sleep_async(discard_weights=True) + worker.llm.sleep.assert_awaited_once_with(level=2) + + def test_configure_generation_config_uses_real_startup_weights_without_draft_refit(): """Speculative training should not start the drafter from dummy weights without refit.""" vllm_config = deepcopy(basic_vllm_test_config)