From 7cbe5648295e363b151d9c0221c4b1adb14511e8 Mon Sep 17 00:00:00 2001 From: ashwini rathi Date: Fri, 28 Aug 2026 04:47:24 +0530 Subject: [PATCH] [Fix][XPU/ROCm/NPU] Defer sgl_kernel.quantization import in expert_pack (#36529) --- .../sglang/srt/layers/quantization/expert_pack.py | 5 ++++- python/sglang/test/runners.py | 15 ++++++++++++--- test/registered/xpu/test_xpu_classification.py | 6 +++++- test/registered/xpu/test_xpu_embedding.py | 7 ++++++- test/registered/xpu/test_xpu_rerank.py | 13 +++---------- test/registered/xpu/test_xpu_reward.py | 7 ++++++- 6 files changed, 36 insertions(+), 17 deletions(-) diff --git a/python/sglang/srt/layers/quantization/expert_pack.py b/python/sglang/srt/layers/quantization/expert_pack.py index 2bc8a4408..7ab5e52da 100644 --- a/python/sglang/srt/layers/quantization/expert_pack.py +++ b/python/sglang/srt/layers/quantization/expert_pack.py @@ -7,7 +7,6 @@ from typing import Optional import torch import torch.nn.functional as F -from sgl_kernel.quantization import ggml_moe_a8_vec from sglang.kernels.ops.moe.expert_pack_mxfp4 import ( mxfp4_matvec, @@ -157,6 +156,10 @@ class ExpertPackMoEMethod(FusedMoEMethodBase): weight_type: int, output_size: int, ) -> torch.Tensor: + # sgl_kernel.quantization is CUDA/MUSA-only; keep the import local so + # this module stays importable on other devices (see gguf.py:44-70). + from sgl_kernel.quantization import ggml_moe_a8_vec + return ggml_moe_a8_vec( inputs, weights, diff --git a/python/sglang/test/runners.py b/python/sglang/test/runners.py index 09de69ee6..03a33e20f 100644 --- a/python/sglang/test/runners.py +++ b/python/sglang/test/runners.py @@ -424,16 +424,25 @@ class HFRunner: f"before producing output" ) - def terminate(self): + def _stop_model_proc(self): + # Fire-and-forget terminate() leaves the child holding the accelerator + # during teardown; a follow-on SRTRunner on the same device can then + # deadlock in driver init (observed on Intel XPU B580). self.model_proc.terminate() + self.model_proc.join(timeout=30) + if self.model_proc.is_alive(): + self.model_proc.kill() + self.model_proc.join() self.in_queue = self.out_queue = None + def terminate(self): + self._stop_model_proc() + def __enter__(self): return self def __exit__(self, exc_type, exc_value, traceback): - self.model_proc.terminate() - self.in_queue = self.out_queue = None + self._stop_model_proc() @staticmethod def forward_generation_raw( diff --git a/test/registered/xpu/test_xpu_classification.py b/test/registered/xpu/test_xpu_classification.py index 1fd2bf367..0018a4c80 100644 --- a/test/registered/xpu/test_xpu_classification.py +++ b/test/registered/xpu/test_xpu_classification.py @@ -4,6 +4,7 @@ Usage: python3 -m unittest test_xpu_classification.TestXPUClassification """ +import gc import multiprocessing as mp import unittest @@ -11,7 +12,7 @@ import torch from sglang.test.ci.ci_register import register_xpu_ci from sglang.test.runners import HFRunner, SRTRunner -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, empty_gpu_cache register_xpu_ci(est_time=120, suite="stage-b-test-1-gpu-xpu") @@ -60,6 +61,7 @@ class TestXPUClassification(CustomTestCase): model_type="embedding", attention_backend="intel_xpu", trust_remote_code=True, + mem_fraction_static=0.55, ) as srt_runner: srt_logits = srt_runner.forward(PROMPTS).embed_logits @@ -72,6 +74,8 @@ class TestXPUClassification(CustomTestCase): def test_classification_logits(self): hf_probs = self._hf_probs() + gc.collect() + empty_gpu_cache() srt_probs = self._srt_probs() self.assertEqual(len(hf_probs), len(PROMPTS)) diff --git a/test/registered/xpu/test_xpu_embedding.py b/test/registered/xpu/test_xpu_embedding.py index b5959bad0..2ff0a11a5 100644 --- a/test/registered/xpu/test_xpu_embedding.py +++ b/test/registered/xpu/test_xpu_embedding.py @@ -5,6 +5,7 @@ Usage: python3 -m unittest test_xpu_embedding.TestXPUEmbedding """ +import gc import multiprocessing as mp import unittest from typing import Optional @@ -14,7 +15,7 @@ from transformers import AutoConfig, AutoTokenizer from sglang.test.ci.ci_register import register_xpu_ci from sglang.test.runners import DEFAULT_PROMPTS, HFRunner, SRTRunner -from sglang.test.test_utils import CustomTestCase, get_similarities +from sglang.test.test_utils import CustomTestCase, empty_gpu_cache, get_similarities register_xpu_ci(est_time=180, suite="stage-b-test-1-gpu-xpu") @@ -65,12 +66,16 @@ class TestXPUEmbedding(CustomTestCase): ) as hf_runner: hf_outputs = hf_runner.forward(truncated_prompts) + gc.collect() + empty_gpu_cache() + with SRTRunner( model_path, tp_size=tp_size, torch_dtype=torch_dtype, model_type="embedding", attention_backend="intel_xpu", + mem_fraction_static=0.55, json_model_override_args=( {"matryoshka_dimensions": [matryoshka_dim]} if matryoshka_dim is not None diff --git a/test/registered/xpu/test_xpu_rerank.py b/test/registered/xpu/test_xpu_rerank.py index f1518788f..afbe8ea99 100644 --- a/test/registered/xpu/test_xpu_rerank.py +++ b/test/registered/xpu/test_xpu_rerank.py @@ -21,7 +21,7 @@ from jinja2.sandbox import ImmutableSandboxedEnvironment from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.test.ci.ci_register import register_xpu_ci from sglang.test.runners import TEST_RERANK_QUERY_DOCS, HFRunner, SRTRunner -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, empty_gpu_cache def _xpu_total_gib() -> float: @@ -30,13 +30,6 @@ def _xpu_total_gib() -> float: return torch.xpu.get_device_properties(0).total_memory / (1024**3) -def _xpu_free_cache() -> None: - gc.collect() - if torch.xpu.is_available(): - torch.xpu.empty_cache() - torch.xpu.synchronize() - - # fp32+Triton fits on B60 (22GiB) but hangs on B580 (~12GiB). _LARGE_XPU_VRAM_GIB = 20.0 _HAS_LARGE_XPU = _xpu_total_gib() >= _LARGE_XPU_VRAM_GIB @@ -214,8 +207,8 @@ class TestXPUCrossEncoderRerank(CustomTestCase): ) as hf_runner: hf_scores = hf_runner.forward(prompts).scores - # HFRunner leaks a ZMQ context on shutdown; free VRAM before SRT starts. - _xpu_free_cache() + gc.collect() + empty_gpu_cache() with SRTRunner( model_path, diff --git a/test/registered/xpu/test_xpu_reward.py b/test/registered/xpu/test_xpu_reward.py index ba4e04a71..710406ed2 100644 --- a/test/registered/xpu/test_xpu_reward.py +++ b/test/registered/xpu/test_xpu_reward.py @@ -5,6 +5,7 @@ Usage: python3 -m unittest test_xpu_reward.TestXPUReward """ +import gc import multiprocessing as mp import unittest @@ -12,7 +13,7 @@ import torch from sglang.test.ci.ci_register import register_xpu_ci from sglang.test.runners import HFRunner, SRTRunner -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, empty_gpu_cache register_xpu_ci(est_time=60, suite="stage-b-test-1-gpu-xpu") @@ -53,12 +54,16 @@ class TestXPUReward(CustomTestCase): ) as hf_runner: hf_outputs = hf_runner.forward(convs) + gc.collect() + empty_gpu_cache() + with SRTRunner( model_path, tp_size=tp_size, torch_dtype=torch_dtype, model_type="reward", attention_backend="intel_xpu", + mem_fraction_static=0.55, ) as srt_runner: prompts = srt_runner.tokenizer.apply_chat_template( convs,