[Ascend] qwen optimization (#12078)
Co-authored-by: Even Zhou <even.y.zhou@outlook.com>
This commit is contained in:
@@ -44,6 +44,9 @@ logger = logging.getLogger(__name__)
|
||||
_is_cuda = is_cuda()
|
||||
_is_npu = is_npu()
|
||||
|
||||
if _is_npu:
|
||||
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import split_qkv_rmsnorm_rope
|
||||
|
||||
|
||||
class Qwen3Attention(nn.Module):
|
||||
def __init__(
|
||||
@@ -161,6 +164,33 @@ class Qwen3Attention(nn.Module):
|
||||
k = k_by_head.view(k.shape)
|
||||
return q, k
|
||||
|
||||
def forward_prepare_native(self, positions, hidden_states):
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
q, k = self._apply_qk_norm(q, k)
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
return q, k, v
|
||||
|
||||
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_norm.weight,
|
||||
self.k_norm.weight,
|
||||
self.q_size,
|
||||
self.kv_size,
|
||||
self.head_dim,
|
||||
self.q_norm.variance_epsilon,
|
||||
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,
|
||||
@@ -170,10 +200,16 @@ class Qwen3Attention(nn.Module):
|
||||
if get_global_server_args().rl_on_policy_target is not None:
|
||||
hidden_states = hidden_states.bfloat16()
|
||||
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
q, k = self._apply_qk_norm(q, k)
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
if not _is_npu:
|
||||
q, k, v = self.forward_prepare_native(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
)
|
||||
else:
|
||||
q, k, v = self.forward_prepare_npu(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
)
|
||||
|
||||
if get_global_server_args().rl_on_policy_target is not None:
|
||||
q = q.to(torch.bfloat16)
|
||||
|
||||
Reference in New Issue
Block a user