diff --git a/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py b/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py index 368db2ebc719..83646155cac7 100644 --- a/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py +++ b/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py @@ -124,6 +124,12 @@ def capture_one( if post_warmup_hook is not None: post_warmup_hook() + # The syncs above order work issued before each warmup, not the last + # warmup's own. A JIT module that finishes loading mid-capture issues + # illegal driver calls on the capturing stream, so drain once more here. + self._device_module.synchronize() + self._tp_group.barrier() + graph = BreakableCUDAGraph(self.deduped_cuda_graph) captured_fn = ( eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn diff --git a/python/sglang/srt/model_executor/runner_backend/full_cuda_graph_backend.py b/python/sglang/srt/model_executor/runner_backend/full_cuda_graph_backend.py index 3a047c8ffa64..508ab1b840d6 100644 --- a/python/sglang/srt/model_executor/runner_backend/full_cuda_graph_backend.py +++ b/python/sglang/srt/model_executor/runner_backend/full_cuda_graph_backend.py @@ -157,6 +157,12 @@ def capture_one( self._reuse_output_buffer = self._output_buffer is not None del warmup_output + # The syncs above order work issued before each warmup, not the last + # warmup's own. A JIT module that finishes loading mid-capture issues + # illegal driver calls on the capturing stream, so drain once more here. + self._device_module.synchronize() + self._tp_group.barrier() + graph = torch.cuda.CUDAGraph() graph_ctx: Callable[..., AbstractContextManager]