[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
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),
@@ -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
# 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
+1
View File
@@ -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",
]