diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index cf0de8f8a..b16fa2b01 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -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, ) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_torch_native_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_torch_native_backend.py index 34bbfc67f..702e01d64 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_torch_native_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_torch_native_backend.py @@ -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) diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index 9c2f9c25d..644acdc90 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -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 diff --git a/python/sglang/srt/models/gemma4_mm.py b/python/sglang/srt/models/gemma4_mm.py index 3629782ca..1858a5d59 100644 --- a/python/sglang/srt/models/gemma4_mm.py +++ b/python/sglang/srt/models/gemma4_mm.py @@ -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. diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index fd3e819db..98d68f101 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/python/sglang/test/ascend/test_ascend_utils.py b/python/sglang/test/ascend/test_ascend_utils.py index 02127c886..c711d6dad 100644 --- a/python/sglang/test/ascend/test_ascend_utils.py +++ b/python/sglang/test/ascend/test_ascend_utils.py @@ -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" diff --git a/test/manual/ascend/llm_models/test_npu_gemma_4_26b_a4b_it_llm.py b/test/manual/ascend/llm_models/test_npu_gemma_4_26b_a4b_it_llm.py new file mode 100644 index 000000000..62fb5f570 --- /dev/null +++ b/test/manual/ascend/llm_models/test_npu_gemma_4_26b_a4b_it_llm.py @@ -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() diff --git a/test/manual/ascend/llm_models/test_npu_gemma_4_31b_llm.py b/test/manual/ascend/llm_models/test_npu_gemma_4_31b_llm.py new file mode 100644 index 000000000..f3302cfc3 --- /dev/null +++ b/test/manual/ascend/llm_models/test_npu_gemma_4_31b_llm.py @@ -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() diff --git a/test/manual/ascend/llm_models/test_npu_gemma_4_e2b_llm.py b/test/manual/ascend/llm_models/test_npu_gemma_4_e2b_llm.py new file mode 100644 index 000000000..9bc6b7ed4 --- /dev/null +++ b/test/manual/ascend/llm_models/test_npu_gemma_4_e2b_llm.py @@ -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() diff --git a/test/manual/ascend/llm_models/test_npu_gemma_4_e4b_llm.py b/test/manual/ascend/llm_models/test_npu_gemma_4_e4b_llm.py new file mode 100644 index 000000000..ac23f7f9f --- /dev/null +++ b/test/manual/ascend/llm_models/test_npu_gemma_4_e4b_llm.py @@ -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() diff --git a/test/manual/ascend/vlm_models/test_npu_gemma_4_26b_a4b_it.py b/test/manual/ascend/vlm_models/test_npu_gemma_4_26b_a4b_it.py new file mode 100644 index 000000000..27f20d2a9 --- /dev/null +++ b/test/manual/ascend/vlm_models/test_npu_gemma_4_26b_a4b_it.py @@ -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() diff --git a/test/manual/ascend/vlm_models/test_npu_gemma_4_31b.py b/test/manual/ascend/vlm_models/test_npu_gemma_4_31b.py new file mode 100644 index 000000000..be84fdce6 --- /dev/null +++ b/test/manual/ascend/vlm_models/test_npu_gemma_4_31b.py @@ -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() diff --git a/test/manual/ascend/vlm_models/test_npu_gemma_4_e2b.py b/test/manual/ascend/vlm_models/test_npu_gemma_4_e2b.py new file mode 100644 index 000000000..30c5849ad --- /dev/null +++ b/test/manual/ascend/vlm_models/test_npu_gemma_4_e2b.py @@ -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() diff --git a/test/manual/ascend/vlm_models/test_npu_gemma_4_e4b.py b/test/manual/ascend/vlm_models/test_npu_gemma_4_e4b.py new file mode 100644 index 000000000..2f48777fe --- /dev/null +++ b/test/manual/ascend/vlm_models/test_npu_gemma_4_e4b.py @@ -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()