[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
|
import numpy as np
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
FULL_ATTENTION_WINDOW = 2147483647
|
||||||
|
|
||||||
|
|
||||||
def _reshape_kv_for_fia_nz(
|
def _reshape_kv_for_fia_nz(
|
||||||
tensor: torch.Tensor, num_heads: int, head_dim: int, page_size: int
|
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)
|
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
|
@dataclass
|
||||||
class ForwardMetadata:
|
class ForwardMetadata:
|
||||||
|
|
||||||
@@ -73,6 +70,9 @@ class ForwardMetadata:
|
|||||||
actual_seq_lengths_q: Optional[torch.Tensor] = None
|
actual_seq_lengths_q: Optional[torch.Tensor] = None
|
||||||
actual_seq_lengths_kv: 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 cache
|
||||||
prefix_lens: Optional[torch.Tensor] = None
|
prefix_lens: Optional[torch.Tensor] = None
|
||||||
flatten_prefix_block_tables: 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 = (
|
self.full_to_swa_index_mapping = (
|
||||||
model_runner.token_to_kv_pool.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 = (
|
self.use_sliding_window_kv_pool = (
|
||||||
isinstance(self.token_to_kv_pool, SWAKVPool)
|
isinstance(self.token_to_kv_pool, SWAKVPool)
|
||||||
and self.token_to_kv_pool.swa_layer_nums > 0
|
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
|
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):
|
def get_verify_buffers_to_fill_after_draft(self):
|
||||||
"""
|
"""
|
||||||
Return buffers for verify attention kernels that needs to be filled after draft.
|
Return buffers for verify attention kernels that needs to be filled after draft.
|
||||||
@@ -516,6 +531,19 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
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:
|
if self.use_sliding_window_kv_pool:
|
||||||
# refilled in place at replay; the captured graph reads this storage
|
# refilled in place at replay; the captured graph reads this storage
|
||||||
self.swa_out_cache_loc_buf = torch.zeros(
|
self.swa_out_cache_loc_buf = torch.zeros(
|
||||||
@@ -536,6 +564,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
|
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
|
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:
|
if self.use_sliding_window_kv_pool and out_cache_loc is not None:
|
||||||
num_tokens = out_cache_loc.shape[0]
|
num_tokens = out_cache_loc.shape[0]
|
||||||
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
|
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, max_seq_pages:].fill_(0)
|
||||||
metadata.block_tables_swa[bs:, :].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_(
|
metadata.block_tables[:bs, :max_seq_pages].copy_(
|
||||||
self.req_to_token[req_pool_indices[:bs], :max_len][:, :: self.page_size]
|
self.req_to_token[req_pool_indices[:bs], :max_len][:, :: self.page_size]
|
||||||
// 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)
|
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)
|
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
if sinks is not None or (
|
if sinks is not None or (self._is_swa_layer(layer) and self.use_fia):
|
||||||
self.is_hybrid_swa and layer.sliding_window_size != -1
|
|
||||||
):
|
|
||||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
# 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
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
else:
|
else:
|
||||||
block_tables = self.forward_metadata.block_tables
|
block_tables = self.forward_metadata.block_tables
|
||||||
if self.use_fia:
|
if self.use_fia:
|
||||||
num_token_padding = q.shape[0]
|
if self._can_use_tnd(layer):
|
||||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
num_token_padding = q.shape[0]
|
||||||
q, k, v = [
|
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||||
data[: forward_batch.num_token_non_padded_cpu]
|
q, k, v = [
|
||||||
for data in [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
|
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||||
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
block_size = self.page_size
|
||||||
query=q,
|
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
|
||||||
key=k_cache.view(
|
query=q,
|
||||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
key=k_cache.view(
|
||||||
),
|
-1,
|
||||||
value=v_cache.view(
|
self.page_size,
|
||||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
layer.tp_k_head_num * layer.qk_head_dim,
|
||||||
),
|
),
|
||||||
pre_tokens=(
|
value=v_cache.view(
|
||||||
layer.sliding_window_size
|
-1,
|
||||||
if layer.sliding_window_size != -1
|
self.page_size,
|
||||||
else FULL_ATTENTION_WINDOW
|
layer.tp_v_head_num * layer.v_head_dim,
|
||||||
),
|
),
|
||||||
next_tokens=(
|
pre_tokens=(
|
||||||
0
|
layer.sliding_window_size
|
||||||
if layer.sliding_window_size != -1
|
if layer.sliding_window_size != -1
|
||||||
else FULL_ATTENTION_WINDOW
|
else FULL_ATTENTION_WINDOW
|
||||||
),
|
),
|
||||||
atten_mask=self.fia_mask,
|
next_tokens=(
|
||||||
block_table=block_tables,
|
0
|
||||||
input_layout="TND",
|
if layer.sliding_window_size != -1
|
||||||
block_size=block_size,
|
else FULL_ATTENTION_WINDOW
|
||||||
num_query_heads=layer.tp_q_head_num,
|
),
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
atten_mask=self.fia_mask,
|
||||||
actual_seq_qlen=self.forward_metadata.seq_lens_list_cumsum,
|
block_table=block_tables,
|
||||||
actual_seq_kvlen=self.forward_metadata.seq_lens_cpu_int,
|
input_layout="TND",
|
||||||
softmax_scale=layer.scaling,
|
block_size=block_size,
|
||||||
sparse_mode=4 if layer.sliding_window_size != -1 else 3,
|
num_query_heads=layer.tp_q_head_num,
|
||||||
learnable_sink=sinks,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
)
|
actual_seq_qlen=self.forward_metadata.seq_lens_list_cumsum,
|
||||||
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
actual_seq_kvlen=self.forward_metadata.seq_lens_cpu_int,
|
||||||
attn_out = torch.cat(
|
softmax_scale=layer.scaling,
|
||||||
[
|
sparse_mode=4 if layer.sliding_window_size != -1 else 3,
|
||||||
attn_out,
|
learnable_sink=sinks,
|
||||||
attn_out.new_zeros(
|
|
||||||
num_token_padding - attn_out.shape[0],
|
|
||||||
*attn_out.shape[1:],
|
|
||||||
),
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
)
|
)
|
||||||
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:
|
else:
|
||||||
attn_out = attention_sinks_prefill_triton(
|
attn_out = attention_sinks_prefill_triton(
|
||||||
q,
|
q,
|
||||||
@@ -1220,46 +1315,94 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
return attn_output
|
return attn_output
|
||||||
|
|
||||||
if self.use_fia:
|
if self.use_fia:
|
||||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
if self._can_use_tnd(layer):
|
||||||
num_token_padding = q.shape[0]
|
"""FIA supports multi-bs in the current version of CANN"""
|
||||||
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||||
q, k, v = [
|
num_token_padding = q.shape[0]
|
||||||
data[: forward_batch.num_token_non_padded_cpu]
|
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||||
for data in [q, k, v]
|
q, k, v = [
|
||||||
]
|
data[: forward_batch.num_token_non_padded_cpu]
|
||||||
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
for data in [q, k, v]
|
||||||
query=q,
|
]
|
||||||
key=k_cache.view(
|
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
||||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
query=q,
|
||||||
),
|
key=k_cache.view(
|
||||||
value=v_cache.view(
|
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||||
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
),
|
||||||
),
|
value=v_cache.view(
|
||||||
block_table=self.forward_metadata.block_tables,
|
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||||
block_size=self.page_size,
|
),
|
||||||
atten_mask=self.fia_mask,
|
block_table=self.forward_metadata.block_tables,
|
||||||
input_layout="TND",
|
block_size=self.page_size,
|
||||||
actual_seq_lengths=self.forward_metadata.seq_lens_list_cumsum,
|
atten_mask=self.fia_mask,
|
||||||
actual_seq_lengths_kv=self.forward_metadata.seq_lens_cpu_int,
|
input_layout="TND",
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
actual_seq_lengths=self.forward_metadata.seq_lens_list_cumsum,
|
||||||
num_heads=layer.tp_q_head_num,
|
actual_seq_lengths_kv=self.forward_metadata.seq_lens_cpu_int,
|
||||||
scale=layer.scaling,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
sparse_mode=3,
|
num_heads=layer.tp_q_head_num,
|
||||||
)
|
scale=layer.scaling,
|
||||||
attn_output = attn_output.view(
|
sparse_mode=3,
|
||||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
)
|
||||||
)
|
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:
|
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||||
attn_output = torch.cat(
|
attn_output = torch.cat(
|
||||||
[
|
[
|
||||||
attn_output,
|
attn_output,
|
||||||
attn_output.new_zeros(
|
attn_output.new_zeros(
|
||||||
num_token_padding - attn_output.shape[0],
|
num_token_padding - attn_output.shape[0],
|
||||||
*attn_output.shape[1:],
|
*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,
|
||||||
),
|
),
|
||||||
],
|
value=v_cache.view(
|
||||||
dim=0,
|
-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:
|
else:
|
||||||
causal = True
|
causal = True
|
||||||
@@ -1340,6 +1483,12 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
scaling=layer.scaling,
|
scaling=layer.scaling,
|
||||||
enable_gqa=use_gqa,
|
enable_gqa=use_gqa,
|
||||||
causal=causal,
|
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_cap=layer.logit_cap,
|
||||||
logit_capping_method=layer.logit_capping_method,
|
logit_capping_method=layer.logit_capping_method,
|
||||||
)
|
)
|
||||||
@@ -1957,7 +2106,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
if sinks is not None:
|
if sinks is not None:
|
||||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
# 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
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
else:
|
else:
|
||||||
block_tables = self.forward_metadata.block_tables
|
block_tables = self.forward_metadata.block_tables
|
||||||
@@ -2041,6 +2190,17 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
return attn_out
|
return attn_out
|
||||||
|
|
||||||
if not self.use_mla:
|
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(
|
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
|
-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
|
-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)
|
query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim)
|
||||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
if seq_lens_cpu_int is None:
|
||||||
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
actual_seq_len_kv = seq_lens_cpu_list
|
||||||
else:
|
else:
|
||||||
actual_seq_len_kv = (
|
actual_seq_len_kv = seq_lens_cpu_int.cpu().int().tolist()
|
||||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
|
||||||
)
|
|
||||||
|
|
||||||
if (layer.qk_head_dim != layer.v_head_dim) and (
|
if (layer.qk_head_dim != layer.v_head_dim) and (
|
||||||
self.is_hybrid_swa and layer.sliding_window_size == -1
|
self.is_hybrid_swa and layer.sliding_window_size == -1
|
||||||
@@ -2113,13 +2271,15 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
query,
|
query,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
block_table=self.forward_metadata.block_tables,
|
block_table=block_tables,
|
||||||
block_size=self.page_size,
|
block_size=self.page_size,
|
||||||
num_heads=layer.tp_q_head_num,
|
num_heads=layer.tp_q_head_num,
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
input_layout="BSH",
|
input_layout="BSH",
|
||||||
scale=layer.scaling,
|
scale=layer.scaling,
|
||||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||||
|
atten_mask=attn_mask,
|
||||||
|
sparse_mode=0,
|
||||||
)
|
)
|
||||||
output = torch.empty(
|
output = torch.empty(
|
||||||
(num_tokens, 1, layer.tp_q_head_num * layer.v_head_dim),
|
(num_tokens, 1, layer.tp_q_head_num * layer.v_head_dim),
|
||||||
@@ -2131,13 +2291,15 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
query,
|
query,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
block_table=self.forward_metadata.block_tables,
|
block_table=block_tables,
|
||||||
block_size=self.page_size,
|
block_size=self.page_size,
|
||||||
num_heads=layer.tp_q_head_num,
|
num_heads=layer.tp_q_head_num,
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
input_layout="BSH",
|
input_layout="BSH",
|
||||||
scale=layer.scaling,
|
scale=layer.scaling,
|
||||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||||
|
atten_mask=attn_mask,
|
||||||
|
sparse_mode=0,
|
||||||
workspace=workspace,
|
workspace=workspace,
|
||||||
out=[output, softmax_lse],
|
out=[output, softmax_lse],
|
||||||
)
|
)
|
||||||
@@ -2298,11 +2460,9 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
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)
|
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
if sinks is not None or (
|
if sinks is not None or (self._is_swa_layer(layer) and self.use_fia):
|
||||||
self.is_hybrid_swa and layer.sliding_window_size != -1
|
|
||||||
):
|
|
||||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
# 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
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
else:
|
else:
|
||||||
block_tables = self.forward_metadata.block_tables
|
block_tables = self.forward_metadata.block_tables
|
||||||
@@ -2473,6 +2633,12 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
scaling=layer.scaling,
|
scaling=layer.scaling,
|
||||||
enable_gqa=use_gqa,
|
enable_gqa=use_gqa,
|
||||||
causal=False,
|
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_cap=layer.logit_cap,
|
||||||
logit_capping_method=layer.logit_capping_method,
|
logit_capping_method=layer.logit_capping_method,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import math
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch.nn.functional import scaled_dot_product_attention
|
from torch.nn.functional import scaled_dot_product_attention
|
||||||
@@ -69,6 +70,8 @@ class AscendTorchNativeAttnBackend:
|
|||||||
scaling=None,
|
scaling=None,
|
||||||
enable_gqa=False,
|
enable_gqa=False,
|
||||||
causal=False,
|
causal=False,
|
||||||
|
sliding_window_size: int = -1,
|
||||||
|
full_to_swa_mapping: Optional[torch.Tensor] = None,
|
||||||
logit_cap: float = 0.0,
|
logit_cap: float = 0.0,
|
||||||
logit_capping_method: str = "tanh",
|
logit_capping_method: str = "tanh",
|
||||||
):
|
):
|
||||||
@@ -89,6 +92,9 @@ class AscendTorchNativeAttnBackend:
|
|||||||
scaling: float or None
|
scaling: float or None
|
||||||
enable_gqa: bool
|
enable_gqa: bool
|
||||||
causal: 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:
|
Returns:
|
||||||
output: [num_tokens, num_heads, head_size]
|
output: [num_tokens, num_heads, head_size]
|
||||||
@@ -104,35 +110,54 @@ class AscendTorchNativeAttnBackend:
|
|||||||
for seq_idx in range(seq_lens.shape[0]):
|
for seq_idx in range(seq_lens.shape[0]):
|
||||||
# Need optimize the performance later.
|
# Need optimize the performance later.
|
||||||
|
|
||||||
extend_seq_len_q = extend_seq_lens[seq_idx]
|
extend_seq_len_q = int(extend_seq_lens[seq_idx].item())
|
||||||
prefill_seq_len_q = extend_prefix_lens[seq_idx]
|
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_q = start_q + extend_seq_len_q
|
||||||
end_kv = start_kv + seq_len_kv
|
end_kv = start_kv + seq_len_kv
|
||||||
atten_start_kv = 0
|
atten_start_kv = 0
|
||||||
atten_end_kv = seq_lens[seq_idx]
|
atten_end_kv = seq_len_kv
|
||||||
# support cross attention
|
# support cross attention
|
||||||
if encoder_lens is not None:
|
if encoder_lens is not None:
|
||||||
if is_cross_attention:
|
if is_cross_attention:
|
||||||
atten_end_kv = encoder_lens[seq_idx]
|
atten_end_kv = int(encoder_lens[seq_idx].item())
|
||||||
else:
|
else:
|
||||||
atten_start_kv = encoder_lens[seq_idx]
|
atten_start_kv = int(encoder_lens[seq_idx].item())
|
||||||
atten_end_kv = encoder_lens[seq_idx] + extend_seq_len_q
|
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 = 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]),
|
(per_req_query.shape[0], seq_len_kv, per_req_query.shape[2]),
|
||||||
dtype=per_req_query.dtype,
|
dtype=per_req_query.dtype,
|
||||||
device=per_req_query.device,
|
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
|
# get key and value from cache. per_req_tokens contains the kv cache
|
||||||
# index for each token in the sequence.
|
# index for each token in the sequence.
|
||||||
req_pool_idx = req_pool_indices[seq_idx]
|
req_pool_idx = req_pool_indices[seq_idx]
|
||||||
per_req_tokens = req_to_token[req_pool_idx, atten_start_kv:atten_end_kv]
|
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_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||||
per_req_value = v_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)
|
per_req_value = per_req_value.to(per_req_query.dtype)
|
||||||
|
|
||||||
if logit_cap > 0:
|
if logit_cap > 0:
|
||||||
per_req_out_redudant = (
|
per_req_out_redundant = (
|
||||||
self.scaled_dot_product_attention_with_softcapping(
|
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_key.unsqueeze(0),
|
||||||
per_req_value.unsqueeze(0),
|
per_req_value.unsqueeze(0),
|
||||||
enable_gqa=enable_gqa,
|
enable_gqa=enable_gqa,
|
||||||
@@ -157,9 +182,9 @@ class AscendTorchNativeAttnBackend:
|
|||||||
.movedim(query.dim() - 2, 0)
|
.movedim(query.dim() - 2, 0)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
per_req_out_redudant = (
|
per_req_out_redundant = (
|
||||||
scaled_dot_product_attention(
|
scaled_dot_product_attention(
|
||||||
per_req_query_redudant.unsqueeze(0),
|
per_req_query_redundant.unsqueeze(0),
|
||||||
per_req_key.unsqueeze(0),
|
per_req_key.unsqueeze(0),
|
||||||
per_req_value.unsqueeze(0),
|
per_req_value.unsqueeze(0),
|
||||||
enable_gqa=enable_gqa,
|
enable_gqa=enable_gqa,
|
||||||
@@ -169,7 +194,9 @@ class AscendTorchNativeAttnBackend:
|
|||||||
.squeeze(0)
|
.squeeze(0)
|
||||||
.movedim(query.dim() - 2, 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
|
start_q, start_kv = end_q, end_kv
|
||||||
return output
|
return output
|
||||||
|
|
||||||
@@ -187,6 +214,8 @@ class AscendTorchNativeAttnBackend:
|
|||||||
scaling=None,
|
scaling=None,
|
||||||
enable_gqa=False,
|
enable_gqa=False,
|
||||||
causal=False,
|
causal=False,
|
||||||
|
sliding_window_size: int = -1,
|
||||||
|
full_to_swa_mapping: Optional[torch.Tensor] = None,
|
||||||
logit_cap: float = 0.0,
|
logit_cap: float = 0.0,
|
||||||
logit_capping_method: str = "tanh",
|
logit_capping_method: str = "tanh",
|
||||||
):
|
):
|
||||||
@@ -218,18 +247,27 @@ class AscendTorchNativeAttnBackend:
|
|||||||
# Need optimize the performance later.
|
# Need optimize the performance later.
|
||||||
|
|
||||||
seq_len_q = 1
|
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_q = start_q + seq_len_q
|
||||||
end_kv = start_kv + seq_len_kv
|
end_kv = start_kv + seq_len_kv
|
||||||
atten_start_kv = 0
|
atten_start_kv = 0
|
||||||
atten_end_kv = seq_lens[seq_idx]
|
atten_end_kv = seq_len_kv
|
||||||
# support cross attention
|
# support cross attention
|
||||||
if encoder_lens is not None:
|
if encoder_lens is not None:
|
||||||
if is_cross_attention:
|
if is_cross_attention:
|
||||||
atten_end_kv = encoder_lens[seq_idx]
|
atten_end_kv = int(encoder_lens[seq_idx].item())
|
||||||
else:
|
else:
|
||||||
atten_start_kv = encoder_lens[seq_idx]
|
atten_start_kv = int(encoder_lens[seq_idx].item())
|
||||||
atten_end_kv = encoder_lens[seq_idx] + seq_len_kv
|
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, :]
|
per_req_query = query[:, start_q:end_q, :]
|
||||||
|
|
||||||
@@ -237,6 +275,10 @@ class AscendTorchNativeAttnBackend:
|
|||||||
# index for each token in the sequence.
|
# index for each token in the sequence.
|
||||||
req_pool_idx = req_pool_indices[seq_idx]
|
req_pool_idx = req_pool_indices[seq_idx]
|
||||||
per_req_tokens = req_to_token[req_pool_idx, atten_start_kv:atten_end_kv]
|
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_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
||||||
per_req_value = v_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]
|
layer_type = config.layer_types[layer_id]
|
||||||
self.sliding_window = (
|
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
|
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:
|
if is_first_rank and input_ids is not None:
|
||||||
ple_ids = input_ids.clone()
|
ple_ids = input_ids.clone()
|
||||||
pad_id = self.config.text_config.pad_token_id
|
pad_id = self.config.text_config.pad_token_id
|
||||||
ple_ids[input_ids == self.config.image_token_id] = pad_id
|
# Use torch.where instead of boolean indexing for NPU graph compatibility
|
||||||
ple_ids[input_ids == self.config.video_token_id] = pad_id
|
ple_ids = torch.where(
|
||||||
ple_ids[input_ids == self.config.audio_token_id] = pad_id
|
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)
|
per_layer_inputs = self.get_per_layer_inputs(ple_ids)
|
||||||
|
|
||||||
# Prepare bidirectional attention masks for image tokens during prefill.
|
# Prepare bidirectional attention masks for image tokens during prefill.
|
||||||
|
|||||||
@@ -2542,7 +2542,7 @@ class ServerArgs:
|
|||||||
self.attention_backend = default_attention_backend
|
self.attention_backend = default_attention_backend
|
||||||
|
|
||||||
prefill_backend, decode_backend = self.get_attention_backends()
|
prefill_backend, decode_backend = self.get_attention_backends()
|
||||||
accepted_backends = ("trtllm_mha", "triton", "intel_xpu")
|
accepted_backends = ("trtllm_mha", "triton", "ascend", "intel_xpu")
|
||||||
assert (
|
assert (
|
||||||
prefill_backend in accepted_backends
|
prefill_backend in accepted_backends
|
||||||
and decode_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"
|
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_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")
|
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(
|
GRANITE_3_0_3B_A800M_INSTRUCT_WEIGHTS_PATH = os.path.join(
|
||||||
MODEL_WEIGHTS_DIR, "ibm-granite/granite-3.0-3b-a800m-instruct"
|
MODEL_WEIGHTS_DIR, "ibm-granite/granite-3.0-3b-a800m-instruct"
|
||||||
|
|||||||
@@ -0,0 +1,30 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_26B_A4B_IT_WEIGHTS_PATH
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma426BA4BIt(GSM8KAscendMixin, CustomTestCase):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-26B-A4B-it model on the GSM8K dataset is no less than 0.35.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-26B-A4B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_26B_A4B_IT_WEIGHTS_PATH
|
||||||
|
accuracy = 0.35
|
||||||
|
other_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--attention-backend",
|
||||||
|
"ascend",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--tp-size",
|
||||||
|
"2",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_31B_WEIGHTS_PATH
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma431B(GSM8KAscendMixin, CustomTestCase):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-31B-it model on the GSM8K dataset is no less than 0.70.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-31B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_31B_WEIGHTS_PATH
|
||||||
|
accuracy = 0.70
|
||||||
|
other_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--attention-backend",
|
||||||
|
"ascend",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--tp-size",
|
||||||
|
"2",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_E2B_WEIGHTS_PATH
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma4E2B(GSM8KAscendMixin, CustomTestCase):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-E2B-it model on the GSM8K dataset is no less than 0.05.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-E2B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_E2B_WEIGHTS_PATH
|
||||||
|
accuracy = 0.05
|
||||||
|
other_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--attention-backend",
|
||||||
|
"ascend",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_E4B_WEIGHTS_PATH
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma4E4B(GSM8KAscendMixin, CustomTestCase):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-E4B-it model on the GSM8K dataset is no less than 0.60.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-E4B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_E4B_WEIGHTS_PATH
|
||||||
|
accuracy = 0.60
|
||||||
|
other_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--attention-backend",
|
||||||
|
"ascend",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_26B_A4B_IT_WEIGHTS_PATH
|
||||||
|
from sglang.test.ascend.vlm_utils import TestVLMModels
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma426BA4BIt(TestVLMModels):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-26B-A4B-it model on the MMMU dataset is no less than 0.40.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-26B-A4B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_26B_A4B_IT_WEIGHTS_PATH
|
||||||
|
mmmu_accuracy = 0.40
|
||||||
|
|
||||||
|
def test_vlm_mmmu_benchmark(self):
|
||||||
|
self._run_vlm_mmmu_test()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_31B_WEIGHTS_PATH
|
||||||
|
from sglang.test.ascend.vlm_utils import TestVLMModels
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma431B(TestVLMModels):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-31B-it model on the MMMU dataset is no less than 0.50.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-31B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_31B_WEIGHTS_PATH
|
||||||
|
mmmu_accuracy = 0.50
|
||||||
|
|
||||||
|
def test_vlm_mmmu_benchmark(self):
|
||||||
|
self._run_vlm_mmmu_test()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_E2B_WEIGHTS_PATH
|
||||||
|
from sglang.test.ascend.vlm_utils import TestVLMModels
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma4E2B(TestVLMModels):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-E2B-it model on the MMMU dataset is no less than 0.15.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-E2B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_E2B_WEIGHTS_PATH
|
||||||
|
mmmu_accuracy = 0.15
|
||||||
|
|
||||||
|
def test_vlm_mmmu_benchmark(self):
|
||||||
|
self._run_vlm_mmmu_test()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ascend.test_ascend_utils import GEMMA_4_E4B_WEIGHTS_PATH
|
||||||
|
from sglang.test.ascend.vlm_utils import TestVLMModels
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma4E4B(TestVLMModels):
|
||||||
|
"""Testcase: Verify that the inference accuracy of the google/gemma-4-E4B-it model on the MMMU dataset is no less than 0.30.
|
||||||
|
|
||||||
|
[Test Category] Model
|
||||||
|
[Test Target] google/gemma-4-E4B-it
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = GEMMA_4_E4B_WEIGHTS_PATH
|
||||||
|
mmmu_accuracy = 0.30
|
||||||
|
|
||||||
|
def test_vlm_mmmu_benchmark(self):
|
||||||
|
self._run_vlm_mmmu_test()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user