[NPU][eagle3] support qwen eagle3 on NPU (#14820)
This commit is contained in:
@@ -889,102 +889,163 @@ class AscendAttnBackend(AttentionBackend):
|
||||
layer, forward_batch.out_cache_loc, k, v
|
||||
)
|
||||
|
||||
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
if is_fia_nz():
|
||||
k_rope_cache = _reshape_kv_for_fia_nz(
|
||||
k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size
|
||||
)
|
||||
c_kv_cache = _reshape_kv_for_fia_nz(
|
||||
c_kv, layer.tp_v_head_num, self.kv_lora_rank, self.page_size
|
||||
)
|
||||
else:
|
||||
k_rope_cache = k_rope.view(
|
||||
-1, layer.tp_k_head_num, self.page_size, self.qk_rope_head_dim
|
||||
)
|
||||
c_kv_cache = c_kv.view(
|
||||
-1, layer.tp_v_head_num, self.page_size, self.kv_lora_rank
|
||||
)
|
||||
if not self.use_mla:
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||
layer.layer_id
|
||||
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
||||
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous()
|
||||
if not self.graph_mode:
|
||||
num_token_padding = query.shape[0]
|
||||
query = query[: forward_batch.num_token_non_padded_cpu]
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list
|
||||
else:
|
||||
actual_seq_lengths_kv = (
|
||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||
)
|
||||
if forward_batch.forward_mode.is_draft_extend():
|
||||
actual_seq_lengths = (
|
||||
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
|
||||
)
|
||||
else:
|
||||
actual_seq_lengths = np.arange(
|
||||
self.speculative_num_draft_tokens,
|
||||
self.speculative_num_draft_tokens + query.shape[0],
|
||||
self.speculative_num_draft_tokens,
|
||||
)
|
||||
|
||||
q_nope = q.view(-1, layer.tp_q_head_num, self.kv_lora_rank).contiguous()
|
||||
q_rope = q_rope.view(-1, layer.tp_q_head_num, self.qk_rope_head_dim)
|
||||
if not self.graph_mode:
|
||||
num_token_padding = q.shape[0]
|
||||
q_nope = q_nope[: forward_batch.num_token_non_padded_cpu]
|
||||
q_rope = q_rope[: forward_batch.num_token_non_padded_cpu]
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list
|
||||
attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score(
|
||||
query,
|
||||
k_cache,
|
||||
v_cache,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="TND",
|
||||
atten_mask=self.mtp_mask,
|
||||
scale=layer.scaling,
|
||||
actual_seq_lengths=actual_seq_lengths,
|
||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
||||
sparse_mode=3,
|
||||
)
|
||||
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 != num_token_padding
|
||||
):
|
||||
attn_output = torch.cat(
|
||||
[
|
||||
attn_output,
|
||||
attn_output.new_zeros(
|
||||
num_token_padding - forward_batch.num_token_non_padded_cpu,
|
||||
*attn_output.shape[1:],
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
return attn_output
|
||||
else:
|
||||
actual_seq_lengths_kv = (
|
||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||
)
|
||||
if forward_batch.forward_mode.is_draft_extend():
|
||||
actual_seq_lengths = (
|
||||
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
|
||||
)
|
||||
else:
|
||||
actual_seq_lengths = np.arange(
|
||||
self.speculative_num_draft_tokens,
|
||||
self.speculative_num_draft_tokens + q_nope.shape[0],
|
||||
self.speculative_num_draft_tokens,
|
||||
)
|
||||
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
if is_fia_nz():
|
||||
k_rope_cache = _reshape_kv_for_fia_nz(
|
||||
k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size
|
||||
)
|
||||
c_kv_cache = _reshape_kv_for_fia_nz(
|
||||
c_kv, layer.tp_v_head_num, self.kv_lora_rank, self.page_size
|
||||
)
|
||||
else:
|
||||
k_rope_cache = k_rope.view(
|
||||
-1, layer.tp_k_head_num, self.page_size, self.qk_rope_head_dim
|
||||
)
|
||||
c_kv_cache = c_kv.view(
|
||||
-1, layer.tp_v_head_num, self.page_size, self.kv_lora_rank
|
||||
)
|
||||
|
||||
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
||||
q_nope,
|
||||
c_kv_cache,
|
||||
c_kv_cache,
|
||||
query_rope=q_rope,
|
||||
key_rope=k_rope_cache,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="TND",
|
||||
scale=layer.scaling,
|
||||
antiquant_mode=0,
|
||||
antiquant_scale=None,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
sparse_mode=3,
|
||||
atten_mask=self.mtp_mask,
|
||||
actual_seq_lengths=actual_seq_lengths,
|
||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
||||
)
|
||||
attn_output = torch.empty_like(q_nope, dtype=q.dtype, device=q.device)
|
||||
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
|
||||
torch_npu.npu_fused_infer_attention_score.out(
|
||||
q_nope,
|
||||
c_kv_cache,
|
||||
c_kv_cache,
|
||||
query_rope=q_rope,
|
||||
key_rope=k_rope_cache,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="TND",
|
||||
scale=layer.scaling,
|
||||
antiquant_mode=0,
|
||||
antiquant_scale=None,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
sparse_mode=3,
|
||||
atten_mask=self.mtp_mask,
|
||||
actual_seq_lengths=actual_seq_lengths,
|
||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
||||
workspace=workspace,
|
||||
out=[attn_output, softmax_lse],
|
||||
)
|
||||
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 != num_token_padding
|
||||
):
|
||||
attn_output = torch.cat(
|
||||
[
|
||||
attn_output,
|
||||
attn_output.new_zeros(
|
||||
num_token_padding - attn_output.shape[0], *attn_output.shape[1:]
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
q_nope = q.view(-1, layer.tp_q_head_num, self.kv_lora_rank).contiguous()
|
||||
q_rope = q_rope.view(-1, layer.tp_q_head_num, self.qk_rope_head_dim)
|
||||
if not self.graph_mode:
|
||||
num_token_padding = q.shape[0]
|
||||
q_nope = q_nope[: forward_batch.num_token_non_padded_cpu]
|
||||
q_rope = q_rope[: forward_batch.num_token_non_padded_cpu]
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list
|
||||
else:
|
||||
actual_seq_lengths_kv = (
|
||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||
)
|
||||
if forward_batch.forward_mode.is_draft_extend():
|
||||
actual_seq_lengths = (
|
||||
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
|
||||
)
|
||||
else:
|
||||
actual_seq_lengths = np.arange(
|
||||
self.speculative_num_draft_tokens,
|
||||
self.speculative_num_draft_tokens + q_nope.shape[0],
|
||||
self.speculative_num_draft_tokens,
|
||||
)
|
||||
|
||||
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
||||
q_nope,
|
||||
c_kv_cache,
|
||||
c_kv_cache,
|
||||
query_rope=q_rope,
|
||||
key_rope=k_rope_cache,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="TND",
|
||||
scale=layer.scaling,
|
||||
antiquant_mode=0,
|
||||
antiquant_scale=None,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
sparse_mode=3,
|
||||
atten_mask=self.mtp_mask,
|
||||
actual_seq_lengths=actual_seq_lengths,
|
||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
||||
)
|
||||
return attn_output
|
||||
attn_output = torch.empty_like(q_nope, dtype=q.dtype, device=q.device)
|
||||
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
|
||||
torch_npu.npu_fused_infer_attention_score.out(
|
||||
q_nope,
|
||||
c_kv_cache,
|
||||
c_kv_cache,
|
||||
query_rope=q_rope,
|
||||
key_rope=k_rope_cache,
|
||||
num_heads=layer.tp_q_head_num,
|
||||
num_key_value_heads=layer.tp_k_head_num,
|
||||
input_layout="TND",
|
||||
scale=layer.scaling,
|
||||
antiquant_mode=0,
|
||||
antiquant_scale=None,
|
||||
block_table=self.forward_metadata.block_tables,
|
||||
block_size=self.page_size,
|
||||
sparse_mode=3,
|
||||
atten_mask=self.mtp_mask,
|
||||
actual_seq_lengths=actual_seq_lengths,
|
||||
actual_seq_lengths_kv=actual_seq_lengths_kv,
|
||||
workspace=workspace,
|
||||
out=[attn_output, softmax_lse],
|
||||
)
|
||||
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 != num_token_padding
|
||||
):
|
||||
attn_output = torch.cat(
|
||||
[
|
||||
attn_output,
|
||||
attn_output.new_zeros(
|
||||
num_token_padding - attn_output.shape[0],
|
||||
*attn_output.shape[1:],
|
||||
),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
return attn_output
|
||||
|
||||
def forward_decode_graph(
|
||||
self,
|
||||
|
||||
+3
@@ -37,6 +37,9 @@ class EAGLEDraftExtendNpuGraphRunner(EAGLEDraftExtendCudaGraphRunner):
|
||||
def _create_graph(self):
|
||||
return torch.npu.NPUGraph()
|
||||
|
||||
def _cache_loc_dtype(self):
|
||||
return torch.int32
|
||||
|
||||
def _capture_init(self, run_once_fn):
|
||||
for _ in range(2):
|
||||
torch.npu.synchronize()
|
||||
|
||||
+36
-3
@@ -17,11 +17,13 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Dict, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import is_deepseek_nsa
|
||||
from sglang.srt.configs.model_config import AttentionArch, is_deepseek_nsa
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
|
||||
EAGLEDraftCudaGraphRunner,
|
||||
@@ -46,6 +48,19 @@ if is_npu():
|
||||
class EAGLEDraftNpuGraphRunner(EAGLEDraftCudaGraphRunner):
|
||||
def __init__(self, eagle_worker: EAGLEWorker):
|
||||
super().__init__(eagle_worker)
|
||||
self.update_attr_name = None
|
||||
self.update_attr_type = None
|
||||
self._init_arch_map()
|
||||
|
||||
def _init_arch_map(self):
|
||||
self.attr_name: Dict[str, str] = {
|
||||
AttentionArch.MLA: "actual_seq_lengths_kv",
|
||||
AttentionArch.MHA: "context_lens",
|
||||
}
|
||||
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
|
||||
AttentionArch.MLA: [],
|
||||
AttentionArch.MHA: torch.Tensor(),
|
||||
}
|
||||
|
||||
def _create_graph(self):
|
||||
return torch.npu.NPUGraph()
|
||||
@@ -63,12 +78,27 @@ class EAGLEDraftNpuGraphRunner(EAGLEDraftCudaGraphRunner):
|
||||
out = run_once_fn()
|
||||
return out
|
||||
|
||||
def _get_update_attr_name(self, model_runner):
|
||||
if self.bs < get_attention_tp_size():
|
||||
return self.attr_name[AttentionArch.MLA]
|
||||
return self.attr_name[model_runner.model_config.attention_arch]
|
||||
|
||||
def _get_update_attr_type(self, model_runner):
|
||||
if self.bs < get_attention_tp_size():
|
||||
return self.attr_type[AttentionArch.MLA]
|
||||
return self.attr_type[model_runner.model_config.attention_arch]
|
||||
|
||||
def _replay_update(self, seq_lens):
|
||||
if isinstance(self.update_attr_type, torch.Tensor):
|
||||
seq_lens = torch.from_numpy(np.array(seq_lens).astype(np.int32))
|
||||
|
||||
self.graphs[self.bs].update(
|
||||
cpu_update_input=[{"actual_seq_lengths_kv": seq_lens}]
|
||||
cpu_update_input=[{self.update_attr_name: seq_lens}]
|
||||
)
|
||||
|
||||
def _replay(self, forward_batch: ForwardBatch):
|
||||
self.update_attr_name = self._get_update_attr_name(self.model_runner)
|
||||
self.update_attr_type = self._get_update_attr_type(self.model_runner)
|
||||
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
|
||||
seq_lens = forward_batch.seq_lens_cpu.tolist() + [0] * (
|
||||
self.bs - self.raw_bs
|
||||
@@ -79,3 +109,6 @@ class EAGLEDraftNpuGraphRunner(EAGLEDraftCudaGraphRunner):
|
||||
thread.join()
|
||||
else:
|
||||
self.graphs[self.bs].replay()
|
||||
|
||||
def _cache_loc_dtype(self):
|
||||
return torch.int32
|
||||
|
||||
@@ -77,13 +77,19 @@ class NPUGraphRunner(CudaGraphRunner):
|
||||
out = run_once_fn()
|
||||
return out
|
||||
|
||||
def _get_update_attr_name(self, model_runner):
|
||||
if self.bs < get_attention_tp_size():
|
||||
def _get_update_attr_name(self, model_runner, forward_batch):
|
||||
if (
|
||||
self.bs < get_attention_tp_size()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
):
|
||||
return self.attr_name[AttentionArch.MLA]
|
||||
return self.attr_name[model_runner.model_config.attention_arch]
|
||||
|
||||
def _get_update_attr_type(self, model_runner):
|
||||
if self.bs < get_attention_tp_size():
|
||||
def _get_update_attr_type(self, model_runner, forward_batch):
|
||||
if (
|
||||
self.bs < get_attention_tp_size()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
):
|
||||
return self.attr_type[AttentionArch.MLA]
|
||||
return self.attr_type[model_runner.model_config.attention_arch]
|
||||
|
||||
@@ -139,8 +145,12 @@ class NPUGraphRunner(CudaGraphRunner):
|
||||
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
|
||||
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
|
||||
|
||||
self.update_attr_name = self._get_update_attr_name(self.model_runner)
|
||||
self.update_attr_type = self._get_update_attr_type(self.model_runner)
|
||||
self.update_attr_name = self._get_update_attr_name(
|
||||
self.model_runner, forward_batch
|
||||
)
|
||||
self.update_attr_type = self._get_update_attr_type(
|
||||
self.model_runner, forward_batch
|
||||
)
|
||||
# Replay
|
||||
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
|
||||
@@ -89,7 +89,7 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
||||
if self.store_dtype != self.dtype:
|
||||
cache_k = cache_k.view(self.store_dtype)
|
||||
cache_v = cache_v.view(self.store_dtype)
|
||||
|
||||
loc = loc.to(torch.int32)
|
||||
torch_npu._npu_reshape_and_cache(
|
||||
key=cache_k,
|
||||
value=cache_v,
|
||||
|
||||
Reference in New Issue
Block a user