From feae615b1146ec7869c203ddf7e997c3a3a59f6e Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Wed, 14 Jan 2026 17:29:23 +0800 Subject: [PATCH] [VLM] Support ViT CUDA Graph for InternVL (#16732) --- python/sglang/srt/models/internvl.py | 30 ++- python/sglang/srt/models/qwen2_5_vl.py | 12 +- .../internvl_vit_cuda_graph_runner.py | 183 ++++++++++++++++++ 3 files changed, 219 insertions(+), 6 deletions(-) create mode 100644 python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py diff --git a/python/sglang/srt/models/internvl.py b/python/sglang/srt/models/internvl.py index a7897cf92..b7be90101 100644 --- a/python/sglang/srt/models/internvl.py +++ b/python/sglang/srt/models/internvl.py @@ -13,6 +13,7 @@ from sglang.srt.distributed import ( get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) +from sglang.srt.environ import envs from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention.vision import SingletonCache, VisionAttention @@ -36,6 +37,9 @@ from sglang.srt.models.internlm2 import InternLM2ForCausalLM from sglang.srt.models.qwen2 import Qwen2ForCausalLM from sglang.srt.models.qwen3 import Qwen3ForCausalLM from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM +from sglang.srt.multimodal.internvl_vit_cuda_graph_runner import ( + InternViTCudaGraphRunner, +) from sglang.srt.multimodal.mm_utils import run_dp_sharded_vision_model from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import is_cuda @@ -82,8 +86,9 @@ class InternAttention(nn.Module): self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, + output_ws: Optional[torch.Tensor] = None, ) -> torch.Tensor: - out = self.attn(hidden_states, cu_seqlens=cu_seqlens) + out = self.attn(hidden_states, cu_seqlens=cu_seqlens, output_ws=output_ws) outs = self.proj_drop(out) return outs @@ -256,6 +261,7 @@ class InternVisionEncoderLayer(nn.Module): self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, + output_ws: Optional[torch.Tensor] = None, ) -> Tuple[ torch.FloatTensor, Optional[torch.FloatTensor], @@ -268,7 +274,9 @@ class InternVisionEncoderLayer(nn.Module): hidden_states = hidden_states + self.drop_path1( self.attn( - self.norm1(hidden_states).to(hidden_states.dtype), cu_seqlens=cu_seqlens + self.norm1(hidden_states).to(hidden_states.dtype), + cu_seqlens=cu_seqlens, + output_ws=output_ws, ) * self.ls1 ) @@ -303,7 +311,11 @@ class InternVisionEncoder(nn.Module): x.item() for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers) ] - aux_stream = torch.cuda.Stream() if _is_cuda else None + + self.enable_cg = _is_cuda and envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get() + aux_stream = ( + None if self.enable_cg else (torch.cuda.Stream() if _is_cuda else None) + ) self.layers = nn.ModuleList( [ InternVisionEncoderLayer( @@ -313,6 +325,10 @@ class InternVisionEncoder(nn.Module): ] ) + self.cuda_graph_runner: Optional[InternViTCudaGraphRunner] = None + if self.enable_cg: + self.cuda_graph_runner = InternViTCudaGraphRunner(self) + def forward( self, inputs_embeds, @@ -329,6 +345,14 @@ class InternVisionEncoder(nn.Module): return_dict (`bool`, *optional*): Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple. """ + if self.enable_cg and (not output_hidden_states): + # graph path only returns last_hidden_state + hidden_states = inputs_embeds.to(device=inputs_embeds.device).contiguous() + hidden_states = self.cuda_graph_runner.run(hidden_states) + if not return_dict: + return (hidden_states,) + return BaseModelOutput(last_hidden_state=hidden_states, hidden_states=None) + output_hidden_states = ( output_hidden_states if output_hidden_states is not None diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index 7f44978fc..da64dc2de 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -76,7 +76,9 @@ from sglang.srt.models.utils import RotaryPosMixin, WeightsMapper, permute_inv from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner from sglang.srt.server_args import get_global_server_args -from sglang.srt.utils import add_prefix, is_npu +from sglang.srt.utils import add_prefix, is_cuda, is_npu + +_is_cuda = is_cuda() logger = logging.getLogger(__name__) @@ -328,7 +330,11 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin): 1 if use_data_parallel else get_tensor_model_parallel_world_size() ) self.max_context_len = max_context_len - self.cuda_graph_runner: Optional[ViTCudaGraphRunner] = ViTCudaGraphRunner(self) + self.enable_cg = _is_cuda and envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get() + + self.cuda_graph_runner: Optional[ViTCudaGraphRunner] = None + if self.enable_cg: + self.cuda_graph_runner = ViTCudaGraphRunner(self) def get_window_index(self, grid_thw): cu_window_seqlens: list = [0] @@ -400,7 +406,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin): x: torch.Tensor, grid_thw: torch.Tensor, ) -> torch.Tensor: - if envs.SGLANG_VIT_ENABLE_CUDA_GRAPH.get(): + if self.enable_cg: return self.forward_with_cuda_graph(x, grid_thw) # patchify diff --git a/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py new file mode 100644 index 000000000..07bc3e77d --- /dev/null +++ b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py @@ -0,0 +1,183 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""ViT CUDA Graph Runner class.""" +from __future__ import annotations + +from typing import Dict, Hashable, Tuple + +import torch +import torch.nn as nn + +from sglang.srt.layers.attention.vision import VisionAttention +from sglang.srt.server_args import get_global_server_args + + +class InternViTCudaGraphRunner: + """CUDA Graph runner for InternVL vision encoder. + + Captures: + y = layer_N(...layer_2(layer_1(x))) + + Keyed by (B, S). This is REQUIRED because InternVL uses [B,S,H]. + """ + + def __init__(self, encoder: nn.Module) -> None: + self.encoder = encoder + + # key -> graph & stable buffers + self.graphs: Dict[Hashable, torch.cuda.CUDAGraph] = {} + self.inp: Dict[Hashable, torch.Tensor] = {} + self.ws: Dict[Hashable, torch.Tensor] = {} + self.out: Dict[Hashable, torch.Tensor] = {} + + # key -> stable cu_seqlens buffers (addresses must be stable) + self.cu: Dict[Hashable, torch.Tensor] = {} + self.cu_kk: Dict[Hashable, torch.Tensor] = {} + + # cache attention metadata + first_layer = encoder.layers[0] + # InternAttention wraps VisionAttention as first_layer.attn.attn + self._attn: VisionAttention = first_layer.attn.attn # type: ignore + + @property + def device(self) -> torch.device: + return next(self.encoder.parameters()).device + + @property + def dtype(self) -> torch.dtype: + return next(self.encoder.parameters()).dtype + + def _graph_key(self, x: torch.Tensor) -> Tuple[int, int]: + # x: [B,S,H] + return (x.shape[0], x.shape[1]) + + def _build_cu(self, B: int, S: int, device: torch.device) -> torch.Tensor: + # [0, S, 2S, ..., B*S] + return torch.arange(0, (B + 1) * S, step=S, device=device, dtype=torch.int32) + + def _alloc_ws( + self, B: int, S: int, H: int, device: torch.device, dtype: torch.dtype + ) -> torch.Tensor: + # InternVL shape: [tokens, nheads, head_dim] + tokens = B * S + + num_heads = getattr(self._attn, "num_attention_heads_per_partition", None) + if num_heads is None: + num_heads = getattr(self._attn, "num_heads", None) + if num_heads is None: + raise RuntimeError("Cannot infer num_heads from VisionAttention") + + head_dim = getattr(self._attn, "head_size", None) + if head_dim is None: + # fallback (should rarely happen) + head_dim = H // int(num_heads) + + return torch.empty( + tokens, + int(num_heads), + int(head_dim), + device=device, + dtype=dtype, + ) + + def _warmup_once(self, key: Hashable) -> None: + """Run a tiny eager warmup on the preallocated buffers to trigger lazy init.""" + override_backend = get_global_server_args().mm_attention_backend + cu = self.cu[key] + cu_kk = self.cu_kk[key] + max_len = int(cu_kk.max().item()) if cu_kk.numel() else 0 + + if override_backend == "triton_attn": + cu_ws = [cu, cu_kk, max_len] + elif override_backend == "fa3": + cu_ws = [cu, max_len] + else: + raise RuntimeError("Not supported ViT attention backend for InternVL CG") + + x = self.inp[key] + y = x + with torch.no_grad(): + for blk in self.encoder.layers: + y = blk(y, cu_seqlens=cu_ws, output_ws=self.ws[key]) + + def _capture_graph(self, key: Hashable) -> None: + g = torch.cuda.CUDAGraph() + override_backend = get_global_server_args().mm_attention_backend + + cu = self.cu[key] + cu_kk = self.cu_kk[key] + max_len = int(cu_kk.max().item()) if cu_kk.numel() else 0 + + if override_backend == "triton_attn": + cu_ws = [cu, cu_kk, max_len] + elif override_backend == "fa3": + cu_ws = [cu, max_len] + else: + raise RuntimeError("Not supported ViT attention backend for InternVL CG") + + torch.cuda.synchronize() + + with torch.cuda.graph(g): + y = self.inp[key] + for blk in self.encoder.layers: + y = blk(y, cu_seqlens=cu_ws, output_ws=self.ws[key]) + # y is a stable output tensor produced during capture; keep reference + self.out[key] = y + + self.graphs[key] = g + + def create_graph(self, x: torch.Tensor) -> Hashable: + # x: [B, S, H] + x = x.contiguous() + key = self._graph_key(x) + if key in self.graphs: + return key + + B, S, H = x.shape + device = x.device + dtype = x.dtype + + # stable input buffer + self.inp[key] = torch.empty_like(x, device=device).contiguous() + + # stable cu buffers + cu = self._build_cu(B, S, device=device) + self.cu[key] = cu + self.cu_kk[key] = cu[1:] - cu[:-1] + + # stable attention workspace + self.ws[key] = self._alloc_ws(B, S, H, device=device, dtype=dtype) + + self.inp[key].copy_(x) + self._warmup_once(key) + + # capture + self._capture_graph(key) + return key + + def run(self, x: torch.Tensor) -> torch.Tensor: + # x: [B, S, H] + x = x.contiguous() + key = self._graph_key(x) + if key not in self.graphs: + self.create_graph(x) + + # update input content (address stable) + self.inp[key].copy_(x) + + # replay + self.graphs[key].replay() + + return self.out[key]