[NPU] Enhance accuracy for model Step3_5 from 0 to 88% (#24582)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user