diff --git a/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py b/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py index c4b890ffa..318cc2b28 100644 --- a/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py +++ b/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py @@ -692,7 +692,12 @@ def _flash_attn_fwd( # SplitKV uses float32 partial output, which doubles the O buffer size # in shared memory, causing OOM for diff-headdim (192, 128) - if arch // 10 in [10, 11] and head_dim != head_dim_v and num_splits > 1: + if ( + arch // 10 in [10, 11] + and head_dim != head_dim_v + and num_splits > 1 + and not (page_table is not None and q_stage == 1) + ): if num_n_blocks >= 64 and head_dim_v != 512: tile_n = 64 num_n_blocks = (seqlen_k_loaded + tile_n - 1) // tile_n diff --git a/python/sglang/srt/models/kimi_vl.py b/python/sglang/srt/models/kimi_vl.py index 27d1e1845..12c1171b4 100644 --- a/python/sglang/srt/models/kimi_vl.py +++ b/python/sglang/srt/models/kimi_vl.py @@ -131,7 +131,6 @@ class KimiVLForConditionalGeneration(nn.Module): config.vision_config, prefix=add_prefix("vision_tower", prefix), use_data_parallel=self.use_data_parallel, - use_tensor_parallel=not self.use_data_parallel, ) self.multi_modal_projector = KimiVLMultiModalProjector(config=config) @@ -263,6 +262,12 @@ class KimiVLForConditionalGeneration(nn.Module): use_default_weight_loading = False if "vision" in name: if self.vision_tower is not None: + # MoonViT's attention is wrapped in sglang's VisionAttention, + # whose sub-modules are named qkv_proj/proj instead of the + # checkpoint's wqkv/wo. + name = name.replace("wqkv.", "attn.qkv_proj.").replace( + "wo.", "attn.proj." + ) use_default_weight_loading = True else: for param_name, weight_name, shard_id in stacked_params_mapping: diff --git a/python/sglang/srt/models/kimi_vl_moonvit.py b/python/sglang/srt/models/kimi_vl_moonvit.py index 48cb383ce..d1bea46ac 100644 --- a/python/sglang/srt/models/kimi_vl_moonvit.py +++ b/python/sglang/srt/models/kimi_vl_moonvit.py @@ -52,18 +52,15 @@ import torch.nn.functional as F from transformers.activations import ACT2FN from transformers.modeling_utils import PreTrainedModel -from sglang.kernel_api_logging import debug_kernel_api - -try: - from flash_attn.flash_attn_interface import flash_attn_varlen_func -except ImportError: - flash_attn_varlen_func = None - from sglang.srt.configs import MoonViTConfig +from sglang.srt.layers.attention.vision import ( + VisionAttention, + VisionAttentionMetadata, + prepare_vision_attention_metadata, +) from sglang.srt.layers.conv import Conv2dLayer from sglang.srt.layers.linear import ( ColumnParallelLinear, - QKVParallelLinear, ReplicatedLinear, RowParallelLinear, ) @@ -75,112 +72,6 @@ from sglang.srt.utils import add_prefix, get_device _MAX_INFERENCE_POS_EMB_CACHE_ENTRIES = 256 -@debug_kernel_api -def multihead_attention( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - q_cu_seqlens: Optional[torch.Tensor] = None, - k_cu_seqlens: Optional[torch.Tensor] = None, - max_seqlen: Optional[int] = None, -): - """Multi-head attention using flash attention 2. - This function is used to handle the case where the query, key, and value are packed. - Args: - q, k, v: tensor of shape (tot_seqlens, num_heads, head_dim). - q_cu_seqlens (torch.Tensor): cumulative sequence lengths of q. - The first element should be 0 and the last element should be q.shape[0]. - k_cu_seqlens (torch.Tensor): cumulative sequence lengths of k. - The first element should be 0 and the last element should be k.shape[0]. - - Returns: - output: shape (batch_size, seqlen, dim) or (tot_seqlens, dim) if packing, - where dim = num_heads * head_dim - """ - if flash_attn_varlen_func is None: - raise ImportError( - "flash_attn is not installed, this function needs flash_attn_varlen_func from flash_attn" - ) - # Unified format legal check - assert q.dim() == k.dim() == v.dim() == 3, "q, k, v must have 3 dims" - # Keep validation on CPU for debugging, but avoid synchronizing the GPU - # once per MoonViT layer in the normal packed CUDA path. - if not q_cu_seqlens.is_cuda: - assert q_cu_seqlens[-1] == q.shape[0], "q_cu_seqlens must sum to q.shape[0]" - assert ( - k_cu_seqlens[-1] == k.shape[0] == v.shape[0] - ), "k_cu_seqlens must sum to k.shape[0]" - assert q.dtype in [ - torch.bfloat16, - torch.float16, - ], f"unsupported dtype {q.dtype} for multihead attn" - - if max_seqlen is None: - max_seqlen = (q_cu_seqlens[1:] - q_cu_seqlens[:-1]).max().item() - attn_out = flash_attn_varlen_func( - q, - k, - v, - q_cu_seqlens, - k_cu_seqlens, - max_seqlen, - max_seqlen, - causal=False, - ) - attn_out = attn_out.flatten(start_dim=-2) - - return attn_out - - -def sdpa_attention( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - q_cu_seqlens: Optional[torch.Tensor] = None, - k_cu_seqlens: Optional[torch.Tensor] = None, - max_seqlen: Optional[int] = None, -) -> torch.Tensor: - """Multi-head attention using torch scaled dot product attention. - This function is used to handle the case where the query, key, and value are packed. - Args: - q, k, v: tensor of shape (tot_seqlens, num_heads, head_dim). - q_cu_seqlens (torch.Tensor): cumulative sequence lengths of q. - The first element should be 0 and the last element should be q.shape[0]. - k_cu_seqlens (torch.Tensor): cumulative sequence lengths of k. - The first element should be 0 and the last element should be k.shape[0]. - - Returns: - output: shape (batch_size, seqlen, dim) or (tot_seqlens, dim) if packing, - where dim = num_heads * head_dim - """ - # Unified format legal check - assert q.dim() == k.dim() == v.dim() == 3, "q, k, v must have 3 dims" - assert q_cu_seqlens[-1] == q.shape[0], "q_cu_seqlens must sum to q.shape[0]" - seq_length = q.shape[0] - attention_mask = torch.zeros( - [1, seq_length, seq_length], device=q.device, dtype=torch.bool - ) - for i in range(1, len(q_cu_seqlens)): - attention_mask[ - ..., - q_cu_seqlens[i - 1] : q_cu_seqlens[i], - q_cu_seqlens[i - 1] : q_cu_seqlens[i], - ] = True - q = q.transpose(0, 1) - k = k.transpose(0, 1) - v = v.transpose(0, 1) - attn_output = F.scaled_dot_product_attention(q, k, v, attention_mask, dropout_p=0.0) - attn_output = attn_output.transpose(0, 1) - attn_output = attn_output.reshape(seq_length, -1) - return attn_output - - -VL_VISION_ATTENTION_FUNCTIONS = { - "flash_attention_2": multihead_attention, - "sdpa": sdpa_attention, -} - - def _apply_rope_input_validation(x, freqs_cis): assert x.ndim == freqs_cis.ndim + 1, (x.shape, freqs_cis.shape) assert x.shape[:-2] == freqs_cis.shape[:-1], (x.shape, freqs_cis.shape) @@ -189,7 +80,7 @@ def _apply_rope_input_validation(x, freqs_cis): def apply_rope( - xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor + xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor, x_shape=None ) -> tuple[torch.Tensor, torch.Tensor]: """ Args: (The leading dimensions of all inputs should be the same) @@ -518,24 +409,17 @@ class MoonVitEncoderLayer(nn.Module): hidden_dim: int, mlp_dim: int, *, - attn_implementation: str = "flash_attention_2", # use fa2 in sglang by default activation=F.gelu, attn_bias: bool = False, prefix: str = "", use_data_parallel: bool = False, - use_tensor_parallel: bool = False, ): super().__init__() self.num_heads = num_heads self.hidden_dim = hidden_dim self.hidden_size_per_attention_head = self.hidden_dim // self.num_heads - self.attn_implementation = attn_implementation - - self.use_tensor_parallel = use_tensor_parallel and not use_data_parallel - tp_size = get_parallel().attn_tp_size if self.use_tensor_parallel else 1 - tp_rank = get_parallel().attn_tp_rank if self.use_tensor_parallel else 0 - self.num_attention_heads_per_partition = self.num_heads // tp_size + use_tensor_parallel = not use_data_parallel self.norm0 = nn.LayerNorm(hidden_dim) self.norm1 = nn.LayerNorm(hidden_dim) self.mlp = MLP2( @@ -543,74 +427,20 @@ class MoonVitEncoderLayer(nn.Module): activation, prefix=add_prefix("mlp", prefix), use_data_parallel=use_data_parallel, - use_tensor_parallel=self.use_tensor_parallel, + use_tensor_parallel=use_tensor_parallel, ) - if self.use_tensor_parallel: - self.wqkv = QKVParallelLinear( - hidden_size=hidden_dim, - head_size=self.hidden_size_per_attention_head, - total_num_heads=num_heads, - total_num_kv_heads=num_heads, - bias=attn_bias, - prefix=add_prefix("wqkv", prefix), - tp_rank=tp_rank, - tp_size=tp_size, - ) - self.wo = RowParallelLinear( - hidden_dim, - hidden_dim, - bias=attn_bias, - prefix=add_prefix("wo", prefix), - tp_rank=tp_rank, - tp_size=tp_size, - ) - else: - self.wqkv = nn.Linear(hidden_dim, hidden_dim * 3, bias=attn_bias) - self.wo = nn.Linear(hidden_dim, hidden_dim, bias=attn_bias) - - def attention_qkvpacked( - self, - x: torch.Tensor, - cu_seqlens: torch.Tensor, - rope_freqs_cis: Optional[torch.Tensor] = None, - max_seqlen: Optional[int] = None, - ): - """ - Args: - x (torch.Tensor): (batch_size, seqlen, hidden_dim) - cu_seqlens (torch.Tensor): - """ - if self.use_tensor_parallel: - xqkv, _ = self.wqkv(x) - else: - xqkv = self.wqkv(x) - - qkv_shape = xqkv.size()[:-1] + ( - 3, - self.num_attention_heads_per_partition, - self.hidden_size_per_attention_head, + self.attn = VisionAttention( + embed_dim=hidden_dim, + num_heads=num_heads, + projection_size=hidden_dim, + use_qkv_parallel=True, + qkv_bias=attn_bias, + proj_bias=attn_bias, + flatten_batch=True, + prefix=add_prefix("attn", prefix), + use_data_parallel=use_data_parallel, + customized_position_embedding_applier=apply_rope, ) - # xqkv: (batch_size, seqlen, 3, nheads, headdim) - xqkv = xqkv.view(*qkv_shape) - xq, xk, xv = torch.unbind(xqkv, dim=-3) - - xq, xk = apply_rope(xq, xk, rope_freqs_cis) - - attn_func = VL_VISION_ATTENTION_FUNCTIONS[self.attn_implementation] - attn_out = attn_func( - xq, - xk, - xv, - q_cu_seqlens=cu_seqlens, - k_cu_seqlens=cu_seqlens, - max_seqlen=max_seqlen, - ) - - if self.use_tensor_parallel: - attn_out, _ = self.wo(attn_out) - else: - attn_out = self.wo(attn_out) - return attn_out def forward( self, @@ -618,6 +448,7 @@ class MoonVitEncoderLayer(nn.Module): cu_seqlens: torch.Tensor, rope_freqs_cis: Union[torch.Tensor, None] = None, max_seqlen: Optional[int] = None, + forward_metadata: Optional[VisionAttentionMetadata] = None, ) -> torch.Tensor: """ Args: @@ -628,13 +459,14 @@ class MoonVitEncoderLayer(nn.Module): """ residual = hidden_states hidden_states = self.norm0(hidden_states) - attn_out = self.attention_qkvpacked( + attn_out = self.attn( hidden_states, - cu_seqlens, - rope_freqs_cis=rope_freqs_cis, + cu_seqlens=cu_seqlens, + position_embeddings=rope_freqs_cis, + forward_metadata=forward_metadata, max_seqlen=max_seqlen, ) - hidden_states = residual + attn_out + hidden_states = residual + attn_out.view(residual.shape) residual = hidden_states hidden_states = self.mlp(self.norm1(hidden_states)) @@ -651,7 +483,6 @@ class MoonVitEncoder(nn.Module): block_cfg: dict, prefix: str = "", use_data_parallel: bool = False, - use_tensor_parallel: bool = False, ) -> None: super().__init__() @@ -663,7 +494,6 @@ class MoonVitEncoder(nn.Module): MoonVitEncoderLayer( prefix=add_prefix(f"blocks.{layer_idx}", prefix), use_data_parallel=use_data_parallel, - use_tensor_parallel=use_tensor_parallel, **block_cfg, ) for layer_idx in range(num_layers) @@ -689,12 +519,17 @@ class MoonVitEncoder(nn.Module): if max_seqlen is None: max_seqlen = (grid_hw[:, 0] * grid_hw[:, 1]).max().item() + forward_metadata = prepare_vision_attention_metadata( + cu_seqlens, device=hidden_states.device + ) + for _, block in enumerate(self.blocks): hidden_states = block( hidden_states, cu_seqlens, rope_freqs_cis=rope_freqs_cis, max_seqlen=max_seqlen, + forward_metadata=forward_metadata, ) hidden_states = self.final_layernorm(hidden_states) @@ -769,7 +604,6 @@ class MoonVitPretrainedModel(PreTrainedModel): config: MoonViTConfig, prefix: str = "", use_data_parallel: bool = False, - use_tensor_parallel: bool = False, *inputs, **kwargs, ): @@ -795,11 +629,9 @@ class MoonVitPretrainedModel(PreTrainedModel): "mlp_dim": config.intermediate_size, "activation": GELUTanh(), "attn_bias": True, - "attn_implementation": config._attn_implementation, }, prefix=add_prefix("encoder", prefix), use_data_parallel=use_data_parallel, - use_tensor_parallel=use_tensor_parallel, ) def forward( diff --git a/python/sglang/srt/models/nvila.py b/python/sglang/srt/models/nvila.py index 64e6d00ca..dce329f77 100644 --- a/python/sglang/srt/models/nvila.py +++ b/python/sglang/srt/models/nvila.py @@ -195,12 +195,10 @@ class NVILAForConditionalGeneration(nn.Module): for x, block_size in zip(vision_features_list, block_sizes) ] - vision_features = torch.stack( + vision_features = torch.cat( [einops.rearrange(x, "1 c h w -> (h w) c") for x in vision_features_list] ) - vision_features = einops.rearrange(vision_features, "n p d -> (n p) d") - return vision_features def load_weights(self, weights: Iterable[tuple[str, Tensor]]) -> None: diff --git a/test/registered/unit/models/test_kimi_vl.py b/test/registered/unit/models/test_kimi_vl.py index 828192f23..1cfa38a1c 100644 --- a/test/registered/unit/models/test_kimi_vl.py +++ b/test/registered/unit/models/test_kimi_vl.py @@ -14,9 +14,9 @@ from sglang.srt.layers.linear import ( ) from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.models.kimi_vl import KimiVLForConditionalGeneration -from sglang.srt.models.kimi_vl_moonvit import MoonVitEncoderLayer, multihead_attention +from sglang.srt.models.kimi_vl_moonvit import MoonVitEncoderLayer from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_context, get_parallel from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -80,17 +80,16 @@ class TestKimiVLEncoderParallelism(CustomTestCase): def test_moonvit_uses_tensor_parallel_layers(self): with get_parallel().override( tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 - ): + ), get_context().override_server_args(): layer = MoonVitEncoderLayer( num_heads=2, hidden_dim=8, mlp_dim=16, prefix="vision_tower.encoder.blocks.0", - use_tensor_parallel=True, ) - self.assertIsInstance(layer.wqkv, QKVParallelLinear) - self.assertIsInstance(layer.wo, RowParallelLinear) + self.assertIsInstance(layer.attn.qkv_proj, QKVParallelLinear) + self.assertIsInstance(layer.attn.proj, RowParallelLinear) self.assertIsInstance(layer.mlp.fc0, ColumnParallelLinear) self.assertIsInstance(layer.mlp.fc1, RowParallelLinear) @@ -142,18 +141,6 @@ class TestKimiVLEncoderParallelism(CustomTestCase): self.assertIsInstance(output, torch.Tensor) self.assertEqual(output.shape, (1, 4, 2)) - def test_moonvit_attention_accepts_precomputed_max_seqlen(self): - q = torch.randn(4, 2, 4, dtype=torch.bfloat16) - cu_seqlens = torch.tensor([0, 4], dtype=torch.int32) - fake_output = torch.randn_like(q) - with patch( - "sglang.srt.models.kimi_vl_moonvit.flash_attn_varlen_func", - return_value=fake_output, - ) as flash_attn: - output = multihead_attention(q, q, q, cu_seqlens, cu_seqlens, max_seqlen=4) - self.assertTrue(torch.equal(output, fake_output.flatten(start_dim=-2))) - self.assertEqual(flash_attn.call_args.args[5:7], (4, 4)) - if __name__ == "__main__": raise SystemExit(pytest.main([__file__, "-v"]))