[Ascend] Support enable-mixed-chunk in non-MLA scenarios (#12491)
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
@@ -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)
|
||||
mask_flag = ~mask_flag
|
||||
if dtype == torch.float16:
|
||||
mask_value = torch.finfo(torch.float32).min
|
||||
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)
|
||||
: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
|
||||
)
|
||||
self.mask_len = max_seq_len
|
||||
|
||||
def get_verify_buffers_to_fill_after_draft(self):
|
||||
# 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):
|
||||
"""
|
||||
Return buffers for verify attention kernels that needs to be filled after draft.
|
||||
Generate a mask flag for attention masks.
|
||||
|
||||
Typically, these are tree mask and position buffers.
|
||||
:param max_seq_len: Maximum sequence length for the mask.
|
||||
:return: A boolean tensor representing the mask flag.
|
||||
"""
|
||||
return [None, None]
|
||||
# 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
|
||||
return mask_flag
|
||||
|
||||
def update_verify_buffers_to_fill_after_draft(
|
||||
self, spec_info: SpecInput, cuda_graph_bs: Optional[int]
|
||||
@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 = 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
|
||||
|
||||
@staticmethod
|
||||
def get_attention_mask_id(seq_lens, extend_lens):
|
||||
"""
|
||||
Generate attention mask IDs based on sequence lengths and extended lengths.
|
||||
|
||||
:param seq_lens: Sequence lengths.
|
||||
:param extend_lens: Extended lengths.
|
||||
:return: A tensor containing the attention mask IDs.
|
||||
"""
|
||||
starts = seq_lens - extend_lens
|
||||
ends = seq_lens
|
||||
|
||||
# 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
|
||||
|
||||
Reference in New Issue
Block a user