[feat]Ascend NPU Gemma-3-12b and Gemma-3-27b support (#8909)
This commit is contained in:
@@ -1876,7 +1876,7 @@ def rotate_half(x):
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def apply_rotary_pos_emb(
|
||||
def apply_rotary_pos_emb_native(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
@@ -1899,6 +1899,33 @@ def apply_rotary_pos_emb(
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
def apply_rotary_pos_emb_npu(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
unsqueeze_dim=1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if q.shape[1] != 128:
|
||||
return apply_rotary_pos_emb_native(q, k, cos, sin, unsqueeze_dim)
|
||||
cos = cos.unsqueeze(unsqueeze_dim)
|
||||
cos = torch.transpose(cos, 1, 2)
|
||||
sin = sin.unsqueeze(unsqueeze_dim)
|
||||
sin = torch.transpose(sin, 1, 2)
|
||||
q = torch.transpose(q, 1, 2)
|
||||
k = torch.transpose(k, 1, 2)
|
||||
q_embed, k_embed = torch_npu.npu_apply_rotary_pos_emb(q, k, cos, sin)
|
||||
q_embed = torch.transpose(q_embed, 1, 2)
|
||||
k_embed = torch.transpose(k_embed, 1, 2)
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
if _is_npu:
|
||||
apply_rotary_pos_emb = apply_rotary_pos_emb_npu
|
||||
else:
|
||||
apply_rotary_pos_emb = apply_rotary_pos_emb_native
|
||||
|
||||
|
||||
def get_rope_cpu(
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
|
||||
Reference in New Issue
Block a user