[NPU]GLM-4.7-Flash optimize with fused kernels (#29509)

This commit is contained in:
Estrella-xx
2026-06-30 19:22:19 +08:00
committed by GitHub
parent 2f730e299f
commit c6a7c98ae4
2 changed files with 37 additions and 15 deletions
@@ -184,40 +184,62 @@ def forward_mla_prepare_npu(
else:
q_lora = None
if m.q_lora_rank is not None:
q, latent_cache = (
get_attn_tp_context()
.fetch_qkv_latent()
.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim],
dim=-1,
)
)
k_nope = latent_cache[..., : m.kv_lora_rank]
q = m.q_a_layernorm(q)
qkv_latent = get_attn_tp_context().fetch_qkv_latent()
if (
_use_ag_after_qlora
and layer_scatter_modes.layer_input_mode == ScatterMode.SCATTERED
and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL
):
q, latent_cache = qkv_latent.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim],
dim=-1,
)
k_nope = latent_cache[..., : m.kv_lora_rank]
q = m.q_a_layernorm(q)
q = scattered_to_tp_attn_full(q, forward_batch)
latent_cache = scattered_to_tp_attn_full(latent_cache, forward_batch)
k_nope = m.kv_a_layernorm(k_nope)
k_nope = m.kv_a_layernorm(k_nope).unsqueeze(1)
k_pe = latent_cache[..., m.kv_lora_rank :].unsqueeze(1)
else:
if qkv_latent.shape[0] < 65536 and not dsa_use_prefill_cp(
forward_batch
):
q, k_nope, k_pe = fused_split_qk_norm(
qkv_latent,
m.q_a_layernorm,
m.kv_a_layernorm,
m.q_lora_rank,
m.kv_lora_rank,
m.qk_rope_head_dim,
eps=m.q_a_layernorm.variance_epsilon,
)
else:
q, latent_cache = qkv_latent.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim],
dim=-1,
)
k_nope = latent_cache[..., : m.kv_lora_rank]
q = m.q_a_layernorm(q)
k_nope = m.kv_a_layernorm(k_nope).unsqueeze(1)
k_pe = latent_cache[..., m.kv_lora_rank :].unsqueeze(1)
# q_lora needed by indexer
if m.use_dsa:
q_lora = q
k_nope = k_nope.unsqueeze(1)
q = m.q_b_proj(q)[0].view(-1, m.num_local_heads, m.qk_head_dim)
else:
q = m.q_proj(hidden_states)[0].view(-1, m.num_local_heads, m.qk_head_dim)
latent_cache = m.kv_a_proj_with_mqa(hidden_states)[0]
k_nope = latent_cache[..., : m.kv_lora_rank]
k_nope = m.kv_a_layernorm(k_nope).unsqueeze(1)
k_pe = latent_cache[..., m.kv_lora_rank :].unsqueeze(1)
q_nope, q_pe = q.split([m.qk_nope_head_dim, m.qk_rope_head_dim], dim=-1)
k_pe = latent_cache[..., m.kv_lora_rank :].unsqueeze(1)
q_nope_out = torch.bmm(q_nope.transpose(0, 1), m.w_kc)
@@ -84,7 +84,7 @@ def fused_topk_npu(
group_select_mode=(1 if use_grouped_topk else 0),
renorm=0,
# 1 for sigmoid, 0 for softmax
norm_type=(0 if topk_config.scoring_func == "softmax" else 1),
norm_type=1,
routed_scaling_factor=(
1 if renormalize else topk_config.routed_scaling_factor
),