From 0f8104ef15ed61140dc783bc69d5c7d22f7f22f2 Mon Sep 17 00:00:00 2001 From: Rohit Kumar Singh <9626333+SKRohit@users.noreply.github.com> Date: Fri, 29 May 2026 07:08:11 +0530 Subject: [PATCH] [XPU] Fix Device Assignment (#26257) --- python/sglang/srt/layers/rotary_embedding/base.py | 1 + python/sglang/srt/models/kimi_vl_moonvit.py | 6 +++--- python/sglang/srt/models/minicpmo.py | 4 ++-- python/sglang/srt/models/minicpmv.py | 10 +++++----- python/sglang/srt/models/transformers.py | 3 ++- .../srt/multimodal/processors/transformers_auto.py | 1 - 6 files changed, 13 insertions(+), 12 deletions(-) diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 2b13c1594..cdb92a894 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -419,6 +419,7 @@ class RotaryEmbedding(MultiPlatformOp): ), "fused_set_kv_buffer_arg is not supported for xpu implementation" 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, diff --git a/python/sglang/srt/models/kimi_vl_moonvit.py b/python/sglang/srt/models/kimi_vl_moonvit.py index 72f25b6b2..7b431987d 100644 --- a/python/sglang/srt/models/kimi_vl_moonvit.py +++ b/python/sglang/srt/models/kimi_vl_moonvit.py @@ -64,7 +64,7 @@ from sglang.srt.layers.conv import Conv2dLayer from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig -from sglang.srt.utils import add_prefix +from sglang.srt.utils import add_prefix, get_device @debug_kernel_api @@ -300,7 +300,7 @@ class Rope2DPosEmb(nn.Module): """ def __init__( - self, dim: int, max_height: int, max_width: int, theta_base=10000, device="cuda" + self, dim: int, max_height: int, max_width: int, theta_base=10000, device=None ): super().__init__() self.dim = dim @@ -308,7 +308,7 @@ class Rope2DPosEmb(nn.Module): self.max_height = max_height self.max_width = max_width self.theta_base = theta_base - self.device = device + self.device = device if device is not None else get_device() def extra_repr(self): return f"dim={self.dim}, max_height={self.max_height}, max_width={self.max_width}, theta_base={self.theta_base}" diff --git a/python/sglang/srt/models/minicpmo.py b/python/sglang/srt/models/minicpmo.py index fc03e29bf..80a3b1b89 100644 --- a/python/sglang/srt/models/minicpmo.py +++ b/python/sglang/srt/models/minicpmo.py @@ -54,7 +54,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.idefics2 import Idefics2VisionTransformer from sglang.srt.models.minicpmv import MiniCPMBaseModel, Resampler2_5 from sglang.srt.models.qwen2 import Qwen2ForCausalLM -from sglang.srt.utils import logger +from sglang.srt.utils import get_device, logger try: from transformers import LogitsWarper @@ -1514,7 +1514,7 @@ class MiniCPMO(MiniCPMBaseModel): prefix=prefix, ) - return resampler.to(device="cuda", dtype=torch.get_default_dtype()) + return resampler.to(device=get_device(), dtype=torch.get_default_dtype()) def pad_input_ids(self, input_ids: List[int], mm_input: MultimodalInputs): # Get all special token IDs diff --git a/python/sglang/srt/models/minicpmv.py b/python/sglang/srt/models/minicpmv.py index f1121a91a..3b8028594 100644 --- a/python/sglang/srt/models/minicpmv.py +++ b/python/sglang/srt/models/minicpmv.py @@ -68,7 +68,7 @@ from sglang.srt.models.minicpmv_vit import ( from sglang.srt.models.qwen2 import Qwen2Config, Qwen2ForCausalLM from sglang.srt.models.qwen3 import Qwen3Config, Qwen3ForCausalLM from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM -from sglang.srt.utils import add_prefix, flatten_nested_list +from sglang.srt.utils import add_prefix, flatten_nested_list, get_device RawImageType = Union[Image.Image, torch.Tensor] @@ -936,7 +936,7 @@ class MiniCPMV2_6(MiniCPMBaseModel): prefix=prefix, ) - return resampler.to(device="cuda", dtype=torch.get_default_dtype()) + return resampler.to(device=get_device(), dtype=torch.get_default_dtype()) def get_vision_embedding( self, @@ -1102,7 +1102,7 @@ class MiniCPMV4_0(MiniCPMBaseModel): prefix=prefix, ) - return resampler.to(device="cuda", dtype=torch.get_default_dtype()) + return resampler.to(device=get_device(), dtype=torch.get_default_dtype()) def get_vision_embedding( self, @@ -1272,7 +1272,7 @@ class MiniCPMV4_5(MiniCPMBaseModel): prefix=prefix, ) - return resampler.to(device="cuda", dtype=torch.get_default_dtype()) + return resampler.to(device=get_device(), dtype=torch.get_default_dtype()) def get_vision_embedding( self, @@ -1490,7 +1490,7 @@ class MiniCPMV4_6(MiniCPMBaseModel): quant_config=quant_config, prefix=prefix, ) - return merger.to(device="cuda", dtype=torch.get_default_dtype()) + return merger.to(device=get_device(), dtype=torch.get_default_dtype()) def get_vision_embedding( self, diff --git a/python/sglang/srt/models/transformers.py b/python/sglang/srt/models/transformers.py index 0c473d43d..79bbba865 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -68,6 +68,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.utils import AutoWeightsLoader, WeightsMapper from sglang.srt.server_args import get_global_server_args +from sglang.srt.utils import get_device from sglang.srt.utils.common import direct_register_custom_op from sglang.srt.utils.hf_transformers_utils import get_hf_text_config @@ -669,7 +670,7 @@ class TransformersBase(nn.Module): new_param = nn.Parameter( torch.empty_like( param.data, - device="cuda", + device=get_device(), ) ) setattr(module, name, new_param) diff --git a/python/sglang/srt/multimodal/processors/transformers_auto.py b/python/sglang/srt/multimodal/processors/transformers_auto.py index 579ae6e24..1c348b3c8 100644 --- a/python/sglang/srt/multimodal/processors/transformers_auto.py +++ b/python/sglang/srt/multimodal/processors/transformers_auto.py @@ -130,7 +130,6 @@ class TransformersAutoMultimodalProcessor(BaseMultimodalProcessor): modality_to_token_id = { Modality.IMAGE: self.mm_tokens.image_token_id, - Modality.MULTI_IMAGES: self.mm_tokens.image_token_id, Modality.VIDEO: self.mm_tokens.video_token_id, Modality.AUDIO: self.mm_tokens.audio_token_id, }