From e6a64920572a5d93f7480c68bd78120c94f73ab0 Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 30 Aug 2026 21:02:14 +0800 Subject: [PATCH] [vlm] fix: preserve per-request vit graph metadata for qwen-vl (#37043) --- python/sglang/srt/models/qwen2_5_vl.py | 11 ++ python/sglang/srt/models/qwen3_vl.py | 2 + .../srt/multimodal/vit_cuda_graph_runner.py | 140 ++++++++++++------ .../test_vit_cuda_graph_metadata_cuda.py | 88 +++++++++++ .../multimodal/test_vit_cuda_graph_runner.py | 34 +++++ 5 files changed, 230 insertions(+), 45 deletions(-) create mode 100644 test/registered/unit/multimodal/test_vit_cuda_graph_metadata_cuda.py diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index 459489bd9..3b58a0282 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -616,6 +616,11 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin): rotary_pos_emb = self.rot_pos_emb(grid_thw) window_index, cu_window_seqlens = self.get_window_index(grid_thw) + cu_window_layout = tuple( + value + for index, value in enumerate(cu_window_seqlens) + if index == 0 or value != cu_window_seqlens[index - 1] + ) cu_window_seqlens = torch.tensor( cu_window_seqlens, device=x.device, @@ -658,6 +663,11 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin): ] ) cu_seqlens = torch.cat([cu_seqlens.new_zeros(1), cu_seqlens]) + full_layout = [0, 0] + total_tokens = 0 + for temporal, height, width in grid_thw.tolist(): + total_tokens += temporal * height * width + full_layout.append(total_tokens) return self.cuda_graph_runner.run( x=x, @@ -665,6 +675,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin): cu_seqlens=cu_seqlens, cu_window_seqlens=cu_window_seqlens, output_indices=reverse_indices, + attention_layout_key=(tuple(full_layout), cu_window_layout), ) diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index c985225c7..2b7f1cc06 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -1053,6 +1053,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): rotary_pos_emb_cos, rotary_pos_emb_sin, ) = self._prepare_graph_inputs(x, grid_thw) + attention_layout_key = (tuple(cu_seqlens.tolist()), None) if not isinstance(cu_seqlens, torch.Tensor): cu_seqlens = torch.tensor(cu_seqlens, device=x.device, dtype=torch.int32) else: @@ -1067,6 +1068,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): cu_seqlens=cu_seqlens, cu_window_seqlens=None, output_indices=None, + attention_layout_key=attention_layout_key, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: diff --git a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py index fd6b5d246..a2e6933ba 100644 --- a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py @@ -60,8 +60,13 @@ class ViTCudaGraphRunner: self.cu_full_len_kk: Dict[Hashable, torch.Tensor] = {} self.cu_window_len_kk: Dict[Hashable, torch.Tensor] = {} - # rotary position buffers shared across graphs + # Current rotary workspace plus older allocations retained by graphs + # captured before the workspace grew. self.sin_cos_ws: Optional[Tuple[torch.Tensor, torch.Tensor]] = None + self._retired_sin_cos_ws: List[Tuple[torch.Tensor, torch.Tensor]] = [] + self._sin_cos_ws_by_graph: Dict[Hashable, Tuple[torch.Tensor, torch.Tensor]] = ( + {} + ) self.max_context_len = getattr(vit, "max_context_len", None) # Qwen2.5-VL specific viarable. @@ -91,31 +96,67 @@ class ViTCudaGraphRunner: def dtype(self) -> torch.dtype: return self.vit.dtype - def _ensure_sin_cos_ws(self, seq_len: int, head_dim: int): - if self.sin_cos_ws is None: - max_shape = self.max_context_len or seq_len - max_shape = max(max_shape, seq_len) + def _get_sin_cos_ws( + self, graph_key: Hashable, seq_len: int, head_dim: int + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Return the stable rotary buffers captured by one graph.""" + graph_ws = self._sin_cos_ws_by_graph.get(graph_key) + if graph_ws is not None: + return graph_ws + + needs_new_workspace = self.sin_cos_ws is None or ( + self.sin_cos_ws[0].size(0) < seq_len + or self.sin_cos_ws[0].size(1) < head_dim + ) + if needs_new_workspace: + previous = self.sin_cos_ws + previous_seq_len = previous[0].size(0) if previous is not None else 0 + previous_head_dim = previous[0].size(1) if previous is not None else 0 + max_shape = max( + self.max_context_len or 0, + previous_seq_len * 2, + seq_len, + ) + max_head_dim = max(previous_head_dim, head_dim) cos_ws = torch.empty( - max_shape, head_dim, dtype=self.dtype, device=self.device + max_shape, max_head_dim, dtype=self.dtype, device=self.device ) sin_ws = torch.empty( - max_shape, head_dim, dtype=self.dtype, device=self.device + max_shape, max_head_dim, dtype=self.dtype, device=self.device ) + if previous is not None: + # CUDA graphs retain captured addresses, so an older allocation + # cannot be freed when a larger request grows the workspace. + self._retired_sin_cos_ws.append(previous) self.sin_cos_ws = (cos_ws, sin_ws) - else: - if self.sin_cos_ws[0].size(0) < seq_len: - max_shape = max(self.sin_cos_ws[0].size(0) * 2, seq_len) - cos_ws = torch.empty( - max_shape, head_dim, dtype=self.dtype, device=self.device - ) - sin_ws = torch.empty( - max_shape, head_dim, dtype=self.dtype, device=self.device - ) - self.sin_cos_ws = (cos_ws, sin_ws) - def _get_graph_key(self, x_3d: torch.Tensor) -> int: - # x_3d: [S, B, H], B=1, S as graph_key - return x_3d.shape[0] + graph_ws = ( + self.sin_cos_ws[0][:seq_len, :head_dim], + self.sin_cos_ws[1][:seq_len, :head_dim], + ) + self._sin_cos_ws_by_graph[graph_key] = graph_ws + return graph_ws + + @staticmethod + def _sequence_layout_key(cu_seqlens: Optional[torch.Tensor]) -> Optional[tuple]: + if cu_seqlens is None: + return None + return tuple(int(value) for value in cu_seqlens.tolist()) + + def _get_graph_key( + self, + x_3d: torch.Tensor, + cu_seqlens: torch.Tensor, + cu_window_seqlens: Optional[torch.Tensor], + attention_layout_key: Optional[Hashable] = None, + ) -> Hashable: + """Include attention boundaries so equal-length batches stay distinct.""" + if attention_layout_key is None: + attention_layout_key = ( + self._sequence_layout_key(cu_seqlens), + self._sequence_layout_key(cu_window_seqlens), + ) + return (x_3d.shape[0], attention_layout_key) def _capture_context(self): # A DP-sharded encoder intentionally lets each rank capture only the @@ -265,15 +306,22 @@ class ViTCudaGraphRunner: self, x_3d: torch.Tensor, # [S, 1, H] cu_seqlens: torch.Tensor, - cu_window_seqlens: torch.Tensor, + cu_window_seqlens: Optional[torch.Tensor], position_embeddings: Optional[ Tuple[torch.Tensor, torch.Tensor] ], # (cos, sin), [S, D] rotary_pos_emb_cos: Optional[torch.Tensor] = None, rotary_pos_emb_sin: Optional[torch.Tensor] = None, - ) -> int: + attention_layout_key: Optional[Hashable] = None, + ) -> Hashable: vit = self.vit - graph_key = self._get_graph_key(x_3d) + graph_key = self._get_graph_key( + x_3d, + cu_seqlens, + cu_window_seqlens, + attention_layout_key, + ) + seq_len = x_3d.shape[0] if graph_key in self.block_graphs: return graph_key @@ -291,7 +339,7 @@ class ViTCudaGraphRunner: x_3d, device=self.device ).contiguous() self.block_ws[graph_key] = torch.empty( - graph_key, + seq_len, num_heads, attn_head_dim, device=self.device, @@ -315,12 +363,10 @@ class ViTCudaGraphRunner: self.block_input[graph_key].copy_(x_3d) if position_embeddings is not None: - # make sure rotary workspace head_dim = position_embeddings[0].shape[1] - self._ensure_sin_cos_ws(graph_key, head_dim) - - used_cos_ws = self.sin_cos_ws[0][:graph_key, :] - used_sin_ws = self.sin_cos_ws[1][:graph_key, :] + used_cos_ws, used_sin_ws = self._get_sin_cos_ws( + graph_key, seq_len, head_dim + ) used_cos_ws.copy_(position_embeddings[0]) used_sin_ws.copy_(position_embeddings[1]) persist_position_embeddings = (used_cos_ws, used_sin_ws) @@ -328,12 +374,10 @@ class ViTCudaGraphRunner: graph_key=graph_key, position_embeddings=persist_position_embeddings ) elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None: - # make sure rotary workspace head_dim = rotary_pos_emb_cos.shape[1] - self._ensure_sin_cos_ws(graph_key, head_dim) - - used_cos_ws = self.sin_cos_ws[0][:graph_key, :] - used_sin_ws = self.sin_cos_ws[1][:graph_key, :] + used_cos_ws, used_sin_ws = self._get_sin_cos_ws( + graph_key, seq_len, head_dim + ) used_cos_ws.copy_(rotary_pos_emb_cos) used_sin_ws.copy_(rotary_pos_emb_sin) self._create_graph( @@ -347,7 +391,7 @@ class ViTCudaGraphRunner: def replay( self, - graph_key: int, + graph_key: Hashable, x_3d: torch.Tensor, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, rotary_pos_emb_cos: Optional[torch.Tensor] = None, @@ -355,20 +399,19 @@ class ViTCudaGraphRunner: output_indices: Optional[torch.Tensor] = None, ) -> torch.Tensor: + seq_len = x_3d.shape[0] if position_embeddings is not None: - # update rotary workspace content head_dim = position_embeddings[0].shape[1] - self._ensure_sin_cos_ws(graph_key, head_dim) - used_cos_ws = self.sin_cos_ws[0][:graph_key, :] - used_sin_ws = self.sin_cos_ws[1][:graph_key, :] + used_cos_ws, used_sin_ws = self._get_sin_cos_ws( + graph_key, seq_len, head_dim + ) used_cos_ws.copy_(position_embeddings[0]) used_sin_ws.copy_(position_embeddings[1]) elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None: - # update rotary workspace content head_dim = rotary_pos_emb_cos.shape[1] - self._ensure_sin_cos_ws(graph_key, head_dim) - used_cos_ws = self.sin_cos_ws[0][:graph_key, :] - used_sin_ws = self.sin_cos_ws[1][:graph_key, :] + used_cos_ws, used_sin_ws = self._get_sin_cos_ws( + graph_key, seq_len, head_dim + ) used_cos_ws.copy_(rotary_pos_emb_cos) used_sin_ws.copy_(rotary_pos_emb_sin) @@ -390,15 +433,21 @@ class ViTCudaGraphRunner: self, x: torch.Tensor, cu_seqlens: torch.Tensor, - cu_window_seqlens: torch.Tensor, + cu_window_seqlens: Optional[torch.Tensor], position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]], rotary_pos_emb_cos: Optional[torch.Tensor] = None, rotary_pos_emb_sin: Optional[torch.Tensor] = None, output_indices: Optional[torch.Tensor] = None, + attention_layout_key: Optional[Hashable] = None, ) -> torch.Tensor: # x: [seq_len, hidden] -> [S, B=1, H] x_3d = x.unsqueeze(1) - graph_key = self._get_graph_key(x_3d) + graph_key = self._get_graph_key( + x_3d, + cu_seqlens, + cu_window_seqlens, + attention_layout_key, + ) if graph_key not in self.block_graphs: self.create_graph( @@ -408,6 +457,7 @@ class ViTCudaGraphRunner: cu_window_seqlens=cu_window_seqlens, rotary_pos_emb_cos=rotary_pos_emb_cos, rotary_pos_emb_sin=rotary_pos_emb_sin, + attention_layout_key=attention_layout_key, ) return self.replay( diff --git a/test/registered/unit/multimodal/test_vit_cuda_graph_metadata_cuda.py b/test/registered/unit/multimodal/test_vit_cuda_graph_metadata_cuda.py new file mode 100644 index 000000000..3dabc128b --- /dev/null +++ b/test/registered/unit/multimodal/test_vit_cuda_graph_metadata_cuda.py @@ -0,0 +1,88 @@ +import sys + +import pytest +import torch +from torch import nn + +from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-large") + + +class _BoundaryBlock(nn.Module): + def __init__(self): + super().__init__() + self.attn = type( + "AttentionConfig", + (), + { + "num_attention_heads_per_partition": 1, + "head_size": 1, + "qkv_backend_name": "triton_attn", + }, + )() + + def forward( + self, + x, + *, + cu_seqlens, + position_embeddings, + output_ws=None, + ): + boundary = cu_seqlens[0][1].to(x.dtype) + position = position_embeddings[0][: x.shape[0], :1].unsqueeze(1) + return x + boundary + position + + +class _Merger(nn.Module): + def forward(self, x): + return x.squeeze(1) + + +class _VisionTower(nn.Module): + def __init__(self): + super().__init__() + self.blocks = nn.ModuleList([_BoundaryBlock()]) + self.merger = _Merger() + self.use_data_parallel = True + self.deepstack_visual_indexes = [] + self.deepstack_merger_list = None + self.max_context_len = None + self.register_buffer("anchor", torch.empty(0, device="cuda")) + + @property + def device(self): + return self.anchor.device + + @property + def dtype(self): + return torch.float32 + + +def test_vit_graph_replays_current_attention_and_position_metadata(): + runner = ViTCudaGraphRunner(_VisionTower()) + + def run(seq_len, boundaries, position): + x = torch.zeros(seq_len, 1, device="cuda") + cu_seqlens = torch.tensor(boundaries, dtype=torch.int32, device="cuda") + positions = torch.full((seq_len, 1), position, device="cuda") + output = runner.run(x, cu_seqlens, None, (positions, positions)) + torch.cuda.synchronize() + return output.cpu() + + first = run(4, [0, 2, 4], 1) + different_layout = run(4, [0, 1, 4], 1) + run(8, [0, 8], 7) + small_after_growth = run(4, [0, 2, 4], 5) + + torch.testing.assert_close(first, torch.full_like(first, 3)) + torch.testing.assert_close(different_layout, torch.full_like(different_layout, 2)) + torch.testing.assert_close( + small_after_growth, torch.full_like(small_after_growth, 7) + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py b/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py index 4d6c15a2f..427e12091 100644 --- a/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py +++ b/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py @@ -3,6 +3,7 @@ from types import SimpleNamespace from unittest.mock import patch import pytest +import torch from sglang.test.ci.ci_register import register_cpu_ci @@ -76,6 +77,39 @@ def test_vit_graph_runner_caches_resolved_backend_name(): assert runner._attn_backend == "fa3" +def test_vit_graph_key_includes_full_and_window_attention_boundaries(): + runner = _runner(use_data_parallel=True) + x = torch.empty(8, 1, 4) + + first = runner._get_graph_key( + x, + torch.tensor([0, 4, 8]), + torch.tensor([0, 2, 4, 8]), + ) + second = runner._get_graph_key( + x, + torch.tensor([0, 2, 8]), + torch.tensor([0, 4, 6, 8]), + ) + + assert first != second + + +def test_vit_graph_keeps_rotary_workspace_address_after_growth(): + runner = _runner(use_data_parallel=True) + runner.vit.device = torch.device("cpu") + runner.vit.dtype = torch.float32 + + small = runner._get_sin_cos_ws("small", seq_len=4, head_dim=2) + small_address = small[0].data_ptr() + runner._get_sin_cos_ws("large", seq_len=16, head_dim=2) + + assert runner._get_sin_cos_ws("small", seq_len=4, head_dim=2)[0].data_ptr() == ( + small_address + ) + assert len(runner._retired_sin_cos_ws) == 1 + + def test_internvl_graph_runner_caches_resolved_backend_name(): attention = SimpleNamespace( qkv_backend_name="triton_attn",