[Ascend] Support enable-mixed-chunk in non-MLA scenarios (#12491)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Michelle Wu
2025-11-27 00:00:21 +08:00
committed by GitHub
co-authored by ronnie_zheng
parent 6c190cbda0
commit 262c3c1fde
3 changed files with 262 additions and 38 deletions
@@ -20,9 +20,12 @@ if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
import logging
import numpy as np
logger = logging.getLogger(__name__)
@dataclass
class ForwardMetadata:
@@ -39,38 +42,166 @@ class ForwardMetadata:
actual_seq_lengths_q: Optional[torch.Tensor] = None
class AscendAttnBackend(AttentionBackend):
class AscendAttnMaskBuilder:
def __init__(self, model_runner: ModelRunner, device, use_fia):
"""
Initialize the AscendAttnMaskBuilder class.
def gen_attention_mask(self, max_seq_len: int, dtype=torch.float16):
mask_flag = torch.tril(
torch.ones((max_seq_len, max_seq_len), dtype=torch.bool)
).view(max_seq_len, max_seq_len)
:param model_runner: ModelRunner instance for model execution.
:param device: Device to run the model on (e.g., 'cuda', 'npu').
:param use_fia: Boolean flag to indicate if environment variable ASCEND_USE_FIA is set to 1.
"""
self.use_fia = use_fia
self.model_runner = model_runner
self.device = device
# Initialize mask
mask_len = 128
self.mask = self.generate_attn_mask(mask_len, "norm", model_runner.dtype).to(
self.device
)
# Initialize FIA mask
fia_mask_len = 2048
self.fia_mask = self.generate_mask_flag(fia_mask_len).to(self.device)
# Initialize MTP mask
mtp_mask_len = 2048
self.mtp_mask = self.generate_mask_flag(mtp_mask_len).to(self.device)
# Initialize mixed chunk mask cache
mixed_chunk_cache_len = 8192
self.mix_mask_cache = self.generate_attn_mask(mixed_chunk_cache_len, "mix")
self.mix_seq_len_cached = self.mix_mask_cache.shape[0]
@staticmethod
def generate_mask_flag(max_seq_len):
"""
Generate a mask flag for attention masks.
:param max_seq_len: Maximum sequence length for the mask.
:return: A boolean tensor representing the mask flag.
"""
# Construct lower triangle matrix.
mask_flag = torch.ones((max_seq_len, max_seq_len), dtype=torch.bool).tril_()
# Create upper triangle matrix used to mark mask positions.
mask_flag = ~mask_flag
if dtype == torch.float16:
mask_value = torch.finfo(torch.float32).min
return mask_flag
@staticmethod
def generate_attn_mask(max_seq_len, mode, dtype=torch.float16):
"""
Generate an attention mask.
:param max_seq_len: Maximum sequence length for the mask.
:param mode: Mode of the mask ('mix' or 'norm').
:param dtype: Data type of the mask tensor.
:return: A tensor representing the attention mask.
"""
mask_flag = AscendAttnMaskBuilder.generate_mask_flag(max_seq_len)
if mode == "mix":
mask_value = (
float("-inf") if dtype in [torch.float16, torch.bfloat16] else 1
)
else:
mask_value = 1
self.mask = (
torch.masked_fill(
torch.zeros(size=(max_seq_len, max_seq_len)), mask_flag, mask_value
)
.to(dtype)
.to(self.device)
)
self.mask_len = max_seq_len
mask_value = torch.finfo(torch.float32).min if dtype == torch.float16 else 1
attn_mask = torch.zeros(
size=(max_seq_len, max_seq_len), dtype=dtype
).masked_fill_(mask_flag, mask_value)
return attn_mask
def get_verify_buffers_to_fill_after_draft(self):
@staticmethod
def get_attention_mask_id(seq_lens, extend_lens):
"""
Return buffers for verify attention kernels that needs to be filled after draft.
Generate attention mask IDs based on sequence lengths and extended lengths.
Typically, these are tree mask and position buffers.
:param seq_lens: Sequence lengths.
:param extend_lens: Extended lengths.
:return: A tensor containing the attention mask IDs.
"""
return [None, None]
starts = seq_lens - extend_lens
ends = seq_lens
def update_verify_buffers_to_fill_after_draft(
self, spec_info: SpecInput, cuda_graph_bs: Optional[int]
# Use torch.stack to stack the start and end indices together
ranges = torch.stack((starts, ends), dim=-1)
# Use list comprehension to generate tensors for each range and concatenate them
attn_mask_id = torch.cat([torch.arange(start, end) for start, end in ranges])
return attn_mask_id
def update_attn_cache(
self,
seqlen: int,
mask_cache: torch.Tensor,
seq_len_cached: int,
dtype: torch.dtype,
mode,
):
pass
"""
Update the attention mask cache.
:param seqlen: Maximum sequence length.
:param mask_cache: Current attention mask cache.
:param seq_len_cached: Cached sequence length.
:param dtype: Data type of the mask tensor.
:param mode: Mode of the mask ('mix' or 'norm').
:return: Updated mask cache and sequence length cache.
"""
if seqlen > seq_len_cached:
seq_len_cached = seqlen
mask_cache = self.generate_attn_mask(seqlen, mode, dtype)
if mask_cache.dtype != dtype:
mask_cache = mask_cache.to(dtype)
return mask_cache, seq_len_cached
def get_splitfuse_attn_mask(
self,
seq_lens: torch.Tensor = None,
position: torch.Tensor = None,
dtype: torch.dtype = None,
device: torch.device = None,
) -> torch.Tensor:
"""
Generate a splitfuse attention mask.
:param seq_lens: Sequence lengths.
:param position: Position indices for the mask.
:param dtype: Data type of the mask tensor.
:param device: Device to run the model on.
:return: A tensor representing the splitfuse attention mask.
"""
if dtype not in [torch.float16, torch.bfloat16]:
raise ValueError("splitfuse_attn_mask now only supports bf16 and fp16")
max_seq_len = max(seq_lens, default=0)
self.mix_mask_cache, self.mix_seq_len_cached = self.update_attn_cache(
max_seq_len, self.mix_mask_cache, self.mix_seq_len_cached, dtype, mode="mix"
)
attn_mask = torch.index_select(self.mix_mask_cache, dim=0, index=position)[
:, :max_seq_len
]
return attn_mask.contiguous().to(device, non_blocking=True)
def update_mask(self, forward_metadata):
"""
Update the splitfuse attention mask based on forward metadata.
:param forward_metadata: Forward metadata containing sequence lengths and extended lengths.
:return: Updated splitfuse attention mask.
"""
attn_mask_id = self.get_attention_mask_id(
forward_metadata.seq_lens_cpu_int,
forward_metadata.extend_seq_lens_cpu_int,
)
mix_mask = self.get_splitfuse_attn_mask(
seq_lens=forward_metadata.seq_lens_cpu_int,
position=attn_mask_id,
dtype=torch.float16,
device=self.device,
).to(torch.bfloat16)
return mix_mask
class AscendAttnBackend(AttentionBackend):
def __init__(self, model_runner: ModelRunner):
super().__init__()
@@ -90,21 +221,31 @@ class AscendAttnBackend(AttentionBackend):
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.graph_mode = False
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
if not self.use_fia:
self.gen_attention_mask(128, model_runner.dtype)
mask_length = 2048
self.fia_mask = ~torch.tril(
torch.ones(
(mask_length, mask_length),
dtype=torch.bool,
device=model_runner.device,
)
)
self.speculative_num_draft_tokens = (
model_runner.server_args.speculative_num_draft_tokens
)
self.mtp_mask = torch.tril(torch.ones(2048, 2048, dtype=torch.bool)).npu()
self.mtp_mask = ~self.mtp_mask
self.ascend_attn_mask_builder = AscendAttnMaskBuilder(
model_runner, self.device, self.use_fia
)
self.mask, self.fia_mask, self.mtp_mask, self.mix_mask = (
self.ascend_attn_mask_builder.mask,
self.ascend_attn_mask_builder.fia_mask,
self.ascend_attn_mask_builder.mtp_mask,
self.ascend_attn_mask_builder.mix_mask_cache,
)
def get_verify_buffers_to_fill_after_draft(self):
"""
Return buffers for verify attention kernels that needs to be filled after draft.
Typically, these are tree mask and position buffers.
"""
return [None, None]
def update_verify_buffers_to_fill_after_draft(
self, spec_info: SpecInput, cuda_graph_bs: Optional[int]
):
pass
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass."""
@@ -134,7 +275,10 @@ class AscendAttnBackend(AttentionBackend):
if forward_batch.forward_mode.is_target_verify():
self.forward_metadata.seq_lens_cpu_int += self.speculative_num_draft_tokens
if forward_batch.forward_mode.is_mixed():
self.mix_mask = self.ascend_attn_mask_builder.update_mask(
self.forward_metadata
)
self.graph_mode = False
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
@@ -939,6 +1083,62 @@ class AscendAttnBackend(AttentionBackend):
)
return attn_output.view(num_tokens, layer.tp_q_head_num * self.kv_lora_rank)
def forward_mixed(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
topk_indices: Optional[torch.Tensor] = None,
):
if (
topk_indices is not None
or self.use_mla
or (not self.use_fia and layer.qk_head_dim > 128)
):
raise NotImplementedError(
"The 'enable-mixed-chunk' feature is currently unsupported in the following scenarios: "
"1. When using the MLA backend on Ascend NPU devices, "
"2. When using the deepseekv3.2 model on Ascend NPU devices, "
"3. When the environment variable ASCEND_USE_FIA is set to 0 and qk_head_dim exceeds 128 on Ascend NPU devices."
)
if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
)
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
# Initialize the output tensor for attention results
attn_output = torch.empty(
(query.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=query.dtype,
device=query.device,
)
torch_npu._npu_paged_attention_splitfuse(
query=query,
key_cache=k_cache,
value_cache=v_cache,
block_table=self.forward_metadata.block_tables,
context_lens=self.forward_metadata.seq_lens_cpu_int,
mask=self.mix_mask,
seq_len=self.forward_metadata.extend_seq_lens_cpu_int,
scale_value=layer.scaling,
num_heads=layer.tp_q_head_num,
num_kv_heads=layer.tp_k_head_num,
out=attn_output,
)
return attn_output.view(
attn_output.shape[0], layer.tp_q_head_num * layer.v_head_dim
)
class AscendAttnMultiStepDraftBackend:
"""
@@ -5,6 +5,8 @@ from typing import TYPE_CHECKING, Optional
import torch
from sglang.srt.utils.common import is_npu
if TYPE_CHECKING:
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.radix_attention import RadixAttention
@@ -97,6 +99,16 @@ class AttentionBackend(ABC):
save_kv_cache=save_kv_cache,
**kwargs,
)
elif forward_batch.forward_mode.is_mixed() and is_npu():
return self.forward_mixed(
q,
k,
v,
layer,
forward_batch,
save_kv_cache=save_kv_cache,
**kwargs,
)
else:
return self.forward_extend(
q,
@@ -132,6 +144,18 @@ class AttentionBackend(ABC):
"""Run a forward for extend."""
raise NotImplementedError()
def forward_mixed(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
):
"""Run a forward for mix."""
raise NotImplementedError()
def support_triton(self):
"""Check if the current backend supports triton."""
return True
@@ -553,7 +553,7 @@ class PrefillAdder:
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
)
else:
if self.rem_chunk_tokens == 0:
if self.rem_chunk_tokens <= 0:
return AddReqResult.OTHER
# Chunked prefill