From f0303fd07eb0fee895632fc5fcc0d53a95fdd5d7 Mon Sep 17 00:00:00 2001 From: Polisetty V R K Jyothendra Varma Date: Mon, 30 Mar 2026 11:05:59 +0530 Subject: [PATCH] [Intel GPU] Enable DeepSeek R1 inference on XPU (#18461) Signed-off-by: P V R K Jyothendra Varma --- .../tuning_fused_moe_triton.py | 13 ++++-- .../quantization/tuning_block_wise_kernel.py | 42 ++++++++++--------- .../layers/moe/token_dispatcher/standard.py | 12 ++++-- .../deepseek_common/deepseek_weight_loader.py | 3 +- .../srt/models/deepseek_common/utils.py | 2 + python/sglang/srt/models/deepseek_v2.py | 2 + 6 files changed, 46 insertions(+), 28 deletions(-) diff --git a/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py b/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py index 34aa83b38..4cc397f65 100644 --- a/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py +++ b/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py @@ -32,9 +32,10 @@ from sglang.srt.server_args import ( ServerArgs, set_global_server_args_for_scheduler, ) -from sglang.srt.utils import is_hip +from sglang.srt.utils import get_device, is_hip, is_xpu _is_hip = is_hip() +_is_xpu = is_xpu() def benchmark_config( @@ -236,8 +237,8 @@ def benchmark_config( class BenchmarkWorker: def __init__(self, seed: int, server_args: ServerArgs) -> None: - torch.set_default_device("cuda") - torch.cuda.manual_seed_all(0) + torch.set_default_device(get_device()) + torch.get_device_module().manual_seed_all(0) self.seed = seed # Get the device ID to allocate tensors and kernels # on the respective GPU. @@ -330,7 +331,11 @@ class BenchmarkWorker: ) -> Dict[str, int]: best_config = None best_time = float("inf") - with torch.cuda.device(self.device_id) if is_hip() else nullcontext(): + with ( + torch.get_device_module().device(self.device_id) + if _is_xpu or _is_hip + else nullcontext() + ): for config in tqdm(search_space): try: kernel_time = benchmark_config( diff --git a/benchmark/kernels/quantization/tuning_block_wise_kernel.py b/benchmark/kernels/quantization/tuning_block_wise_kernel.py index 396b14a75..9e4368043 100644 --- a/benchmark/kernels/quantization/tuning_block_wise_kernel.py +++ b/benchmark/kernels/quantization/tuning_block_wise_kernel.py @@ -31,7 +31,13 @@ from sglang.srt.layers.quantization.fp8_kernel import ( _w8a8_block_fp8_matmul_unrolledx4, ) from sglang.srt.layers.quantization.int8_kernel import _w8a8_block_int8_matmul -from sglang.srt.utils import get_device_core_count, get_device_name, is_hip +from sglang.srt.utils import ( + get_device, + get_device_core_count, + get_device_count, + get_device_name, + is_hip, +) _is_hip = is_hip() @@ -221,18 +227,18 @@ def benchmark_config( def run(): w8a8_block_matmul(A, B, As, Bs, block_size, config, out_dtype) - torch.cuda.synchronize() + torch.get_device_module().synchronize() # JIT complication & warmup for _ in range(5): run() - torch.cuda.synchronize() + torch.get_device_module().synchronize() - start_event = torch.cuda.Event(enable_timing=True) - end_event = torch.cuda.Event(enable_timing=True) + start_event = torch.get_device_module().Event(enable_timing=True) + end_event = torch.get_device_module().Event(enable_timing=True) latencies: List[float] = [] for i in range(num_iters): - torch.cuda.synchronize() + torch.get_device_module().synchronize() start_event.record() run() end_event.record() @@ -244,6 +250,7 @@ def benchmark_config( def tune(M, N, K, block_size, out_dtype, search_space, input_type): factor_for_scale = 1e-2 + device = get_device() if input_type == "fp8": fp8_info = torch.finfo( @@ -252,14 +259,14 @@ def tune(M, N, K, block_size, out_dtype, search_space, input_type): fp8_max, fp8_min = fp8_info.max, fp8_info.min A_fp32 = ( - (torch.rand(M, K, dtype=torch.float32, device="cuda") - 0.5) * 2 * fp8_max + (torch.rand(M, K, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max ) A = A_fp32.clamp(min=fp8_min, max=fp8_max).to( torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn ) B_fp32 = ( - (torch.rand(N, K, dtype=torch.float32, device="cuda") - 0.5) * 2 * fp8_max + (torch.rand(N, K, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max ) B = B_fp32.clamp(min=fp8_min, max=fp8_max).to( torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn @@ -269,12 +276,12 @@ def tune(M, N, K, block_size, out_dtype, search_space, input_type): int8_max, int8_min = int8_info.max, int8_info.min A_fp32 = ( - (torch.rand(M, K, dtype=torch.float32, device="cuda") - 0.5) * 2 * int8_max + (torch.rand(M, K, dtype=torch.float32, device=device) - 0.5) * 2 * int8_max ) A = A_fp32.clamp(min=int8_min, max=int8_max).to(torch.int8) B_fp32 = ( - (torch.rand(N, K, dtype=torch.float32, device="cuda") - 0.5) * 2 * int8_max + (torch.rand(N, K, dtype=torch.float32, device=device) - 0.5) * 2 * int8_max ) B = B_fp32.clamp(min=int8_min, max=int8_max).to(torch.int8) @@ -282,9 +289,9 @@ def tune(M, N, K, block_size, out_dtype, search_space, input_type): n_tiles = (N + block_n - 1) // block_n k_tiles = (K + block_k - 1) // block_k - As = torch.rand(M, k_tiles, dtype=torch.float32, device="cuda") * factor_for_scale + As = torch.rand(M, k_tiles, dtype=torch.float32, device=device) * factor_for_scale Bs = ( - torch.rand(n_tiles, k_tiles, dtype=torch.float32, device="cuda") + torch.rand(n_tiles, k_tiles, dtype=torch.float32, device=device) * factor_for_scale ) @@ -351,11 +358,6 @@ def save_configs( lock.release() -def get_available_gpu_count(): - """Get the number of available GPUs.""" - return torch.cuda.device_count() - - def tune_on_gpu(args_dict): """Run tuning on a specific GPU.""" gpu_id = args_dict["gpu_id"] @@ -364,7 +366,7 @@ def tune_on_gpu(args_dict): args = args_dict["args"] lock = args_dict["lock"] - torch.cuda.set_device(gpu_id) + torch.get_device_module().set_device(gpu_id) print(f"Starting tuning on GPU {gpu_id} with batch sizes {batch_sizes}") block_n = args.block_n @@ -415,12 +417,12 @@ def distribute_batch_sizes(batch_sizes, num_gpus): def main(args): print(args) - num_gpus = get_available_gpu_count() + num_gpus = get_device_count() if num_gpus == 0: raise RuntimeError("No GPU available for tuning") print(f"Found {num_gpus} GPUs for parallel tuning") - torch.cuda.init() + torch.get_device_module().init() if args.batch_size is None: batch_sizes = [ diff --git a/python/sglang/srt/layers/moe/token_dispatcher/standard.py b/python/sglang/srt/layers/moe/token_dispatcher/standard.py index b77c19f83..ef78ca9f3 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/standard.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/standard.py @@ -30,7 +30,12 @@ from sglang.srt.layers.moe.utils import ( get_moe_runner_backend, should_use_flashinfer_cutlass_moe_fp4_allgather, ) -from sglang.srt.utils.common import get_bool_env_var, is_hip, is_sm120_supported +from sglang.srt.utils.common import ( + get_bool_env_var, + get_device, + is_hip, + is_sm120_supported, +) _is_hip = is_hip() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -149,15 +154,16 @@ class StandardDispatcher(BaseDispatcher): and TopKOutputChecker.format_is_standard(topk_output) ): if self.local_expert_mapping is None: + device = get_device() self.local_expert_mapping = torch.full( - (self.num_experts,), -1, dtype=torch.int32, device="cuda" + (self.num_experts,), -1, dtype=torch.int32, device=device ) self.local_expert_mapping[ self.moe_ep_rank * self.num_local_routed_experts : (self.moe_ep_rank + 1) * self.num_local_routed_experts ] = torch.arange( - 0, self.num_local_routed_experts, dtype=torch.int32, device="cuda" + 0, self.num_local_routed_experts, dtype=torch.int32, device=device ) if self.num_local_shared_experts > 0: diff --git a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py index 12ce382ed..b72e8290d 100644 --- a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py +++ b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py @@ -50,6 +50,7 @@ from sglang.srt.models.deepseek_common.utils import ( _is_fp8_fnuz, _is_hip, _is_npu, + _is_xpu, _use_aiter_gfx95, awq_dequantize_func, enable_nextn_moe_bf16_cast_to_fp8, @@ -497,7 +498,7 @@ class DeepseekV2WeightLoaderMixin: ) if ( - _is_cuda + (_is_cuda or _is_xpu) and weight_block_size[0] == 128 and weight_block_size[1] == 128 ): diff --git a/python/sglang/srt/models/deepseek_common/utils.py b/python/sglang/srt/models/deepseek_common/utils.py index 73f26c2b0..a5579d528 100644 --- a/python/sglang/srt/models/deepseek_common/utils.py +++ b/python/sglang/srt/models/deepseek_common/utils.py @@ -31,6 +31,7 @@ from sglang.srt.utils import ( is_hip, is_npu, is_nvidia_cublas_version_ge_12_9, + is_xpu, ) _is_hip = is_hip() @@ -40,6 +41,7 @@ _is_fp8_fnuz = is_fp8_fnuz() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _is_cpu_amx_available = cpu_has_amx_support() _is_cpu = is_cpu() +_is_xpu = is_xpu() _device_sm = get_device_sm() _is_gfx95_supported = is_gfx95_supported() _use_aiter_gfx95 = _use_aiter and _is_gfx95_supported diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 0d88f541a..0a2aa07e0 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -137,6 +137,7 @@ from sglang.srt.models.deepseek_common.utils import ( _is_gfx95_supported, _is_hip, _is_npu, + _is_xpu, _use_aiter, _use_aiter_gfx95, ) @@ -677,6 +678,7 @@ class DeepseekV2MoE(nn.Module): ) if ( not _is_cuda + and not _is_xpu and not _use_aiter or isinstance(self.experts.quant_method, KTEPWrapperMethod) ):