[VLM] Support ViT CUDA Graph for InternVL (#16732)
This commit is contained in:
@@ -13,6 +13,7 @@ from sglang.srt.distributed import (
|
|||||||
get_tensor_model_parallel_rank,
|
get_tensor_model_parallel_rank,
|
||||||
get_tensor_model_parallel_world_size,
|
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.activation import get_act_fn
|
||||||
from sglang.srt.layers.attention import vision_utils
|
from sglang.srt.layers.attention import vision_utils
|
||||||
from sglang.srt.layers.attention.vision import SingletonCache, VisionAttention
|
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.qwen2 import Qwen2ForCausalLM
|
||||||
from sglang.srt.models.qwen3 import Qwen3ForCausalLM
|
from sglang.srt.models.qwen3 import Qwen3ForCausalLM
|
||||||
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM
|
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.multimodal.mm_utils import run_dp_sharded_vision_model
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
@@ -82,8 +86,9 @@ class InternAttention(nn.Module):
|
|||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
cu_seqlens: torch.Tensor,
|
cu_seqlens: torch.Tensor,
|
||||||
|
output_ws: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> 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)
|
outs = self.proj_drop(out)
|
||||||
return outs
|
return outs
|
||||||
|
|
||||||
@@ -256,6 +261,7 @@ class InternVisionEncoderLayer(nn.Module):
|
|||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
cu_seqlens: torch.Tensor,
|
cu_seqlens: torch.Tensor,
|
||||||
|
output_ws: Optional[torch.Tensor] = None,
|
||||||
) -> Tuple[
|
) -> Tuple[
|
||||||
torch.FloatTensor,
|
torch.FloatTensor,
|
||||||
Optional[torch.FloatTensor],
|
Optional[torch.FloatTensor],
|
||||||
@@ -268,7 +274,9 @@ class InternVisionEncoderLayer(nn.Module):
|
|||||||
|
|
||||||
hidden_states = hidden_states + self.drop_path1(
|
hidden_states = hidden_states + self.drop_path1(
|
||||||
self.attn(
|
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
|
* self.ls1
|
||||||
)
|
)
|
||||||
@@ -303,7 +311,11 @@ class InternVisionEncoder(nn.Module):
|
|||||||
x.item()
|
x.item()
|
||||||
for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers)
|
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(
|
self.layers = nn.ModuleList(
|
||||||
[
|
[
|
||||||
InternVisionEncoderLayer(
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
inputs_embeds,
|
inputs_embeds,
|
||||||
@@ -329,6 +345,14 @@ class InternVisionEncoder(nn.Module):
|
|||||||
return_dict (`bool`, *optional*):
|
return_dict (`bool`, *optional*):
|
||||||
Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
|
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 = (
|
||||||
output_hidden_states
|
output_hidden_states
|
||||||
if output_hidden_states is not None
|
if output_hidden_states is not None
|
||||||
|
|||||||
@@ -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.mm_utils import run_dp_sharded_mrope_vision_model
|
||||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||||
from sglang.srt.server_args import get_global_server_args
|
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__)
|
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()
|
1 if use_data_parallel else get_tensor_model_parallel_world_size()
|
||||||
)
|
)
|
||||||
self.max_context_len = max_context_len
|
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):
|
def get_window_index(self, grid_thw):
|
||||||
cu_window_seqlens: list = [0]
|
cu_window_seqlens: list = [0]
|
||||||
@@ -400,7 +406,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
grid_thw: torch.Tensor,
|
grid_thw: torch.Tensor,
|
||||||
) -> 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)
|
return self.forward_with_cuda_graph(x, grid_thw)
|
||||||
|
|
||||||
# patchify
|
# patchify
|
||||||
|
|||||||
@@ -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]
|
||||||
Reference in New Issue
Block a user