use sgl_kernel_npu rmsrope accelerate llada2 (#27127)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot]
ronnie_zheng
parent
ff6c755952
commit
e856eae921
@@ -96,6 +96,15 @@ logger = logging.getLogger(__name__)
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
split_qkv_rmsnorm_rope_pos_cache_half_npu = None
|
||||||
|
if _is_npu:
|
||||||
|
try:
|
||||||
|
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope_pos_cache_half_npu import (
|
||||||
|
split_qkv_rmsnorm_rope_pos_cache_half_npu,
|
||||||
|
)
|
||||||
|
except (ImportError, OSError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class LLaDA2MoeMLP(nn.Module):
|
class LLaDA2MoeMLP(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -521,6 +530,31 @@ class LLaDA2MoeAttention(nn.Module):
|
|||||||
if hidden_states.shape[0] == 0:
|
if hidden_states.shape[0] == 0:
|
||||||
return hidden_states
|
return hidden_states
|
||||||
qkv, _ = self.query_key_value(hidden_states)
|
qkv, _ = self.query_key_value(hidden_states)
|
||||||
|
if _is_npu and split_qkv_rmsnorm_rope_pos_cache_half_npu is not None:
|
||||||
|
q, k, v = split_qkv_rmsnorm_rope_pos_cache_half_npu(
|
||||||
|
qkv,
|
||||||
|
positions,
|
||||||
|
self.rotary_emb.cos_sin_cache,
|
||||||
|
self.q_size,
|
||||||
|
self.kv_size,
|
||||||
|
self.head_dim,
|
||||||
|
eps=self.query_layernorm.variance_epsilon if self.use_qk_norm else None,
|
||||||
|
q_weight=self.query_layernorm.weight if self.use_qk_norm else None,
|
||||||
|
k_weight=self.key_layernorm.weight if self.use_qk_norm else None,
|
||||||
|
q_bias=(
|
||||||
|
getattr(self.query_layernorm, "bias", None)
|
||||||
|
if self.use_qk_norm
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
k_bias=(
|
||||||
|
getattr(self.key_layernorm, "bias", None)
|
||||||
|
if self.use_qk_norm
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
rope_dim=self.rotary_dim,
|
||||||
|
)
|
||||||
|
can_fuse_set_kv = False
|
||||||
|
else:
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
if self.use_qk_norm:
|
if self.use_qk_norm:
|
||||||
q, k = apply_qk_norm(
|
q, k = apply_qk_norm(
|
||||||
|
|||||||
Reference in New Issue
Block a user