[NPU] Enhance accuracy for model Step3_5 from 0 to 88% (#24582)

This commit is contained in:
McZyWu
2026-05-29 11:29:30 +08:00
committed by GitHub
parent 36d0a6e08e
commit b1173c8c14
2 changed files with 155 additions and 28 deletions
@@ -45,6 +45,7 @@ def _reshape_kv_for_fia_nz(
logger = logging.getLogger(__name__)
SWA_INT_MAX = 2147483647
@dataclass
@@ -207,6 +208,16 @@ class AscendAttnMaskBuilder:
)
return attn_mask
def get_swa_mask(self, seq_lens: torch.Tensor, s2: int, left_context=512):
if seq_lens.dim() == 1:
seq_lens = seq_lens.unsqueeze(1)
b = seq_lens.size(0)
device = seq_lens.device
indices = torch.arange(s2, device=device).unsqueeze(0).expand(b, -1)
start_indices = torch.clamp(seq_lens - left_context, min=0)
mask = (indices < start_indices) | (indices >= seq_lens)
return mask.unsqueeze(1).to(self.device, non_blocking=True)
def _cp_allgather_and_save_kv_npu(
forward_batch, layer, k, v, cp_size, token_to_kv_pool
@@ -1066,25 +1077,77 @@ 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:
if sinks is not None or (
self.is_hybrid_swa and layer.sliding_window_size != -1
):
# Use SWA block tables if hybrid SWA is enabled for this layer
if self.is_hybrid_swa and layer.sliding_window_size != -1:
block_tables = self.forward_metadata.block_tables_swa
else:
block_tables = self.forward_metadata.block_tables
attn_out = attention_sinks_prefill_triton(
q,
k_cache,
v_cache,
sinks,
self.forward_metadata.extend_seq_lens,
block_tables,
self.forward_metadata.seq_lens,
layer.scaling,
layer.sliding_window_size,
layer.tp_q_head_num,
layer.tp_k_head_num,
)
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 SWA_INT_MAX
),
next_tokens=(
0 if layer.sliding_window_size != -1 else SWA_INT_MAX
),
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,
)
attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
else:
attn_out = attention_sinks_prefill_triton(
q,
k_cache,
v_cache,
sinks,
self.forward_metadata.extend_seq_lens,
block_tables,
self.forward_metadata.seq_lens,
layer.scaling,
layer.sliding_window_size,
layer.tp_q_head_num,
layer.tp_k_head_num,
)
return attn_out
if is_cp_mode:
@@ -1986,24 +2049,69 @@ 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:
if sinks is not None or (
self.is_hybrid_swa and layer.sliding_window_size != -1
):
# Use SWA block tables if hybrid SWA is enabled for this layer
if self.is_hybrid_swa and layer.sliding_window_size != -1:
block_tables = self.forward_metadata.block_tables_swa
else:
block_tables = self.forward_metadata.block_tables
attn_out = attention_sinks_triton(
q,
k_cache,
v_cache,
sinks,
block_tables,
self.forward_metadata.seq_lens,
layer.scaling,
layer.sliding_window_size,
layer.tp_q_head_num,
layer.tp_k_head_num,
)
if self.use_fia:
if self.forward_metadata.seq_lens_cpu_int is None:
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
else:
actual_seq_len_kv = (
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
)
block_size = self.page_size
max_model_len = block_tables.shape[-1] * block_size
swa_mask = self.ascend_attn_mask_builder.get_swa_mask(
self.forward_metadata.seq_lens,
max_model_len,
layer.sliding_window_size,
)
attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2(
q.view(
forward_batch.batch_size,
-1,
layer.tp_q_head_num,
layer.qk_head_dim,
),
k_cache.view(
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
),
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_size=block_size,
atten_mask=(
swa_mask if layer.sliding_window_size != -1 else None
),
sparse_mode=4 if layer.sliding_window_size != -1 else 0,
softmax_scale=layer.scaling,
block_table=block_tables,
actual_seq_qlen=[1] * len(self.forward_metadata.seq_lens),
actual_seq_kvlen=actual_seq_len_kv,
learnable_sink=sinks,
)
attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
else:
attn_out = attention_sinks_triton(
q,
k_cache,
v_cache,
sinks,
block_tables,
self.forward_metadata.seq_lens,
layer.scaling,
layer.sliding_window_size,
layer.tp_q_head_num,
layer.tp_k_head_num,
)
return attn_out
if self.use_fia:
@@ -69,6 +69,20 @@ except ImportError:
flashinfer_cutlass_fused_moe = None
def swiglustep_and_mul(x: torch.Tensor, limit: float = 7.0) -> torch.Tensor:
"""Out-variant of swiglustep activation.
Writes into `out`:
silu(x[:d]).clamp(max=limit) * x[d:].clamp(-limit, limit)
"""
gate, up = x.chunk(2, dim=-1)
gate = F.silu(gate)
gate = gate.clamp(max=limit)
up = up.clamp(min=-limit, max=limit)
out = gate * up
return out
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
"""Unquantized method for embeddings."""
@@ -693,7 +707,12 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
hidden_states = swiglu_oai(layer, hidden_states)
elif self.moe_runner_config.activation == "silu":
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
if self.moe_runner_config.gemm1_clamp_limit is not None:
hidden_states = swiglustep_and_mul(
hidden_states, self.moe_runner_config.gemm1_clamp_limit
)
else:
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
else:
from sglang.srt.layers.activation import GeluAndMul