[NPU] Add Gemma4 Sliding Window Attention support on Ascend backend (#26147)
This commit is contained in:
@@ -38,6 +38,9 @@ import logging
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
FULL_ATTENTION_WINDOW = 2147483647
|
||||
|
||||
|
||||
def _reshape_kv_for_fia_nz(
|
||||
tensor: torch.Tensor, num_heads: int, head_dim: int, page_size: int
|
||||
@@ -46,12 +49,6 @@ def _reshape_kv_for_fia_nz(
|
||||
return tensor.view(-1, 1, num_heads * head_dim // 16, page_size, 16)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# default max value of full attention window size
|
||||
FULL_ATTENTION_WINDOW = 2147483647
|
||||
|
||||
|
||||
@dataclass
|
||||
class ForwardMetadata:
|
||||
|
||||
@@ -73,6 +70,9 @@ class ForwardMetadata:
|
||||
actual_seq_lengths_q: Optional[torch.Tensor] = None
|
||||
actual_seq_lengths_kv: Optional[torch.Tensor] = None
|
||||
|
||||
# swa attention mask for graph mode decode
|
||||
swa_mask: Optional[torch.Tensor] = None
|
||||
|
||||
# prefix cache
|
||||
prefix_lens: Optional[torch.Tensor] = None
|
||||
flatten_prefix_block_tables: Optional[torch.Tensor] = None
|
||||
@@ -337,6 +337,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.full_to_swa_index_mapping = (
|
||||
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
||||
)
|
||||
self.sliding_window_size = model_runner.sliding_window_size
|
||||
self.use_sliding_window_kv_pool = (
|
||||
isinstance(self.token_to_kv_pool, SWAKVPool)
|
||||
and self.token_to_kv_pool.swa_layer_nums > 0
|
||||
@@ -363,6 +364,20 @@ class AscendAttnBackend(AttentionBackend):
|
||||
|
||||
self.attn_cp_size = model_runner.attn_cp_size
|
||||
|
||||
def _is_swa_layer(self, layer: RadixAttention) -> bool:
|
||||
return (
|
||||
self.is_hybrid_swa
|
||||
and layer.sliding_window_size is not None
|
||||
and layer.sliding_window_size > -1
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _can_use_tnd(layer: RadixAttention) -> bool:
|
||||
"""Check if TND layout is supported."""
|
||||
d = layer.qk_head_dim
|
||||
v = layer.v_head_dim
|
||||
return (d == v and d in (128, 192)) or (d == 192 and v == 128)
|
||||
|
||||
def get_verify_buffers_to_fill_after_draft(self):
|
||||
"""
|
||||
Return buffers for verify attention kernels that needs to be filled after draft.
|
||||
@@ -516,6 +531,19 @@ class AscendAttnBackend(AttentionBackend):
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
# SWA mask: True = masked out (don't attend), False = attend.
|
||||
# Pre-allocated at max size, sliced per batch size during capture,
|
||||
# content updated via copy_() during replay.
|
||||
self.graph_metadata["swa_mask"] = torch.ones(
|
||||
(max_bs, 1, total_context_len),
|
||||
dtype=torch.bool,
|
||||
device=self.device,
|
||||
)
|
||||
# Pre-allocated index buffer for mask generation during replay,
|
||||
# avoids torch.arange allocation on every replay step.
|
||||
self.graph_metadata["swa_indices"] = torch.arange(
|
||||
total_context_len, device=self.device, dtype=torch.int32
|
||||
)
|
||||
if self.use_sliding_window_kv_pool:
|
||||
# refilled in place at replay; the captured graph reads this storage
|
||||
self.swa_out_cache_loc_buf = torch.zeros(
|
||||
@@ -536,6 +564,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
|
||||
if self.is_hybrid_swa:
|
||||
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
|
||||
metadata.swa_mask = self.graph_metadata["swa_mask"][:bs, :, :]
|
||||
if self.use_sliding_window_kv_pool and out_cache_loc is not None:
|
||||
num_tokens = out_cache_loc.shape[0]
|
||||
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
|
||||
@@ -636,6 +665,18 @@ class AscendAttnBackend(AttentionBackend):
|
||||
)
|
||||
metadata.block_tables_swa[:bs, max_seq_pages:].fill_(0)
|
||||
metadata.block_tables_swa[bs:, :].fill_(0)
|
||||
|
||||
# Update SWA mask: True = masked out (don't attend), False = attend
|
||||
seq_lens_int = seq_lens_cpu[:bs].int()
|
||||
starts = torch.clamp(seq_lens_int - self.sliding_window_size, min=0)
|
||||
indices = self.graph_metadata["swa_indices"]
|
||||
start_exp = starts.unsqueeze(1).to(self.device)
|
||||
seq_exp = seq_lens_int.unsqueeze(1).to(self.device)
|
||||
mask = (indices.unsqueeze(0) < start_exp) | (
|
||||
indices.unsqueeze(0) >= seq_exp
|
||||
)
|
||||
metadata.swa_mask[:bs, 0, :].copy_(mask)
|
||||
metadata.swa_mask[bs:, :, :].fill_(True)
|
||||
metadata.block_tables[:bs, :max_seq_pages].copy_(
|
||||
self.req_to_token[req_pool_indices[:bs], :max_len][:, :: self.page_size]
|
||||
// self.page_size
|
||||
@@ -1132,65 +1173,119 @@ class AscendAttnBackend(AttentionBackend):
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
|
||||
if sinks is not None or (
|
||||
self.is_hybrid_swa and layer.sliding_window_size != -1
|
||||
):
|
||||
if sinks is not None or (self._is_swa_layer(layer) and self.use_fia):
|
||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||
if self._is_swa_layer(layer):
|
||||
block_tables = self.forward_metadata.block_tables_swa
|
||||
else:
|
||||
block_tables = self.forward_metadata.block_tables
|
||||
if self.use_fia:
|
||||
num_token_padding = q.shape[0]
|
||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||
q, k, v = [
|
||||
data[: forward_batch.num_token_non_padded_cpu]
|
||||
for data in [q, k, v]
|
||||
]
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
block_size = self.page_size
|
||||
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||
query=q,
|
||||
key=k_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
),
|
||||
value=v_cache.view(
|
||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||
),
|
||||
pre_tokens=(
|
||||
layer.sliding_window_size
|
||||
if layer.sliding_window_size != -1
|
||||
else FULL_ATTENTION_WINDOW
|
||||
),
|
||||
next_tokens=(
|
||||
0
|
||||
if layer.sliding_window_size != -1
|
||||
else FULL_ATTENTION_WINDOW
|
||||
),
|
||||
atten_mask=self.fia_mask,
|
||||
block_table=block_tables,
|
||||
input_layout="TND",
|
||||
block_size=block_size,
|
||||
num_query_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
actual_seq_qlen=self.forward_metadata.seq_lens_list_cumsum,
|
||||
actual_seq_kvlen=self.forward_metadata.seq_lens_cpu_int,
|
||||
softmax_scale=layer.scaling,
|
||||
sparse_mode=4 if layer.sliding_window_size != -1 else 3,
|
||||
learnable_sink=sinks,
|
||||
)
|
||||
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||
attn_out = torch.cat(
|
||||
[
|
||||
attn_out,
|
||||
attn_out.new_zeros(
|
||||
num_token_padding - attn_out.shape[0],
|
||||
*attn_out.shape[1:],
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
if self._can_use_tnd(layer):
|
||||
num_token_padding = q.shape[0]
|
||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||
q, k, v = [
|
||||
data[: forward_batch.num_token_non_padded_cpu]
|
||||
for data in [q, k, v]
|
||||
]
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
block_size = self.page_size
|
||||
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||
query=q,
|
||||
key=k_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_k_head_num * layer.qk_head_dim,
|
||||
),
|
||||
value=v_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_v_head_num * layer.v_head_dim,
|
||||
),
|
||||
pre_tokens=(
|
||||
layer.sliding_window_size
|
||||
if layer.sliding_window_size != -1
|
||||
else FULL_ATTENTION_WINDOW
|
||||
),
|
||||
next_tokens=(
|
||||
0
|
||||
if layer.sliding_window_size != -1
|
||||
else FULL_ATTENTION_WINDOW
|
||||
),
|
||||
atten_mask=self.fia_mask,
|
||||
block_table=block_tables,
|
||||
input_layout="TND",
|
||||
block_size=block_size,
|
||||
num_query_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
actual_seq_qlen=self.forward_metadata.seq_lens_list_cumsum,
|
||||
actual_seq_kvlen=self.forward_metadata.seq_lens_cpu_int,
|
||||
softmax_scale=layer.scaling,
|
||||
sparse_mode=4 if layer.sliding_window_size != -1 else 3,
|
||||
learnable_sink=sinks,
|
||||
)
|
||||
attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||
attn_out = attn_out.view(
|
||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||
)
|
||||
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||
attn_out = torch.cat(
|
||||
[
|
||||
attn_out,
|
||||
attn_out.new_zeros(
|
||||
num_token_padding - attn_out.shape[0],
|
||||
*attn_out.shape[1:],
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
else:
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
|
||||
# FIA BSND with paged KV cache (reads prefix tokens from cache)
|
||||
seq_lens_cpu = forward_batch.seq_lens.cpu().tolist()
|
||||
attn_out = torch.empty(
|
||||
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
|
||||
device=q.device,
|
||||
dtype=q.dtype,
|
||||
)
|
||||
q_len_offset = 0
|
||||
for seq_idx, q_len in enumerate(
|
||||
forward_batch.extend_seq_lens_cpu
|
||||
):
|
||||
if q_len == 0:
|
||||
continue
|
||||
total_kv_len = seq_lens_cpu[seq_idx]
|
||||
result, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||
query=q[None, q_len_offset : q_len_offset + q_len],
|
||||
key=k_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_k_head_num * layer.qk_head_dim,
|
||||
),
|
||||
value=v_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_v_head_num * layer.v_head_dim,
|
||||
),
|
||||
num_query_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="BSND",
|
||||
block_table=block_tables[seq_idx : seq_idx + 1],
|
||||
block_size=self.page_size,
|
||||
actual_seq_qlen=[q_len],
|
||||
actual_seq_kvlen=[total_kv_len],
|
||||
atten_mask=self.fia_mask.unsqueeze(0),
|
||||
sparse_mode=4,
|
||||
softmax_scale=layer.scaling,
|
||||
pre_tokens=layer.sliding_window_size,
|
||||
next_tokens=0,
|
||||
)
|
||||
attn_out[q_len_offset : q_len_offset + q_len] = result[0]
|
||||
q_len_offset += q_len
|
||||
|
||||
attn_out = attn_out.view(
|
||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||
)
|
||||
|
||||
else:
|
||||
attn_out = attention_sinks_prefill_triton(
|
||||
q,
|
||||
@@ -1220,46 +1315,94 @@ class AscendAttnBackend(AttentionBackend):
|
||||
return attn_output
|
||||
|
||||
if self.use_fia:
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
num_token_padding = q.shape[0]
|
||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||
q, k, v = [
|
||||
data[: forward_batch.num_token_non_padded_cpu]
|
||||
for data in [q, k, v]
|
||||
]
|
||||
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
||||
query=q,
|
||||
key=k_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
),
|
||||
value=v_cache.view(
|
||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||
),
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
atten_mask=self.fia_mask,
|
||||
input_layout="TND",
|
||||
actual_seq_lengths=self.forward_metadata.seq_lens_list_cumsum,
|
||||
actual_seq_lengths_kv=self.forward_metadata.seq_lens_cpu_int,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
scale=layer.scaling,
|
||||
sparse_mode=3,
|
||||
)
|
||||
attn_output = attn_output.view(
|
||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||
)
|
||||
if self._can_use_tnd(layer):
|
||||
"""FIA supports multi-bs in the current version of CANN"""
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
num_token_padding = q.shape[0]
|
||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||
q, k, v = [
|
||||
data[: forward_batch.num_token_non_padded_cpu]
|
||||
for data in [q, k, v]
|
||||
]
|
||||
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
||||
query=q,
|
||||
key=k_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
),
|
||||
value=v_cache.view(
|
||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||
),
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
atten_mask=self.fia_mask,
|
||||
input_layout="TND",
|
||||
actual_seq_lengths=self.forward_metadata.seq_lens_list_cumsum,
|
||||
actual_seq_lengths_kv=self.forward_metadata.seq_lens_cpu_int,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
scale=layer.scaling,
|
||||
sparse_mode=3,
|
||||
)
|
||||
attn_output = attn_output.view(
|
||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||
)
|
||||
|
||||
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||
attn_output = torch.cat(
|
||||
[
|
||||
attn_output,
|
||||
attn_output.new_zeros(
|
||||
num_token_padding - attn_output.shape[0],
|
||||
*attn_output.shape[1:],
|
||||
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||
attn_output = torch.cat(
|
||||
[
|
||||
attn_output,
|
||||
attn_output.new_zeros(
|
||||
num_token_padding - attn_output.shape[0],
|
||||
*attn_output.shape[1:],
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
else:
|
||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||
|
||||
# FIA BSND with paged KV cache (reads prefix tokens from cache)
|
||||
seq_lens_cpu = forward_batch.seq_lens.cpu().tolist()
|
||||
attn_output = torch.empty(
|
||||
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
|
||||
device=q.device,
|
||||
dtype=q.dtype,
|
||||
)
|
||||
q_len_offset = 0
|
||||
for seq_idx, q_len in enumerate(forward_batch.extend_seq_lens_cpu):
|
||||
if q_len == 0:
|
||||
continue
|
||||
total_kv_len = seq_lens_cpu[seq_idx]
|
||||
result, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||
query=q[None, q_len_offset : q_len_offset + q_len],
|
||||
key=k_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_k_head_num * layer.qk_head_dim,
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
value=v_cache.view(
|
||||
-1,
|
||||
self.page_size,
|
||||
layer.tp_v_head_num * layer.v_head_dim,
|
||||
),
|
||||
num_query_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="BSND",
|
||||
block_table=self.forward_metadata.block_tables[
|
||||
seq_idx : seq_idx + 1
|
||||
],
|
||||
block_size=self.page_size,
|
||||
actual_seq_qlen=[q_len],
|
||||
actual_seq_kvlen=[total_kv_len],
|
||||
atten_mask=self.fia_mask.unsqueeze(0),
|
||||
sparse_mode=3,
|
||||
softmax_scale=layer.scaling,
|
||||
)
|
||||
attn_output[q_len_offset : q_len_offset + q_len] = result[0]
|
||||
q_len_offset += q_len
|
||||
|
||||
attn_output = attn_output.view(
|
||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||
)
|
||||
else:
|
||||
causal = True
|
||||
@@ -1340,6 +1483,12 @@ class AscendAttnBackend(AttentionBackend):
|
||||
scaling=layer.scaling,
|
||||
enable_gqa=use_gqa,
|
||||
causal=causal,
|
||||
sliding_window_size=layer.sliding_window_size,
|
||||
full_to_swa_mapping=(
|
||||
self.full_to_swa_index_mapping
|
||||
if self._is_swa_layer(layer)
|
||||
else None
|
||||
),
|
||||
logit_cap=layer.logit_cap,
|
||||
logit_capping_method=layer.logit_capping_method,
|
||||
)
|
||||
@@ -1957,7 +2106,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
|
||||
if sinks is not None:
|
||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||
if self._is_swa_layer(layer):
|
||||
block_tables = self.forward_metadata.block_tables_swa
|
||||
else:
|
||||
block_tables = self.forward_metadata.block_tables
|
||||
@@ -2041,6 +2190,17 @@ class AscendAttnBackend(AttentionBackend):
|
||||
return attn_out
|
||||
|
||||
if not self.use_mla:
|
||||
seq_lens_cpu_int = self.forward_metadata.seq_lens_cpu_int
|
||||
seq_lens_cpu_list = self.forward_metadata.seq_lens_cpu_list
|
||||
if self._is_swa_layer(layer):
|
||||
# CUDA/NPU graph capture uses seq_len fill value 0 on Ascend.
|
||||
# Avoid dynamic window block-table construction during capture,
|
||||
# because it can create a zero-width block table and break tiling.
|
||||
block_tables = self.forward_metadata.block_tables_swa
|
||||
attn_mask = self.forward_metadata.swa_mask
|
||||
else:
|
||||
block_tables = self.forward_metadata.block_tables
|
||||
attn_mask = None
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
)
|
||||
@@ -2048,12 +2208,10 @@ class AscendAttnBackend(AttentionBackend):
|
||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||
)
|
||||
query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim)
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
||||
if seq_lens_cpu_int is None:
|
||||
actual_seq_len_kv = seq_lens_cpu_list
|
||||
else:
|
||||
actual_seq_len_kv = (
|
||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||
)
|
||||
actual_seq_len_kv = seq_lens_cpu_int.cpu().int().tolist()
|
||||
|
||||
if (layer.qk_head_dim != layer.v_head_dim) and (
|
||||
self.is_hybrid_swa and layer.sliding_window_size == -1
|
||||
@@ -2113,13 +2271,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
query,
|
||||
k_cache,
|
||||
v_cache,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_table=block_tables,
|
||||
block_size=self.page_size,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="BSH",
|
||||
scale=layer.scaling,
|
||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||
atten_mask=attn_mask,
|
||||
sparse_mode=0,
|
||||
)
|
||||
output = torch.empty(
|
||||
(num_tokens, 1, layer.tp_q_head_num * layer.v_head_dim),
|
||||
@@ -2131,13 +2291,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
query,
|
||||
k_cache,
|
||||
v_cache,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_table=block_tables,
|
||||
block_size=self.page_size,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="BSH",
|
||||
scale=layer.scaling,
|
||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||
atten_mask=attn_mask,
|
||||
sparse_mode=0,
|
||||
workspace=workspace,
|
||||
out=[output, softmax_lse],
|
||||
)
|
||||
@@ -2298,11 +2460,9 @@ class AscendAttnBackend(AttentionBackend):
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
|
||||
if sinks is not None or (
|
||||
self.is_hybrid_swa and layer.sliding_window_size != -1
|
||||
):
|
||||
if sinks is not None or (self._is_swa_layer(layer) and self.use_fia):
|
||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||
if self._is_swa_layer(layer):
|
||||
block_tables = self.forward_metadata.block_tables_swa
|
||||
else:
|
||||
block_tables = self.forward_metadata.block_tables
|
||||
@@ -2473,6 +2633,12 @@ class AscendAttnBackend(AttentionBackend):
|
||||
scaling=layer.scaling,
|
||||
enable_gqa=use_gqa,
|
||||
causal=False,
|
||||
sliding_window_size=layer.sliding_window_size,
|
||||
full_to_swa_mapping=(
|
||||
self.full_to_swa_index_mapping
|
||||
if self._is_swa_layer(layer)
|
||||
else None
|
||||
),
|
||||
logit_cap=layer.logit_cap,
|
||||
logit_capping_method=layer.logit_capping_method,
|
||||
)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch.nn.functional import scaled_dot_product_attention
|
||||
@@ -69,6 +70,8 @@ class AscendTorchNativeAttnBackend:
|
||||
scaling=None,
|
||||
enable_gqa=False,
|
||||
causal=False,
|
||||
sliding_window_size: int = -1,
|
||||
full_to_swa_mapping: Optional[torch.Tensor] = None,
|
||||
logit_cap: float = 0.0,
|
||||
logit_capping_method: str = "tanh",
|
||||
):
|
||||
@@ -89,6 +92,9 @@ class AscendTorchNativeAttnBackend:
|
||||
scaling: float or None
|
||||
enable_gqa: bool
|
||||
causal: bool
|
||||
sliding_window_size: int, -1 means no sliding window
|
||||
full_to_swa_mapping: mapping from full pool index to SWA pool index,
|
||||
required for SWA layers to translate req_to_token indices
|
||||
|
||||
Returns:
|
||||
output: [num_tokens, num_heads, head_size]
|
||||
@@ -104,35 +110,54 @@ class AscendTorchNativeAttnBackend:
|
||||
for seq_idx in range(seq_lens.shape[0]):
|
||||
# Need optimize the performance later.
|
||||
|
||||
extend_seq_len_q = extend_seq_lens[seq_idx]
|
||||
prefill_seq_len_q = extend_prefix_lens[seq_idx]
|
||||
extend_seq_len_q = int(extend_seq_lens[seq_idx].item())
|
||||
prefill_seq_len_q = int(extend_prefix_lens[seq_idx].item())
|
||||
|
||||
seq_len_kv = seq_lens[seq_idx]
|
||||
seq_len_kv = int(seq_lens[seq_idx].item())
|
||||
end_q = start_q + extend_seq_len_q
|
||||
end_kv = start_kv + seq_len_kv
|
||||
atten_start_kv = 0
|
||||
atten_end_kv = seq_lens[seq_idx]
|
||||
atten_end_kv = seq_len_kv
|
||||
# support cross attention
|
||||
if encoder_lens is not None:
|
||||
if is_cross_attention:
|
||||
atten_end_kv = encoder_lens[seq_idx]
|
||||
atten_end_kv = int(encoder_lens[seq_idx].item())
|
||||
else:
|
||||
atten_start_kv = encoder_lens[seq_idx]
|
||||
atten_end_kv = encoder_lens[seq_idx] + extend_seq_len_q
|
||||
atten_start_kv = int(encoder_lens[seq_idx].item())
|
||||
atten_end_kv = atten_start_kv + extend_seq_len_q
|
||||
|
||||
if (
|
||||
sliding_window_size is not None
|
||||
and sliding_window_size > -1
|
||||
and encoder_lens is None
|
||||
):
|
||||
# For extend, the sliding window must be anchored at the first
|
||||
# query token in this chunk rather than the final sequence
|
||||
# length. Otherwise a large extend chunk can no longer fit in
|
||||
# the cropped query suffix, which breaks the native fallback.
|
||||
atten_start_kv = max(
|
||||
prefill_seq_len_q - sliding_window_size, atten_start_kv
|
||||
)
|
||||
|
||||
per_req_query = query[:, start_q:end_q, :]
|
||||
per_req_query_redudant = torch.empty(
|
||||
query_start_idx = max(prefill_seq_len_q - atten_start_kv, 0)
|
||||
seq_len_kv = atten_end_kv - atten_start_kv
|
||||
per_req_query_redundant = torch.zeros(
|
||||
(per_req_query.shape[0], seq_len_kv, per_req_query.shape[2]),
|
||||
dtype=per_req_query.dtype,
|
||||
device=per_req_query.device,
|
||||
)
|
||||
|
||||
per_req_query_redudant[:, prefill_seq_len_q:, :] = per_req_query
|
||||
per_req_query_redundant[:, query_start_idx:, :] = per_req_query
|
||||
|
||||
# get key and value from cache. per_req_tokens contains the kv cache
|
||||
# index for each token in the sequence.
|
||||
req_pool_idx = req_pool_indices[seq_idx]
|
||||
per_req_tokens = req_to_token[req_pool_idx, atten_start_kv:atten_end_kv]
|
||||
# For SWA layers, k_cache/v_cache are from the SWA pool but
|
||||
# req_to_token stores full pool indices. Translate before indexing.
|
||||
if full_to_swa_mapping is not None:
|
||||
per_req_tokens = full_to_swa_mapping[per_req_tokens]
|
||||
per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||
per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||
|
||||
@@ -142,9 +167,9 @@ class AscendTorchNativeAttnBackend:
|
||||
per_req_value = per_req_value.to(per_req_query.dtype)
|
||||
|
||||
if logit_cap > 0:
|
||||
per_req_out_redudant = (
|
||||
per_req_out_redundant = (
|
||||
self.scaled_dot_product_attention_with_softcapping(
|
||||
per_req_query_redudant.unsqueeze(0),
|
||||
per_req_query_redundant.unsqueeze(0),
|
||||
per_req_key.unsqueeze(0),
|
||||
per_req_value.unsqueeze(0),
|
||||
enable_gqa=enable_gqa,
|
||||
@@ -157,9 +182,9 @@ class AscendTorchNativeAttnBackend:
|
||||
.movedim(query.dim() - 2, 0)
|
||||
)
|
||||
else:
|
||||
per_req_out_redudant = (
|
||||
per_req_out_redundant = (
|
||||
scaled_dot_product_attention(
|
||||
per_req_query_redudant.unsqueeze(0),
|
||||
per_req_query_redundant.unsqueeze(0),
|
||||
per_req_key.unsqueeze(0),
|
||||
per_req_value.unsqueeze(0),
|
||||
enable_gqa=enable_gqa,
|
||||
@@ -169,7 +194,9 @@ class AscendTorchNativeAttnBackend:
|
||||
.squeeze(0)
|
||||
.movedim(query.dim() - 2, 0)
|
||||
)
|
||||
output[start_q:end_q, :, :] = per_req_out_redudant[prefill_seq_len_q:, :, :]
|
||||
output[start_q:end_q, :, :] = per_req_out_redundant[
|
||||
query_start_idx : query_start_idx + extend_seq_len_q, :, :
|
||||
]
|
||||
start_q, start_kv = end_q, end_kv
|
||||
return output
|
||||
|
||||
@@ -187,6 +214,8 @@ class AscendTorchNativeAttnBackend:
|
||||
scaling=None,
|
||||
enable_gqa=False,
|
||||
causal=False,
|
||||
sliding_window_size: int = -1,
|
||||
full_to_swa_mapping: Optional[torch.Tensor] = None,
|
||||
logit_cap: float = 0.0,
|
||||
logit_capping_method: str = "tanh",
|
||||
):
|
||||
@@ -218,18 +247,27 @@ class AscendTorchNativeAttnBackend:
|
||||
# Need optimize the performance later.
|
||||
|
||||
seq_len_q = 1
|
||||
seq_len_kv = seq_lens[seq_idx]
|
||||
seq_len_kv = int(seq_lens[seq_idx].item())
|
||||
end_q = start_q + seq_len_q
|
||||
end_kv = start_kv + seq_len_kv
|
||||
atten_start_kv = 0
|
||||
atten_end_kv = seq_lens[seq_idx]
|
||||
atten_end_kv = seq_len_kv
|
||||
# support cross attention
|
||||
if encoder_lens is not None:
|
||||
if is_cross_attention:
|
||||
atten_end_kv = encoder_lens[seq_idx]
|
||||
atten_end_kv = int(encoder_lens[seq_idx].item())
|
||||
else:
|
||||
atten_start_kv = encoder_lens[seq_idx]
|
||||
atten_end_kv = encoder_lens[seq_idx] + seq_len_kv
|
||||
atten_start_kv = int(encoder_lens[seq_idx].item())
|
||||
atten_end_kv = atten_start_kv + seq_len_kv
|
||||
|
||||
if (
|
||||
sliding_window_size is not None
|
||||
and sliding_window_size > -1
|
||||
and encoder_lens is None
|
||||
):
|
||||
atten_start_kv = max(
|
||||
atten_end_kv - (sliding_window_size + 1), atten_start_kv
|
||||
)
|
||||
|
||||
per_req_query = query[:, start_q:end_q, :]
|
||||
|
||||
@@ -237,6 +275,10 @@ class AscendTorchNativeAttnBackend:
|
||||
# index for each token in the sequence.
|
||||
req_pool_idx = req_pool_indices[seq_idx]
|
||||
per_req_tokens = req_to_token[req_pool_idx, atten_start_kv:atten_end_kv]
|
||||
# For SWA layers, k_cache/v_cache are from the SWA pool but
|
||||
# req_to_token stores full pool indices. Translate before indexing.
|
||||
if full_to_swa_mapping is not None:
|
||||
per_req_tokens = full_to_swa_mapping[per_req_tokens]
|
||||
per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||
per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||
|
||||
|
||||
@@ -295,7 +295,9 @@ class Gemma4Attention(nn.Module):
|
||||
|
||||
layer_type = config.layer_types[layer_id]
|
||||
self.sliding_window = (
|
||||
config.sliding_window if layer_type == "sliding_attention" else None
|
||||
get_attention_sliding_window_size(config)
|
||||
if layer_type == "sliding_attention"
|
||||
else -1
|
||||
)
|
||||
|
||||
self.total_num_heads = config.num_attention_heads
|
||||
|
||||
@@ -608,9 +608,16 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
|
||||
if is_first_rank and input_ids is not None:
|
||||
ple_ids = input_ids.clone()
|
||||
pad_id = self.config.text_config.pad_token_id
|
||||
ple_ids[input_ids == self.config.image_token_id] = pad_id
|
||||
ple_ids[input_ids == self.config.video_token_id] = pad_id
|
||||
ple_ids[input_ids == self.config.audio_token_id] = pad_id
|
||||
# Use torch.where instead of boolean indexing for NPU graph compatibility
|
||||
ple_ids = torch.where(
|
||||
input_ids == self.config.image_token_id, pad_id, ple_ids
|
||||
)
|
||||
ple_ids = torch.where(
|
||||
input_ids == self.config.video_token_id, pad_id, ple_ids
|
||||
)
|
||||
ple_ids = torch.where(
|
||||
input_ids == self.config.audio_token_id, pad_id, ple_ids
|
||||
)
|
||||
per_layer_inputs = self.get_per_layer_inputs(ple_ids)
|
||||
|
||||
# Prepare bidirectional attention masks for image tokens during prefill.
|
||||
|
||||
@@ -2542,7 +2542,7 @@ class ServerArgs:
|
||||
self.attention_backend = default_attention_backend
|
||||
|
||||
prefill_backend, decode_backend = self.get_attention_backends()
|
||||
accepted_backends = ("trtllm_mha", "triton", "intel_xpu")
|
||||
accepted_backends = ("trtllm_mha", "triton", "ascend", "intel_xpu")
|
||||
assert (
|
||||
prefill_backend in accepted_backends
|
||||
and decode_backend in accepted_backends
|
||||
|
||||
@@ -68,6 +68,12 @@ EXAONE_3_5_7_8B_INSTRUCT_WEIGHTS_PATH = os.path.join(
|
||||
MODEL_WEIGHTS_DIR, "LGAI-EXAONE/EXAONE-3.5-7.8B-Instruct"
|
||||
)
|
||||
GEMMA_3_4B_IT_WEIGHTS_PATH = os.path.join(MODEL_WEIGHTS_DIR, "google/gemma-3-4b-it")
|
||||
GEMMA_4_E2B_WEIGHTS_PATH = os.path.join(MODEL_WEIGHTS_DIR, "google/gemma-4-E2B-it")
|
||||
GEMMA_4_E4B_WEIGHTS_PATH = os.path.join(MODEL_WEIGHTS_DIR, "google/gemma-4-E4B-it")
|
||||
GEMMA_4_26B_A4B_IT_WEIGHTS_PATH = os.path.join(
|
||||
MODEL_WEIGHTS_DIR, "google/gemma-4-26B-A4B-it"
|
||||
)
|
||||
GEMMA_4_31B_WEIGHTS_PATH = os.path.join(MODEL_WEIGHTS_DIR, "google/gemma-4-31B-it")
|
||||
GLM_4_9B_CHAT_WEIGHTS_PATH = os.path.join(MODEL_WEIGHTS_DIR, "ZhipuAI/glm-4-9b-chat")
|
||||
GRANITE_3_0_3B_A800M_INSTRUCT_WEIGHTS_PATH = os.path.join(
|
||||
MODEL_WEIGHTS_DIR, "ibm-granite/granite-3.0-3b-a800m-instruct"
|
||||
|
||||
Reference in New Issue
Block a user