From 5b76f55d900a33069ce1fbb78e7b51c5867aa201 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Wed, 1 Jul 2026 08:49:58 +0800 Subject: [PATCH] [Fix]: Inline H2D during CUDA graph capture to avoid stream isolation in Offloader (#29166) Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- python/sglang/srt/utils/offloader.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/utils/offloader.py b/python/sglang/srt/utils/offloader.py index 56a96ecc0..a2e1df8ac 100644 --- a/python/sglang/srt/utils/offloader.py +++ b/python/sglang/srt/utils/offloader.py @@ -306,6 +306,10 @@ class _ModuleOffloader(ABC): param_offloader.post_init() def start_onload(self): + if torch.cuda.is_current_stream_capturing(): + self._device_tensors = self._create_device_tensors() + self._load_event = None + return self.alt_stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(self.alt_stream): self._device_tensors = self._create_device_tensors() @@ -318,7 +322,13 @@ class _ModuleOffloader(ABC): def wait_and_get_device_tensors(self): assert self._device_tensors is not None - self._load_event.wait() + if torch.cuda.is_current_stream_capturing(): + if self._load_event is not None: + self._device_tensors = self._create_device_tensors() + self._load_event = None + return self._device_tensors + if self._load_event is not None: + self._load_event.wait() return self._device_tensors def _create_device_tensors(self):