Add fused_rope and for xpu (#25773)
This commit is contained in:
@@ -46,6 +46,9 @@ if _is_hip:
|
||||
fused_qk_rope_reshape_and_cache,
|
||||
)
|
||||
|
||||
if _is_xpu:
|
||||
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
|
||||
|
||||
|
||||
class RotaryEmbedding(MultiPlatformOp):
|
||||
"""Original rotary positional embedding."""
|
||||
@@ -420,14 +423,34 @@ class RotaryEmbedding(MultiPlatformOp):
|
||||
positions = torch.add(positions, offsets) if offsets is not None else positions
|
||||
|
||||
self._match_cos_sin_cache_dtype(query)
|
||||
return torch.ops.sgl_kernel.rotary_embedding(
|
||||
positions,
|
||||
query,
|
||||
key,
|
||||
self.head_size,
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style,
|
||||
)
|
||||
|
||||
# Fused_qk_rope only supports aligned head_size
|
||||
if self.head_size in [128, 256, 512]:
|
||||
num_tokens = positions.size(0)
|
||||
q_rope = query.view(num_tokens, -1, self.head_size)
|
||||
k_rope = key.view(num_tokens, -1, self.head_size)
|
||||
if self.head_size != self.rotary_dim:
|
||||
q_rope = q_rope[..., : self.rotary_dim]
|
||||
k_rope = k_rope[..., : self.rotary_dim]
|
||||
fused_qk_rope_with_cos_sin_cache_inplace(
|
||||
q_rope,
|
||||
k_rope,
|
||||
self.cos_sin_cache,
|
||||
positions,
|
||||
self.rotary_dim,
|
||||
self.is_neox_style,
|
||||
)
|
||||
return query, key
|
||||
else:
|
||||
# Use fallback kernel of 'rotary_embedding'
|
||||
return torch.ops.sgl_kernel.rotary_embedding(
|
||||
positions,
|
||||
query,
|
||||
key,
|
||||
self.head_size,
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style,
|
||||
)
|
||||
|
||||
|
||||
class LinearScalingRotaryEmbedding(RotaryEmbedding):
|
||||
|
||||
Reference in New Issue
Block a user