From dc1e46ec8ffe8124f6fda52990ca5c63e67489a9 Mon Sep 17 00:00:00 2001 From: gaopengff Date: Thu, 11 Jun 2026 12:25:55 +0800 Subject: [PATCH] [Intel GPU]Add sycl mrope pass for xpu device (#27646) --- .../srt/layers/rotary_embedding/mrope.py | 26 +++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 775fcc0c7..10b0d99c6 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -19,10 +19,17 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_linear_ramp_mask, ) from sglang.srt.server_args import get_global_server_args -from sglang.srt.utils import cpu_has_amx_support, is_cuda, is_npu, support_triton +from sglang.srt.utils import ( + cpu_has_amx_support, + is_cuda, + is_npu, + is_xpu, + support_triton, +) _is_cuda = is_cuda() _is_npu = is_npu() +_is_xpu = is_xpu() _is_cpu_amx_available = cpu_has_amx_support() if _is_cuda: @@ -31,6 +38,8 @@ if _is_cuda: if _is_npu: import torch_npu +if _is_xpu: + from sgl_kernel import multimodal_rotary_embedding import triton import triton.language as tl @@ -377,7 +386,20 @@ class MRotaryEmbedding(RotaryEmbedding): ) -> Tuple[torch.Tensor, torch.Tensor]: assert positions.ndim in (1, 2) if positions.ndim == 2 and self.mrope_section: - return self.forward_triton(positions, query, key) + multimodal_rotary_embedding( + query, + key, + self.cos_sin_cache, + positions, + self.mrope_section, + self.head_size, + self.rotary_dim, + self.mrope_interleaved, + self.mrope_interleaved_glm, + self.is_neox_style, + self.axis_map, + ) + return query, key return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) @staticmethod