[NPU] DFlash Speculative Decoding Support NPU (#23122)

This commit is contained in:
gjsheu
2026-05-30 15:13:59 +08:00
committed by GitHub
parent fe4b29d391
commit 02aeed5387
4 changed files with 68 additions and 12 deletions
@@ -1675,6 +1675,12 @@ class AscendAttnBackend(AttentionBackend):
self.speculative_num_draft_tokens + query.shape[0],
self.speculative_num_draft_tokens,
)
if layer.attn_type == AttentionType.ENCODER_ONLY:
mask = None
sparse_mode = 0
else:
mask = self.mtp_mask
sparse_mode = 3
attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score(
query,
@@ -1685,15 +1691,16 @@ class AscendAttnBackend(AttentionBackend):
num_heads=layer.tp_q_head_num,
num_key_value_heads=layer.tp_k_head_num,
input_layout="TND",
atten_mask=self.mtp_mask,
atten_mask=mask,
scale=layer.scaling,
actual_seq_lengths=actual_seq_lengths,
actual_seq_lengths_kv=actual_seq_lengths_kv,
sparse_mode=3,
sparse_mode=sparse_mode,
)
attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim)
if (
not self.graph_mode
and forward_batch.num_token_non_padded_cpu is not None
and forward_batch.num_token_non_padded_cpu != num_token_padding
):
attn_output = torch.cat(
@@ -177,6 +177,14 @@ class NPUGraphRunner(CudaGraphRunner):
# In speculative decoding, these two fields are still needed.
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
if (
self.model_runner.spec_algorithm.is_dflash()
and self.model_runner.is_draft_worker
and forward_batch.input_embeds is not None
):
self.buffers.input_embeds[: self.raw_num_token].copy_(
forward_batch.input_embeds
)
if (
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get()
and forward_batch.mrope_positions is not None
@@ -207,10 +215,18 @@ class NPUGraphRunner(CudaGraphRunner):
if isinstance(output, LogitsProcessorOutput):
if self.is_dllm:
next_token_logits = None
full_logits = output.full_logits[: self.raw_num_token]
full_logits = (
output.full_logits[: self.raw_num_token]
if output.full_logits is not None
else None
)
else:
full_logits = None
next_token_logits = output.next_token_logits[: self.raw_num_token]
next_token_logits = (
output.next_token_logits[: self.raw_num_token]
if output.next_token_logits is not None
else None
)
return LogitsProcessorOutput(
next_token_logits=next_token_logits,
full_logits=full_logits,
+30 -3
View File
@@ -30,7 +30,11 @@ from sglang.srt.speculative.dflash_utils import (
can_dflash_slice_qkv_weight,
parse_dflash_draft_config,
)
from sglang.srt.utils import is_npu
_is_npu = is_npu()
if _is_npu:
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import split_qkv_rmsnorm_rope
logger = logging.getLogger(__name__)
@@ -118,6 +122,26 @@ class DFlashAttention(nn.Module):
attn_type=AttentionType.ENCODER_ONLY,
)
def forward_prepare_npu(self, positions, hidden_states):
qkv, _ = self.qkv_proj(hidden_states)
if self.attn.layer_id == 0:
self.rotary_emb.get_cos_sin_with_position(positions)
q, k, v = split_qkv_rmsnorm_rope(
qkv,
self.rotary_emb.position_sin,
self.rotary_emb.position_cos,
self.q_size,
self.kv_size,
self.head_dim,
eps=self.q_norm.variance_epsilon,
q_weight=self.q_norm.weight,
k_weight=self.k_norm.weight,
q_bias=getattr(self.q_norm, "bias", None),
k_bias=getattr(self.k_norm, "bias", None),
)
return q, k, v
def forward(
self,
positions: torch.Tensor,
@@ -125,9 +149,12 @@ class DFlashAttention(nn.Module):
forward_batch: ForwardBatch,
) -> torch.Tensor:
qkv, _ = self.qkv_proj(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = apply_qk_norm(q, k, self.q_norm, self.k_norm, self.head_dim)
q, k = self.rotary_emb(positions, q, k)
if _is_npu:
q, k, v = self.forward_prepare_npu(positions, hidden_states)
else:
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
q, k = apply_qk_norm(q, k, self.q_norm, self.k_norm, self.head_dim)
q, k = self.rotary_emb(positions, q, k)
attn_output = self.attn(q, k, v, forward_batch)
output, _ = self.o_proj(attn_output)
return output
+11 -5
View File
@@ -29,7 +29,10 @@ from sglang.srt.speculative.dflash_utils import (
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
from sglang.srt.utils import is_cuda
from sglang.srt.utils import is_cuda, is_npu
_is_npu = is_npu()
logger = logging.getLogger(__name__)
@@ -98,7 +101,7 @@ class DFlashWorker:
draft_server_args = deepcopy(server_args)
draft_server_args.skip_tokenizer_init = True
draft_backend = draft_server_args.speculative_draft_attention_backend
supported_draft_backends = ("flashinfer", "fa3", "fa4", "triton")
supported_draft_backends = ("flashinfer", "fa3", "fa4", "triton", "ascend")
if draft_backend is None:
draft_backend, _ = draft_server_args.get_attention_backends()
if draft_backend is None:
@@ -1016,9 +1019,12 @@ class DFlashWorker:
) -> None:
for layer in self.draft_model.layers:
attn = layer.self_attn
k, v = attn.kv_proj_only(ctx_hidden)
k = attn.apply_k_norm(k)
k = attn.apply_k_rope(ctx_positions, k)
if _is_npu:
_, k, v = attn.forward_prepare_npu(ctx_positions, ctx_hidden)
else:
k, v = attn.kv_proj_only(ctx_hidden)
k = attn.apply_k_norm(k)
k = attn.apply_k_rope(ctx_positions, k)
k = k.view(-1, attn.num_kv_heads, attn.head_dim)
v = v.view(-1, attn.num_kv_heads, attn.head_dim)
self.draft_model_runner.token_to_kv_pool.set_kv_buffer(