[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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
SWA_INT_MAX = 2147483647
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -207,6 +208,16 @@ class AscendAttnMaskBuilder:
|
|||||||
)
|
)
|
||||||
return attn_mask
|
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(
|
def _cp_allgather_and_save_kv_npu(
|
||||||
forward_batch, layer, k, v, cp_size, token_to_kv_pool
|
forward_batch, layer, k, v, cp_size, token_to_kv_pool
|
||||||
@@ -1066,12 +1077,64 @@ 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:
|
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
|
# 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_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
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:
|
||||||
|
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(
|
attn_out = attention_sinks_prefill_triton(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
@@ -1986,12 +2049,57 @@ 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:
|
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
|
# 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_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
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.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(
|
attn_out = attention_sinks_triton(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
|
|||||||
@@ -69,6 +69,20 @@ except ImportError:
|
|||||||
flashinfer_cutlass_fused_moe = None
|
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):
|
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||||
"""Unquantized method for embeddings."""
|
"""Unquantized method for embeddings."""
|
||||||
|
|
||||||
@@ -693,6 +707,11 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
|
|
||||||
hidden_states = swiglu_oai(layer, hidden_states)
|
hidden_states = swiglu_oai(layer, hidden_states)
|
||||||
elif self.moe_runner_config.activation == "silu":
|
elif self.moe_runner_config.activation == "silu":
|
||||||
|
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)
|
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
|
||||||
else:
|
else:
|
||||||
from sglang.srt.layers.activation import GeluAndMul
|
from sglang.srt.layers.activation import GeluAndMul
|
||||||
|
|||||||
Reference in New Issue
Block a user