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:
wenxuewuhd
2026-07-20 22:46:17 +03:00
committed by GitHub
co-authored by gemini-code-assist[bot] ronnie_zheng
parent ff6c755952
commit e856eae921
+34
View File
@@ -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(