[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:
co-authored by
wunhuang
sogalin
parent
244531bc4f
commit
c2db19ffa4
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user