[NPU]DeepSeek-V3.2 support npu mlaprolog (#15381)

Co-authored-by: Zhengda Qin <zhengdqin@gmail.com>
Co-authored-by: richhuan <huan_rz@qq.com>
This commit is contained in:
lawtherWu
2026-01-26 20:42:37 +08:00
committed by GitHub
parent 738b1ac988
commit b56366f827
5 changed files with 195 additions and 61 deletions

View File

@@ -713,6 +713,9 @@ class AscendAttnBackend(AttentionBackend):
sinks: Optional[torch.Tensor] = None,
slopes: Optional[torch.Tensor] = None,
):
if is_mla_preprocess_enabled():
# MLAPO and MLAPROLOG do save kv_cache
save_kv_cache = False
if topk_indices is not None:
return self.forward_sparse(
q,
@@ -730,9 +733,6 @@ class AscendAttnBackend(AttentionBackend):
or forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2()
):
if is_mla_preprocess_enabled():
save_kv_cache = False
return self.forward_mtp(
q,
k,

View File

@@ -1,3 +1,4 @@
import re
from functools import lru_cache
from typing import TYPE_CHECKING, Optional
@@ -72,6 +73,7 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
num_local_heads,
qk_nope_head_dim,
qk_rope_head_dim,
v_head_dim,
quant_config: Optional["QuantizationConfig"] = None,
):
super().__init__()
@@ -92,6 +94,10 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
self.qk_nope_head_dim = qk_nope_head_dim # 128
self.qk_rope_head_dim = qk_rope_head_dim # 64
self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
self.v_head_dim = v_head_dim
self.q_b_proj_weight_scale = self.q_b_proj.weight_scale.view(1, -1).to(
torch.float
)
def preprocess_weights(self, hidden_states):
self.dummy = torch.zeros(
@@ -230,6 +236,15 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
self.num_local_heads * (self.qk_nope_head_dim + self.qk_rope_head_dim)
)
def mlaprolog_preprocess_weight(self):
self.qkv_a_proj.weight.data = self.qkv_a_proj.weight.data.transpose(0, 1)
qkv_a_proj_weight_q = self.qkv_a_proj.weight.data[:, : self.q_lora_rank].clone()
qkv_a_proj_weight_kv = self.qkv_a_proj.weight.data[
:, self.q_lora_rank :
].clone()
self.q_a_proj_weight = npu_format_cast(qkv_a_proj_weight_q)
self.kv_a_proj_weight = npu_format_cast(qkv_a_proj_weight_kv)
def get_sin_cos(self, positions):
cos_sin = self.rotary_emb.cos_sin_cache[positions]
cos, sin = cos_sin.chunk(2, dim=-1)
@@ -399,17 +414,64 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
positions,
)
def forward_mlaprolog(self, positions, hidden_states, forward_batch):
if not self.has_preprocess_weights:
self.mlaprolog_preprocess_weight()
self.has_preprocess_weights = True
self.cos, self.sin = self.get_sin_cos(positions)
k_cache, v_cache, slot_mapping = self.get_kv_cache_and_cache_idx(forward_batch)
mla_prolog_input_args = {
"token_x": hidden_states,
"weight_dq": self.q_a_proj_weight,
"weight_uq_qr": self.q_b_proj.weight,
"weight_uk": self.w_kc,
"weight_dkv_kr": self.kv_a_proj_weight,
"rmsnorm_gamma_cq": self.q_a_layernorm.weight,
"rmsnorm_gamma_ckv": self.kv_a_layernorm.weight,
"rope_sin": self.sin,
"rope_cos": self.cos,
"kv_cache": k_cache,
"kr_cache": v_cache,
"cache_index": slot_mapping.to(dtype=torch.int64),
"dequant_scale_w_uq_qr": self.q_b_proj_weight_scale,
"rmsnorm_epsilon_cq": self.q_a_layernorm.variance_epsilon,
"rmsnorm_epsilon_ckv": self.kv_a_layernorm.variance_epsilon,
"cache_mode": "PA_BSND",
"query_norm_flag": True,
"weight_quant_mode": 1, # 0:no quant; 1:uq_qr: quant; 2: weight_dq,weight_uq_qr,weight_dkv_kr: quant
}
q_nope, q_pe, dequant_scale_q_nope, qr, dequant_q_norm = (
torch.ops.custom.npu_mla_prolog_v3(**mla_prolog_input_args)
)
dequant_q_norm = dequant_q_norm.view(hidden_states.shape[0])
return (
q_pe,
v_cache,
q_nope,
k_cache,
qr,
forward_batch,
positions,
dequant_q_norm,
)
def forward(self, positions, hidden_states, forward_batch, zero_allocator):
assert self.quant_config and self.quant_config.get_name() == "modelslim"
# assert self.quant_config and self.quant_config.get_name() == "modelslim"
# route by `qkv_a_proj` quant type as MTP layers can be unquantized
_is_w8a8 = (
hasattr(self.qkv_a_proj.quant_method, "quant_config")
and self.qkv_a_proj.quant_method.quant_config.get_name() == "modelslim"
)
# with the mlaprolog enabled, the kv_b_proj layers are unquantized
_is_mlaprolog = hasattr(self.quant_config, "ignore") and any(
re.fullmatch(r".*kv_b_proj", l) for l in self.quant_config.ignore
)
if _is_w8a8:
return self.forward_mlapo(
positions, hidden_states, forward_batch, zero_allocator
)
elif _is_mlaprolog:
return self.forward_mlaprolog(positions, hidden_states, forward_batch)
else:
return self.forward_absorb_prepare_npu_rms_norm_cache(
positions, hidden_states, forward_batch, zero_allocator

View File

@@ -1,3 +1,4 @@
import re
from typing import TYPE_CHECKING
import torch
@@ -281,61 +282,25 @@ def forward_dsa_prepare_npu(
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
):
dynamic_scale = None
if is_mla_preprocess_enabled() and forward_batch.forward_mode.is_decode():
if not hasattr(m, "mla_preprocess"):
m.mla_preprocess = NPUFusedMLAPreprocess(
m.fused_qkv_a_proj_with_mqa,
m.q_a_layernorm,
m.kv_a_layernorm,
m.q_b_proj,
m.w_kc,
m.rotary_emb,
m.layer_id,
m.num_local_heads,
m.qk_nope_head_dim,
m.qk_rope_head_dim,
m.quant_config,
)
if m.alt_stream is not None:
mla_event = torch.npu.Event()
mla_event.record()
with torch.npu.stream(m.alt_stream):
# alt stream waits for the completion of the event on the main stream to ensure data dependency is complete
torch.npu.current_stream().wait_event(mla_event)
(
q_pe,
k_pe,
q_nope_out,
k_nope,
forward_batch,
zero_allocator,
positions,
) = m.mla_preprocess.forward(
positions, hidden_states, forward_batch, zero_allocator
)
fused_qkv_a_proj_out = m.fused_qkv_a_proj_with_mqa(hidden_states)[0]
q, _ = fused_qkv_a_proj_out.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim], dim=-1
)
q_lora = m.q_a_layernorm(q)
torch.npu.current_stream().wait_stream(m.alt_stream)
else:
(
q_pe,
k_pe,
q_nope_out,
k_nope,
forward_batch,
zero_allocator,
positions,
) = m.mla_preprocess.forward(
positions, hidden_states, forward_batch, zero_allocator
)
fused_qkv_a_proj_out = m.fused_qkv_a_proj_with_mqa(hidden_states)[0]
q, _ = fused_qkv_a_proj_out.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim], dim=-1
)
q_lora = m.q_a_layernorm(q)
(
q_pe,
k_pe,
q_nope_out,
k_nope,
q_lora,
forward_batch,
zero_allocator,
positions,
dynamic_scale,
) = npu_mla_preprocess(
m,
hidden_states,
positions,
forward_batch,
zero_allocator,
)
else:
fused_qkv_a_proj_out = m.fused_qkv_a_proj_with_mqa(hidden_states)[0]
q, latent_cache = fused_qkv_a_proj_out.split(
@@ -384,7 +349,7 @@ def forward_dsa_prepare_npu(
)
topk_indices = m.indexer(
hidden_states, q_lora, positions, forward_batch, m.layer_id
hidden_states, q_lora, positions, forward_batch, m.layer_id, dynamic_scale
)
return (
@@ -451,4 +416,100 @@ def forward_dsa_core_npu(
return output
def npu_mla_preprocess(
m: "DeepseekV2AttentionMLA",
hidden_states: torch.Tensor,
positions: torch.Tensor,
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
):
dynamic_scale = None
if not hasattr(m, "mla_preprocess"):
m.mla_preprocess = NPUFusedMLAPreprocess(
m.fused_qkv_a_proj_with_mqa,
m.q_a_layernorm,
m.kv_a_layernorm,
m.q_b_proj,
m.w_kc,
m.rotary_emb,
m.layer_id,
m.num_local_heads,
m.qk_nope_head_dim,
m.qk_rope_head_dim,
m.v_head_dim,
m.quant_config,
)
# mlaprolog does not require additional calculation of q_lora
_is_mlaprolog = hasattr(m.quant_config, "ignore") and any(
re.fullmatch(r".*kv_b_proj", l) for l in m.quant_config.ignore
)
if _is_mlaprolog:
(
q_pe,
k_pe,
q_nope_out,
k_nope,
q_lora,
forward_batch,
positions,
dynamic_scale,
) = m.mla_preprocess.forward(
positions, hidden_states, forward_batch, zero_allocator
)
else:
if m.alt_stream is not None:
mla_event = torch.npu.Event()
mla_event.record()
with torch.npu.stream(m.alt_stream):
# alt stream waits for the completion of the event on the main stream to ensure data dependency is complete
torch.npu.current_stream().wait_event(mla_event)
(
q_pe,
k_pe,
q_nope_out,
k_nope,
forward_batch,
zero_allocator,
positions,
) = m.mla_preprocess.forward(
positions, hidden_states, forward_batch, zero_allocator
)
fused_qkv_a_proj_out = m.fused_qkv_a_proj_with_mqa(hidden_states)[0]
q, _ = fused_qkv_a_proj_out.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim], dim=-1
)
q_lora = m.q_a_layernorm(q)
torch.npu.current_stream().wait_event(m.alt_stream)
else:
(
q_pe,
k_pe,
q_nope_out,
k_nope,
forward_batch,
zero_allocator,
positions,
) = m.mla_preprocess.forward(
positions, hidden_states, forward_batch, zero_allocator
)
fused_qkv_a_proj_out = m.fused_qkv_a_proj_with_mqa(hidden_states)[0]
q, _ = fused_qkv_a_proj_out.split(
[m.q_lora_rank, m.kv_lora_rank + m.qk_rope_head_dim], dim=-1
)
q_lora = m.q_a_layernorm(q)
return (
q_pe,
k_pe,
q_nope_out,
k_nope,
q_lora,
forward_batch,
zero_allocator,
positions,
dynamic_scale,
)
# endregion

View File

@@ -93,8 +93,14 @@ class NPUW8A8Int8DynamicLinearMethod(_NPULinearMethodBase):
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
original_dtype = x.dtype
quant_out, dynamic_scale = torch.ops.npu.npu_dynamic_quant(x)
if isinstance(x, tuple):
"""dynamic_scale is calculated in malprolog kernel"""
original_dtype = torch.bfloat16
quant_out, dynamic_scale = x
else:
original_dtype = x.dtype
quant_out, dynamic_scale = torch.ops.npu.npu_dynamic_quant(x)
return torch.ops.npu.npu_quant_matmul(
quant_out,
layer.weight,

View File

@@ -1113,6 +1113,7 @@ class Indexer(MultiPlatformOp):
positions: torch.Tensor,
forward_batch: ForwardBatch,
layer_id: int,
dynamic_scale: torch.Tensor = None,
) -> torch.Tensor:
if forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int is None:
actual_seq_lengths_kv = forward_batch.attn_backend.forward_metadata.seq_lens
@@ -1136,6 +1137,9 @@ class Indexer(MultiPlatformOp):
if self.alt_stream is not None:
self.alt_stream.wait_stream(torch.npu.current_stream())
with torch.npu.stream(self.alt_stream):
q_lora = (
(q_lora, dynamic_scale) if dynamic_scale is not None else q_lora
)
q = self.wq_b(q_lora)[
0
] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128]
@@ -1154,6 +1158,7 @@ class Indexer(MultiPlatformOp):
q.record_stream(self.alt_stream)
q_rope_event = self.alt_stream.record_event()
else:
q_lora = (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora
q = self.wq_b(q_lora)[0] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128]
q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128]
q_pe, q_nope = torch.split(