[AMD] Enable EAGLE speculative decoding for Qwen3.5 FP8 and MXFP4 models with aiter's unified attention (#23146)

Co-authored-by: wunhuang <wunhuang@amd.com>
Co-authored-by: sogalin <39478626+sogalin@users.noreply.github.com>
This commit is contained in:
Hubert Lu
2026-05-05 00:09:40 -07:00
committed by GitHub
co-authored by wunhuang sogalin
parent 244531bc4f
commit c2db19ffa4
4 changed files with 592 additions and 152 deletions
@@ -13,6 +13,10 @@ import torch
import triton import triton
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend 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 ( from sglang.srt.layers.attention.utils import (
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
create_flashmla_kv_indices_triton, create_flashmla_kv_indices_triton,
@@ -106,6 +110,7 @@ class ForwardMetadata:
global_workspace_buffer = None global_workspace_buffer = None
_AITER_PARTITION_SIZE_ROCM = 256 _AITER_PARTITION_SIZE_ROCM = 256
@@ -179,6 +184,11 @@ class AiterAttnBackend(AttentionBackend):
self.qo_indptr = torch.zeros( self.qo_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=model_runner.device (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( self.mask_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int64, device=model_runner.device (max_bs + 1,), dtype=torch.int64, device=model_runner.device
) )
@@ -208,6 +218,17 @@ class AiterAttnBackend(AttentionBackend):
"SGLANG_USE_AITER_UNIFIED_ATTN" "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 # aiter kernel related initialization
self.max_num_partitions = ( self.max_num_partitions = (
self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1 self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1
@@ -485,6 +506,128 @@ class AiterAttnBackend(AttentionBackend):
) )
return page_table[:, strided_indices] // page_size 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( def _resolve_v2_num_draft_tokens(
self, self,
extend_seq_lens: Optional[torch.Tensor] = None, extend_seq_lens: Optional[torch.Tensor] = None,
@@ -729,11 +872,17 @@ class AiterAttnBackend(AttentionBackend):
elif self.page_size > 1: elif self.page_size > 1:
kv_indices = self._transform_table_1_to_real(kv_indices) kv_indices = self._transform_table_1_to_real(kv_indices)
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr_unified_decode[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(
self.kv_last_page_len[:bs], dim=0
)
else:
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: else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
bs = kv_indptr.shape[0] - 1 bs = kv_indptr.shape[0] - 1
@@ -1038,10 +1187,30 @@ class AiterAttnBackend(AttentionBackend):
run_graph=False, run_graph=False,
) )
else: else:
# Non-MLA target_verify: use triton extend kernel with custom mask
bs = len(forward_batch.req_pool_indices) bs = len(forward_batch.req_pool_indices)
draft_num = spec_info.draft_token_num draft_num = spec_info.draft_token_num
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( qo_indptr = torch.arange(
0, 0,
(1 + bs) * draft_num, (1 + bs) * draft_num,
@@ -1299,7 +1468,12 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = None kv_last_page_len = None
max_q_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: if not self.use_triton_unified_attention:
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
@@ -1317,14 +1491,24 @@ class AiterAttnBackend(AttentionBackend):
) )
else: else:
max_q_len = 1 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( kv_indices = self.cuda_graph_kv_indices.view(
-1, max_num_blocks_per_seq -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_table = self.cuda_graph_swa_page_table
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: if self.use_sliding_window_kv_pool:
swa_page_indices = ( swa_page_indices = (
@@ -1350,10 +1534,7 @@ class AiterAttnBackend(AttentionBackend):
new_cols = page_indices.shape[1] new_cols = page_indices.shape[1]
kv_indices[:new_rows, :new_cols].copy_(page_indices) kv_indices[:new_rows, :new_cols].copy_(page_indices)
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr_unified_decode[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(
self.cuda_graph_kv_last_page_len[:bs], dim=0
)
kv_indptr = None kv_indptr = None
else: else:
@@ -1441,7 +1622,6 @@ class AiterAttnBackend(AttentionBackend):
if self.use_mla: if self.use_mla:
if _use_mla_ps_kernel: if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -1483,9 +1663,48 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map=reduce_partial_map, reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits, num_kv_splits=num_kv_splits,
) )
else:
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]
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: else:
custom_mask = self.cuda_graph_custom_mask custom_mask = self.cuda_graph_custom_mask
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.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) seq_mask_len = max_q_len * (seq_lens + max_q_len)
mask_indptr = self.mask_indptr mask_indptr = self.mask_indptr
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0)
@@ -1593,7 +1812,6 @@ class AiterAttnBackend(AttentionBackend):
max_q_len = num_tokens_per_bs max_q_len = num_tokens_per_bs
if _use_mla_ps_kernel: if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -1682,7 +1900,13 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = None kv_last_page_len = None
max_q_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: if not self.use_triton_unified_attention:
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
@@ -1699,14 +1923,24 @@ class AiterAttnBackend(AttentionBackend):
) )
else: else:
max_q_len = 1 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( kv_indices = self.cuda_graph_kv_indices.view(
-1, max_num_blocks_per_seq -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_table = self.cuda_graph_swa_page_table
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: if self.use_sliding_window_kv_pool:
swa_page_indices = ( swa_page_indices = (
@@ -1732,10 +1966,7 @@ class AiterAttnBackend(AttentionBackend):
new_cols = page_indices.shape[1] new_cols = page_indices.shape[1]
kv_indices[:new_rows, :new_cols].copy_(page_indices) kv_indices[:new_rows, :new_cols].copy_(page_indices)
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr_unified_decode[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(
self.cuda_graph_kv_last_page_len[:bs], dim=0
)
kv_indptr = None kv_indptr = None
else: else:
@@ -1825,7 +2056,6 @@ class AiterAttnBackend(AttentionBackend):
if self.use_mla: if self.use_mla:
if _use_mla_ps_kernel: if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -1867,9 +2097,49 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map=reduce_partial_map, reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits, num_kv_splits=num_kv_splits,
) )
else:
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]
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: else:
custom_mask = self.cuda_graph_custom_mask custom_mask = self.cuda_graph_custom_mask
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.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) seq_mask_len = max_q_len * (seq_lens + max_q_len)
mask_indptr = self.mask_indptr[: bs + 1] mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
@@ -1913,7 +2183,6 @@ class AiterAttnBackend(AttentionBackend):
max_q_len = num_tokens_per_bs max_q_len = num_tokens_per_bs
if self.use_mla and _use_mla_ps_kernel: if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -1979,7 +2248,6 @@ class AiterAttnBackend(AttentionBackend):
max_q_len = num_tokens_per_bs max_q_len = num_tokens_per_bs
if self.use_mla and _use_mla_ps_kernel: if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -2072,7 +2340,6 @@ class AiterAttnBackend(AttentionBackend):
self.use_triton_unified_attention self.use_triton_unified_attention
and self.use_sliding_window_kv_pool and self.use_sliding_window_kv_pool
): ):
token_to_kv_pool = forward_batch.token_to_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( k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
layer.layer_id layer.layer_id
@@ -2275,7 +2542,6 @@ class AiterAttnBackend(AttentionBackend):
forward_batch.forward_mode.is_draft_extend() forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2() or forward_batch.forward_mode.is_draft_extend_v2()
): ):
work_metadata = self.forward_metadata.work_metadata work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set work_info_set = self.forward_metadata.work_info_set
@@ -2351,7 +2617,6 @@ class AiterAttnBackend(AttentionBackend):
forward_batch.forward_mode.is_target_verify() forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend() 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: if layer.qk_head_dim != layer.v_head_dim:
o = q.new_empty( o = q.new_empty(
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim) (q.shape[0], layer.tp_q_head_num * layer.v_head_dim)
@@ -2359,6 +2624,67 @@ class AiterAttnBackend(AttentionBackend):
else: else:
o = torch.empty_like(q) 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( self.extend_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k.contiguous(), k.contiguous(),
@@ -2457,7 +2783,6 @@ class AiterAttnBackend(AttentionBackend):
# use standard set_kv_buffer, as they lack SWA-specific attributes # use standard set_kv_buffer, as they lack SWA-specific attributes
# like full_to_swa_index_mapping. # like full_to_swa_index_mapping.
if self.use_triton_unified_attention and self.use_sliding_window_kv_pool: if self.use_triton_unified_attention and self.use_sliding_window_kv_pool:
token_to_kv_pool = forward_batch.token_to_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( k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
layer.layer_id layer.layer_id
@@ -2542,10 +2867,15 @@ class AiterAttnBackend(AttentionBackend):
layer.layer_id layer.layer_id
) )
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) o = torch.empty_like(q, dtype=self.input_dtype)
if self.use_triton_unified_attention: if self.use_triton_unified_attention:
bs = forward_batch.batch_size bs = forward_batch.batch_size
window_size = (-1, -1) window_size = (-1, -1)
page_table = self.forward_metadata.kv_indices page_table = self.forward_metadata.kv_indices
@@ -2568,7 +2898,7 @@ class AiterAttnBackend(AttentionBackend):
v=v_cache.view( v=v_cache.view(
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim -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, cu_seqlens_q=self.forward_metadata.qo_indptr,
seqused_k=forward_batch.seq_lens, seqused_k=forward_batch.seq_lens,
max_seqlen_q=self.forward_metadata.max_q_len, max_seqlen_q=self.forward_metadata.max_q_len,
@@ -2589,7 +2919,7 @@ class AiterAttnBackend(AttentionBackend):
v_cache = v_cache.to(self.input_dtype) v_cache = v_cache.to(self.input_dtype)
paged_attention_ragged( 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, self.workspace_buffer,
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), 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), k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim),
@@ -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,
)
+12
View File
@@ -62,6 +62,18 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
): ):
quant_config = None 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.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_tensor_model_parallel_world_size()
self.quant_config = quant_config self.quant_config = quant_config
+1
View File
@@ -131,6 +131,7 @@ QUANTIZATION_CHOICES = [
"auto-round", "auto-round",
"compressed-tensors", # for Ktransformers "compressed-tensors", # for Ktransformers
"modelslim", # for NPU "modelslim", # for NPU
"quark", # AMD Quark quantizer (FP8 / MXFP4 / Int4FP8 etc.)
"quark_int4fp8_moe", "quark_int4fp8_moe",
"unquant", "unquant",
] ]