diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index ed3f045274e..7f35ad010e7 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -1269,7 +1269,7 @@ def reset_prefix_cache(self): "reset_prefix_cache can only be used with async_engine=False. Use reset_prefix_cache_async instead." ) - self.llm.llm_engine.reset_prefix_cache() + self.llm.llm_engine.reset_prefix_cache(reset_connector=True) gc.collect() torch.cuda.empty_cache() @@ -1285,7 +1285,7 @@ def sleep(self): ) # Reset the prefix cache to ensure that prefix cache is not reused after weights are updated - self.llm.llm_engine.reset_prefix_cache() + self.llm.llm_engine.reset_prefix_cache(reset_connector=True) # Clear the renderer's multimodal processor cache (sender side) so it # stays in sync with the receiver cache that vLLM clears internally # during sleep. Without this, the sender thinks images are already diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index 201d0790044..f3e581fd2be 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -188,6 +188,13 @@ def _create_engine(self, llm_kwargs: dict[str, Any]) -> None: ) llm_kwargs["compilation_config"] = CompilationConfig(**compilation_config) + if isinstance(llm_kwargs.get("kv_transfer_config"), dict): + from vllm.config import KVTransferConfig + + llm_kwargs["kv_transfer_config"] = KVTransferConfig( + **llm_kwargs["kv_transfer_config"] + ) + self.llm_async_engine_args = AsyncEngineArgs(**llm_kwargs) self.stat_loggers = ( [PrometheusStatLogger] @@ -1459,7 +1466,7 @@ async def reset_prefix_cache_async(self): "reset_prefix_cache_async can only be used with async_engine=True. Use reset_prefix_cache instead." ) - await self.llm.reset_prefix_cache() + await self.llm.reset_prefix_cache(reset_connector=True) gc.collect() torch.cuda.empty_cache() @@ -1475,7 +1482,7 @@ async def sleep_async(self): ) # Reset the prefix cache to ensure that prefix cache is not reused after weights are updated - await self.llm.reset_prefix_cache() + await self.llm.reset_prefix_cache(reset_connector=True) # Reset the multimodal processor cache (sender side) so it stays in # sync with the receiver cache that vLLM clears internally during # sleep. Without this, the sender thinks images are already cached on