diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 0350038e9..cbd14f618 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -13,6 +13,10 @@ import torch import triton from sglang.srt.layers.attention.base_attn_backend import AttentionBackend +from sglang.srt.layers.attention.triton_ops.aiter_unified_attention import ( + scatter_ragged_to_page_table_kernel, + scatter_req_to_token_to_page_table_kernel, +) from sglang.srt.layers.attention.utils import ( create_flashinfer_kv_indices_triton, create_flashmla_kv_indices_triton, @@ -106,6 +110,7 @@ class ForwardMetadata: global_workspace_buffer = None + _AITER_PARTITION_SIZE_ROCM = 256 @@ -179,6 +184,11 @@ class AiterAttnBackend(AttentionBackend): self.qo_indptr = torch.zeros( (max_bs + 1,), dtype=torch.int32, device=model_runner.device ) + # qo_indptr for the unified-attn decode path (q_len == 1 per request) + # is always arange(0, bs+1); precompute once to avoid a per-step cumsum. + self.qo_indptr_unified_decode = torch.arange( + 0, max_bs + 1, dtype=torch.int32, device=model_runner.device + ) self.mask_indptr = torch.zeros( (max_bs + 1,), dtype=torch.int64, device=model_runner.device ) @@ -208,6 +218,17 @@ class AiterAttnBackend(AttentionBackend): "SGLANG_USE_AITER_UNIFIED_ATTN" ) + # When topk == 1 the EAGLE draft chain is linear, so target_verify's + # mask reduces to pure causal and can go through unified_attention + # instead of the legacy triton extend_attention_fwd. Gated on non-MLA + # (MLA has its own verify path) and env var for opt-out. + self._use_unified_verify = ( + self.use_triton_unified_attention + and not self.use_mla + and self.topk == 1 + and get_bool_env_var("SGLANG_AITER_UNIFIED_VERIFY", "1") + ) + # aiter kernel related initialization self.max_num_partitions = ( self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1 @@ -485,6 +506,128 @@ class AiterAttnBackend(AttentionBackend): ) return page_table[:, strided_indices] // page_size + def _build_unified_page_table_from_spec( + self, + spec_info, + bs: int, + dest_buf: Optional[torch.Tensor] = None, + swa_dest_buf: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Convert ragged (token-level) kv_indices from spec_info into a 2D + block-level page_table of shape (bs, max_num_blocks_per_seq). + unified_attention expects max_seqlen_k = page_table.shape[1] * + page_size to be a captured constant, so rows are sized to the + backend-level max_num_blocks_per_seq regardless of seqused_k. + """ + kv_indptr = spec_info.kv_indptr + kv_flat = spec_info.kv_indices + page_size = self.page_size + max_blocks = (self.max_context_len + page_size - 1) // page_size + + swa_slot_mapping = None + swa_page_table = None + + if dest_buf is not None: + # The scatter kernel fills [0, num_blocks) and loads past that use + # other=0, so the tail is 0-filled. Under graph replay rows > bs + # are stale but unified_attention only walks rows [0, bs). + page_table = dest_buf + else: + page_table = torch.zeros( + bs, max_blocks, dtype=torch.int32, device=self.device + ) + + if self.use_sliding_window_kv_pool: + swa_slot_mapping = self.token_to_kv_pool.full_to_swa_index_mapping.long() + + if swa_dest_buf is not None: + swa_page_table = swa_dest_buf + else: + swa_page_table = torch.zeros( + bs, max_blocks, dtype=torch.int32, device=self.device + ) + + BLOCK_SIZE = 1024 + grid = (bs, triton.cdiv(max(max_blocks, 1), BLOCK_SIZE)) + scatter_ragged_to_page_table_kernel[grid]( + kv_flat, + kv_indptr, + page_table, + page_table.stride(0), + swa_page_table, + swa_slot_mapping, + PAGE_SIZE=page_size, + BLOCK_SIZE=BLOCK_SIZE, + HAS_SWA=(swa_slot_mapping is not None), + ) + + return page_table, swa_page_table + + def _build_verify_unified_metadata( + self, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + draft_num: int, + page_table_dest: Optional[torch.Tensor] = None, + swa_page_table_dest: Optional[torch.Tensor] = None, + ): + """Build the 2D block page_table + qo_indptr for EAGLE target_verify + through unified_attention. Assumes the new draft K/V have already been + written by set_kv_buffer, so req_to_token[rp, :seq_lens[i]+draft_num] + covers both the prefix and the freshly committed draft tokens. Returns + (page_table, qo_indptr, max_q_len=draft_num). + """ + device = seq_lens.device + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( + 0, + (1 + bs) * draft_num, + step=draft_num, + dtype=torch.int32, + device=device, + ) + + page_size = self.page_size + max_blocks = (self.max_context_len + page_size - 1) // page_size + + swa_slot_mapping = None + swa_page_table = None + + if page_table_dest is not None: + page_table = page_table_dest + else: + page_table = torch.zeros(bs, max_blocks, dtype=torch.int32, device=device) + + if self.use_sliding_window_kv_pool: + swa_slot_mapping = self.token_to_kv_pool.full_to_swa_index_mapping.long() + + if swa_page_table_dest is not None: + swa_page_table = swa_page_table_dest + else: + swa_page_table = torch.zeros( + bs, max_blocks, dtype=torch.int32, device=device + ) + + BLOCK_SIZE = 1024 + grid = (bs, triton.cdiv(max(max_blocks, 1), BLOCK_SIZE)) + scatter_req_to_token_to_page_table_kernel[grid]( + self.req_to_token, + req_pool_indices, + seq_lens, + page_table, + self.req_to_token.stride(0), + page_table.stride(0), + swa_page_table, + swa_slot_mapping, + DRAFT_NUM=draft_num, + PAGE_SIZE=page_size, + BLOCK_SIZE=BLOCK_SIZE, + HAS_SWA=(swa_slot_mapping is not None), + ) + + return page_table, qo_indptr, draft_num, swa_page_table + def _resolve_v2_num_draft_tokens( self, extend_seq_lens: Optional[torch.Tensor] = None, @@ -729,14 +872,20 @@ class AiterAttnBackend(AttentionBackend): elif self.page_size > 1: kv_indices = self._transform_table_1_to_real(kv_indices) - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[1 : bs + 1] = torch.cumsum( - self.kv_last_page_len[:bs], dim=0 - ) + qo_indptr = self.qo_indptr_unified_decode[: bs + 1] else: - kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices - bs = kv_indptr.shape[0] - 1 + if self.use_triton_unified_attention and not self.use_mla: + bs = spec_info.kv_indptr.shape[0] - 1 + kv_indices, swa_page_table = ( + self._build_unified_page_table_from_spec(spec_info, bs) + ) + max_q_len = 1 + qo_indptr = self.qo_indptr_unified_decode[: bs + 1] + kv_indptr = None + else: + kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices + bs = kv_indptr.shape[0] - 1 if self.use_mla: qo_indptr = self.qo_indptr_[: bs + 1] @@ -1038,51 +1187,71 @@ class AiterAttnBackend(AttentionBackend): run_graph=False, ) else: - # Non-MLA target_verify: use triton extend kernel with custom mask bs = len(forward_batch.req_pool_indices) draft_num = spec_info.draft_token_num - qo_indptr = torch.arange( - 0, - (1 + bs) * draft_num, - step=draft_num, - dtype=torch.int32, - device=self.device, - ) + if self._use_unified_verify: + page_table, qo_indptr, max_q_len, swa_page_table = ( + self._build_verify_unified_metadata( + bs, + forward_batch.seq_lens, + forward_batch.req_pool_indices, + draft_num, + ) + ) + max_kv_len = page_table.shape[1] * self.page_size + self.forward_metadata = ForwardMetadata( + None, # kv_indptr unused in unified-verify path + page_table, # 2D block page_table stored in kv_indices + qo_indptr, + None, + max_q_len, + max_kv_len, + max_extend_len=max_q_len, + swa_page_table=swa_page_table, + ) + else: + qo_indptr = torch.arange( + 0, + (1 + bs) * draft_num, + step=draft_num, + dtype=torch.int32, + device=self.device, + ) - kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] + kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) + kv_indptr = kv_indptr[: bs + 1] - kv_indices = torch.empty( - kv_indptr[-1], dtype=torch.int64, device=self.device - ) - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - forward_batch.req_pool_indices, - forward_batch.seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) + kv_indices = torch.empty( + kv_indptr[-1], dtype=torch.int64, device=self.device + ) + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + forward_batch.req_pool_indices, + forward_batch.seq_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) - custom_mask = spec_info.custom_mask - seq_mask_len = draft_num * (forward_batch.seq_lens + draft_num) - mask_indptr = self.mask_indptr - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) - mask_indptr = mask_indptr[: bs + 1] + custom_mask = spec_info.custom_mask + seq_mask_len = draft_num * (forward_batch.seq_lens + draft_num) + mask_indptr = self.mask_indptr + mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) + mask_indptr = mask_indptr[: bs + 1] - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - None, - draft_num, - None, - custom_mask=custom_mask, - mask_indptr=mask_indptr, - max_extend_len=draft_num, - ) + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + None, + draft_num, + None, + custom_mask=custom_mask, + mask_indptr=mask_indptr, + max_extend_len=draft_num, + ) else: prefix_lens = forward_batch.extend_prefix_lens @@ -1299,7 +1468,12 @@ class AiterAttnBackend(AttentionBackend): kv_last_page_len = None max_q_len = None - if spec_info is None: + if spec_info is None or ( + self.use_triton_unified_attention and not self.use_mla + ): + max_num_blocks_per_seq = ( + self.max_context_len + self.page_size - 1 + ) // self.page_size if not self.use_triton_unified_attention: kv_indptr = self.kv_indptr @@ -1317,43 +1491,50 @@ class AiterAttnBackend(AttentionBackend): ) else: max_q_len = 1 - max_num_blocks_per_seq = ( - self.max_context_len + self.page_size - 1 - ) // self.page_size kv_indices = self.cuda_graph_kv_indices.view( -1, max_num_blocks_per_seq ) - page_indices = self.req_to_token[req_pool_indices[:bs], :max_kv_len] - if self.use_sliding_window_kv_pool: - swa_page_indices = ( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - page_indices - ) - ) - - page_indices = self._transform_table_1_to_real(page_indices) - swa_page_indices = self._transform_table_1_to_real( - swa_page_indices - ) - - new_rows = swa_page_indices.shape[0] - new_cols = swa_page_indices.shape[1] - - kv_indices[:new_rows, :new_cols].copy_(page_indices) swa_page_table = self.cuda_graph_swa_page_table - swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices) - elif self.page_size > 1: - page_indices = self._transform_table_1_to_real(page_indices) - new_rows = page_indices.shape[0] - new_cols = page_indices.shape[1] - kv_indices[:new_rows, :new_cols].copy_(page_indices) - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[1 : bs + 1] = torch.cumsum( - self.cuda_graph_kv_last_page_len[:bs], dim=0 - ) + if spec_info is not None: + self._build_unified_page_table_from_spec( + spec_info, + bs, + dest_buf=kv_indices, + swa_dest_buf=swa_page_table, + ) + else: + page_indices = self.req_to_token[ + req_pool_indices[:bs], :max_kv_len + ] + + if self.use_sliding_window_kv_pool: + swa_page_indices = ( + self.token_to_kv_pool.translate_loc_from_full_to_swa( + page_indices + ) + ) + + page_indices = self._transform_table_1_to_real(page_indices) + swa_page_indices = self._transform_table_1_to_real( + swa_page_indices + ) + + new_rows = swa_page_indices.shape[0] + new_cols = swa_page_indices.shape[1] + + kv_indices[:new_rows, :new_cols].copy_(page_indices) + swa_page_table = self.cuda_graph_swa_page_table + swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices) + elif self.page_size > 1: + page_indices = self._transform_table_1_to_real(page_indices) + new_rows = page_indices.shape[0] + new_cols = page_indices.shape[1] + kv_indices[:new_rows, :new_cols].copy_(page_indices) + + qo_indptr = self.qo_indptr_unified_decode[: bs + 1] kv_indptr = None else: @@ -1441,7 +1622,6 @@ class AiterAttnBackend(AttentionBackend): if self.use_mla: if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch self.make_mla_meta_data( @@ -1484,24 +1664,63 @@ class AiterAttnBackend(AttentionBackend): num_kv_splits=num_kv_splits, ) else: - custom_mask = self.cuda_graph_custom_mask - custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask - seq_mask_len = max_q_len * (seq_lens + max_q_len) - mask_indptr = self.mask_indptr - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) - mask_indptr = mask_indptr[: bs + 1] + if self._use_unified_verify: + max_num_blocks_per_seq = ( + self.max_context_len + self.page_size - 1 + ) // self.page_size + page_table = self.cuda_graph_kv_indices.view( + -1, max_num_blocks_per_seq + )[:bs] - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - custom_mask=custom_mask, - mask_indptr=mask_indptr, - max_extend_len=max_q_len, - ) + swa_page_table = None + + if self.use_sliding_window_kv_pool: + swa_page_table = self.cuda_graph_swa_page_table.view( + -1, max_num_blocks_per_seq + )[:bs] + + _page_table, _qo_indptr, _max_q_len, _swa_page_table = ( + self._build_verify_unified_metadata( + bs, + seq_lens, + req_pool_indices, + self.num_draft_tokens, + page_table_dest=page_table, + swa_page_table_dest=swa_page_table, + ) + ) + max_kv_len = max_num_blocks_per_seq * self.page_size + self.forward_metadata = ForwardMetadata( + None, + _page_table, + _qo_indptr, + kv_last_page_len, + _max_q_len, + max_kv_len, + max_extend_len=_max_q_len, + swa_page_table=_swa_page_table, + ) + else: + custom_mask = self.cuda_graph_custom_mask + custom_mask[: spec_info.custom_mask.shape[0]] = ( + spec_info.custom_mask + ) + seq_mask_len = max_q_len * (seq_lens + max_q_len) + mask_indptr = self.mask_indptr + mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) + mask_indptr = mask_indptr[: bs + 1] + + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + kv_last_page_len, + max_q_len, + max_kv_len, + custom_mask=custom_mask, + mask_indptr=mask_indptr, + max_extend_len=max_q_len, + ) elif forward_mode.is_draft_extend_v2(): # EAGLE V2: Uses fixed num_draft_tokens per batch self._ensure_spec_v2_topk_supported() @@ -1593,7 +1812,6 @@ class AiterAttnBackend(AttentionBackend): max_q_len = num_tokens_per_bs if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch self.make_mla_meta_data( @@ -1682,7 +1900,13 @@ class AiterAttnBackend(AttentionBackend): kv_last_page_len = None max_q_len = None - if spec_info is None: + if spec_info is None or ( + self.use_triton_unified_attention and not self.use_mla + ): + max_num_blocks_per_seq = ( + self.max_context_len + self.page_size - 1 + ) // self.page_size + if not self.use_triton_unified_attention: kv_indptr = self.kv_indptr kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) @@ -1699,43 +1923,50 @@ class AiterAttnBackend(AttentionBackend): ) else: max_q_len = 1 - max_num_blocks_per_seq = ( - self.max_context_len + self.page_size - 1 - ) // self.page_size kv_indices = self.cuda_graph_kv_indices.view( -1, max_num_blocks_per_seq ) - page_indices = self.req_to_token[req_pool_indices[:bs], :max_kv_len] - if self.use_sliding_window_kv_pool: - swa_page_indices = ( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - page_indices - ) - ) - - page_indices = self._transform_table_1_to_real(page_indices) - swa_page_indices = self._transform_table_1_to_real( - swa_page_indices - ) - - new_rows = swa_page_indices.shape[0] - new_cols = swa_page_indices.shape[1] - - kv_indices[:new_rows, :new_cols].copy_(page_indices) swa_page_table = self.cuda_graph_swa_page_table - swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices) - elif self.page_size > 1: - page_indices = self._transform_table_1_to_real(page_indices) - new_rows = page_indices.shape[0] - new_cols = page_indices.shape[1] - kv_indices[:new_rows, :new_cols].copy_(page_indices) - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[1 : bs + 1] = torch.cumsum( - self.cuda_graph_kv_last_page_len[:bs], dim=0 - ) + if spec_info is not None: + self._build_unified_page_table_from_spec( + spec_info, + bs, + dest_buf=kv_indices, + swa_dest_buf=swa_page_table, + ) + else: + page_indices = self.req_to_token[ + req_pool_indices[:bs], :max_kv_len + ] + + if self.use_sliding_window_kv_pool: + swa_page_indices = ( + self.token_to_kv_pool.translate_loc_from_full_to_swa( + page_indices + ) + ) + + page_indices = self._transform_table_1_to_real(page_indices) + swa_page_indices = self._transform_table_1_to_real( + swa_page_indices + ) + + new_rows = swa_page_indices.shape[0] + new_cols = swa_page_indices.shape[1] + + kv_indices[:new_rows, :new_cols].copy_(page_indices) + swa_page_table = self.cuda_graph_swa_page_table + swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices) + elif self.page_size > 1: + page_indices = self._transform_table_1_to_real(page_indices) + new_rows = page_indices.shape[0] + new_cols = page_indices.shape[1] + kv_indices[:new_rows, :new_cols].copy_(page_indices) + + qo_indptr = self.qo_indptr_unified_decode[: bs + 1] kv_indptr = None else: @@ -1825,7 +2056,6 @@ class AiterAttnBackend(AttentionBackend): if self.use_mla: if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch self.make_mla_meta_data( @@ -1868,23 +2098,63 @@ class AiterAttnBackend(AttentionBackend): num_kv_splits=num_kv_splits, ) else: - custom_mask = self.cuda_graph_custom_mask - custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask - seq_mask_len = max_q_len * (seq_lens + max_q_len) - mask_indptr = self.mask_indptr[: bs + 1] - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) + if self._use_unified_verify: + max_num_blocks_per_seq = ( + self.max_context_len + self.page_size - 1 + ) // self.page_size + page_table = self.cuda_graph_kv_indices.view( + -1, max_num_blocks_per_seq + )[:bs] - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - custom_mask=custom_mask, - mask_indptr=mask_indptr, - max_extend_len=max_q_len, - ) + swa_page_table = None + + if self.use_sliding_window_kv_pool: + swa_page_table = self.cuda_graph_swa_page_table.view( + -1, max_num_blocks_per_seq + )[:bs] + + _page_table, _qo_indptr, _max_q_len, _swa_page_table = ( + self._build_verify_unified_metadata( + bs, + seq_lens, + req_pool_indices, + self.num_draft_tokens, + page_table_dest=page_table, + swa_page_table_dest=swa_page_table, + ) + ) + + max_kv_len_unified = max_num_blocks_per_seq * self.page_size + self.forward_metadata = ForwardMetadata( + None, + _page_table, + _qo_indptr, + kv_last_page_len, + _max_q_len, + max_kv_len_unified, + max_extend_len=_max_q_len, + swa_page_table=_swa_page_table, + ) + else: + custom_mask = self.cuda_graph_custom_mask + custom_mask[: spec_info.custom_mask.shape[0]] = ( + spec_info.custom_mask + ) + seq_mask_len = max_q_len * (seq_lens + max_q_len) + mask_indptr = self.mask_indptr[: bs + 1] + mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) + + self.forward_metadata = ForwardMetadata( + kv_indptr, + kv_indices, + qo_indptr, + kv_last_page_len, + max_q_len, + max_kv_len, + custom_mask=custom_mask, + mask_indptr=mask_indptr, + max_extend_len=max_q_len, + ) elif forward_mode.is_draft_extend_v2(): # EAGLE V2: Fixed num_draft_tokens per batch self._ensure_spec_v2_topk_supported() @@ -1913,7 +2183,6 @@ class AiterAttnBackend(AttentionBackend): max_q_len = num_tokens_per_bs if self.use_mla and _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch self.make_mla_meta_data( @@ -1979,7 +2248,6 @@ class AiterAttnBackend(AttentionBackend): max_q_len = num_tokens_per_bs if self.use_mla and _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch self.make_mla_meta_data( @@ -2072,7 +2340,6 @@ class AiterAttnBackend(AttentionBackend): self.use_triton_unified_attention and self.use_sliding_window_kv_pool ): - token_to_kv_pool = forward_batch.token_to_kv_pool k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( layer.layer_id @@ -2275,7 +2542,6 @@ class AiterAttnBackend(AttentionBackend): forward_batch.forward_mode.is_draft_extend() or forward_batch.forward_mode.is_draft_extend_v2() ): - work_metadata = self.forward_metadata.work_metadata work_indptr = self.forward_metadata.work_indptr work_info_set = self.forward_metadata.work_info_set @@ -2351,7 +2617,6 @@ class AiterAttnBackend(AttentionBackend): forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend() ): - # Use triton extend kernel which supports custom masks and causal masking if layer.qk_head_dim != layer.v_head_dim: o = q.new_empty( (q.shape[0], layer.tp_q_head_num * layer.v_head_dim) @@ -2359,6 +2624,67 @@ class AiterAttnBackend(AttentionBackend): else: o = torch.empty_like(q) + # target_verify goes through unified_attention when topk == 1 + # (the linear draft chain gives a pure causal mask). MLA and + # draft_extend still use the legacy extend_attention_fwd path. + if ( + self._use_unified_verify + and forward_batch.forward_mode.is_target_verify() + ): + k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( + layer.layer_id + ) + page_table = self.forward_metadata.kv_indices + max_kv_len = page_table.shape[1] * self.page_size + + window_size = (-1, -1) + + if ( + layer.sliding_window_size is not None + and layer.sliding_window_size > -1 + ): + window_size = (layer.sliding_window_size - 1, 0) + if self.forward_metadata.swa_page_table is not None: + page_table = self.forward_metadata.swa_page_table + + q_unified = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim) + k_unified = k_cache.view( + -1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim + ) + v_unified = v_cache.view( + -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim + ) + if layer.tp_k_head_num == 1 and layer.tp_q_head_num > 1: + # Qwen3.5 can replicate one KV head across multiple TP ranks. + # Present the local KV head as per-Q-head stride-0 views so + # target_verify uses the same local head mapping as the model. + k_unified = k_unified.expand(-1, -1, layer.tp_q_head_num, -1) + v_unified = v_unified.expand(-1, -1, layer.tp_q_head_num, -1) + + # The seq_lens + draft_num add has to run INSIDE the graph + # region; a host-side pre-add would allocate a new tensor + # each replay and break the captured pointer. + unified_attention( + q=q_unified, + k=k_unified, + v=v_unified, + out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim), + cu_seqlens_q=self.forward_metadata.qo_indptr, + seqused_k=forward_batch.seq_lens + self.num_draft_tokens, + max_seqlen_q=self.forward_metadata.max_q_len, + max_seqlen_k=max_kv_len, + softmax_scale=layer.scaling, + causal=True, + window_size=window_size, + block_table=page_table, + softcap=layer.logit_cap, + q_descale=None, + k_descale=k_descale, + v_descale=v_descale, + sinks=sinks, + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + self.extend_attention_fwd( q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), k.contiguous(), @@ -2457,7 +2783,6 @@ class AiterAttnBackend(AttentionBackend): # use standard set_kv_buffer, as they lack SWA-specific attributes # like full_to_swa_index_mapping. if self.use_triton_unified_attention and self.use_sliding_window_kv_pool: - token_to_kv_pool = forward_batch.token_to_kv_pool k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( layer.layer_id @@ -2542,10 +2867,15 @@ class AiterAttnBackend(AttentionBackend): layer.layer_id ) - o = torch.empty_like(q, dtype=self.input_dtype) + if layer.qk_head_dim != layer.v_head_dim: + o = q.new_empty( + (q.shape[0], layer.tp_q_head_num * layer.v_head_dim), + dtype=self.input_dtype, + ) + else: + o = torch.empty_like(q, dtype=self.input_dtype) if self.use_triton_unified_attention: - bs = forward_batch.batch_size window_size = (-1, -1) page_table = self.forward_metadata.kv_indices @@ -2568,7 +2898,7 @@ class AiterAttnBackend(AttentionBackend): v=v_cache.view( -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim ), - out=o.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim), cu_seqlens_q=self.forward_metadata.qo_indptr, seqused_k=forward_batch.seq_lens, max_seqlen_q=self.forward_metadata.max_q_len, @@ -2589,7 +2919,7 @@ class AiterAttnBackend(AttentionBackend): v_cache = v_cache.to(self.input_dtype) paged_attention_ragged( - o.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + o.view(-1, layer.tp_q_head_num, layer.v_head_dim), self.workspace_buffer, q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim), diff --git a/python/sglang/srt/layers/attention/triton_ops/aiter_unified_attention.py b/python/sglang/srt/layers/attention/triton_ops/aiter_unified_attention.py new file mode 100644 index 000000000..ed790a483 --- /dev/null +++ b/python/sglang/srt/layers/attention/triton_ops/aiter_unified_attention.py @@ -0,0 +1,97 @@ +import triton +import triton.language as tl + + +@triton.jit +def scatter_ragged_to_page_table_kernel( + kv_flat_ptr, + kv_indptr_ptr, + dest_ptr, + dest_stride, + sw_page_table_ptr, + swa_slot_mapping_ptr, + PAGE_SIZE: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + HAS_SWA: tl.constexpr, +): + """Scatter ragged token-level kv_indices into a 2D block-level page table.""" + pid = tl.program_id(0) + block_id = tl.program_id(1) + + start = tl.load(kv_indptr_ptr + pid).to(tl.int64) + kv_len = tl.load(kv_indptr_ptr + pid + 1).to(tl.int64) - start + num_blocks = (kv_len + PAGE_SIZE - 1) // PAGE_SIZE + + offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + if block_id * BLOCK_SIZE >= num_blocks: + return + mask = offsets < num_blocks + token_idx = offsets.to(tl.int64) * PAGE_SIZE + vals = tl.load(kv_flat_ptr + start + token_idx, mask=mask, other=0) + block_vals = vals // PAGE_SIZE + tl.store( + dest_ptr + pid.to(tl.int64) * dest_stride + offsets, + block_vals, + mask=mask, + ) + + if HAS_SWA: + sw_vals = tl.load(swa_slot_mapping_ptr + vals) + block_vals = sw_vals // PAGE_SIZE + tl.store( + sw_page_table_ptr + pid.to(tl.int64) * dest_stride + offsets, + block_vals, + mask=mask, + ) + + +@triton.jit +def scatter_req_to_token_to_page_table_kernel( + req_to_token_ptr, + req_pool_indices_ptr, + seq_lens_ptr, + page_table_ptr, + req_to_token_stride, + page_table_stride, + sw_page_table_ptr, + swa_slot_mapping_ptr, + DRAFT_NUM: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + HAS_SWA: tl.constexpr, +): + """Build the 2D block-level page_table for target_verify from req_to_token.""" + pid = tl.program_id(0) + block_id = tl.program_id(1) + + seq_len = tl.load(seq_lens_ptr + pid).to(tl.int64) + kv_len = seq_len + DRAFT_NUM + num_blocks = (kv_len + PAGE_SIZE - 1) // PAGE_SIZE + + offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + if block_id * BLOCK_SIZE >= num_blocks: + return + mask = offsets < num_blocks + + rp = tl.load(req_pool_indices_ptr + pid).to(tl.int64) + token_idx = offsets.to(tl.int64) * PAGE_SIZE + vals = tl.load( + req_to_token_ptr + rp * req_to_token_stride + token_idx, + mask=mask, + other=0, + ) + block_vals = vals // PAGE_SIZE + tl.store( + page_table_ptr + pid.to(tl.int64) * page_table_stride + offsets, + block_vals, + mask=mask, + ) + + if HAS_SWA: + sw_vals = tl.load(swa_slot_mapping_ptr + vals) + block_vals = sw_vals // PAGE_SIZE + tl.store( + sw_page_table_ptr + pid.to(tl.int64) * page_table_stride + offsets, + block_vals, + mask=mask, + ) diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index c60737c56..a548900af 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -62,6 +62,18 @@ class Qwen3_5ForCausalLMMTP(nn.Module): ): quant_config = None + # Quark-quantized Qwen3.5 MXFP4 checkpoints ship the MTP module in + # bf16; every `mtp.*` layer appears under the quantization exclude + # list. Detect that and skip quantization here so linear/MoE weight + # loaders allocate bf16 shapes (see sgl-project/sglang#23113). + if quant_config and quant_config.get_name() == "quark": + exclude_layers = getattr(quant_config, "exclude_layers", []) + if any( + isinstance(layer, str) and layer.startswith("mtp.") + for layer in exclude_layers + ): + quant_config = None + self.config = config self.tp_size = get_tensor_model_parallel_world_size() self.quant_config = quant_config diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 9d46514c5..67d50dc9b 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -131,6 +131,7 @@ QUANTIZATION_CHOICES = [ "auto-round", "compressed-tensors", # for Ktransformers "modelslim", # for NPU + "quark", # AMD Quark quantizer (FP8 / MXFP4 / Int4FP8 etc.) "quark_int4fp8_moe", "unquant", ]