From 0e934cd88141a72a80d1077837f5107056f16b08 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Tue, 8 Sep 2026 05:34:04 -0700 Subject: [PATCH 01/10] [XPU] Use sgl-kernel-xpu fused sampling kernels On XPU, `sampling_backend` resolved to "pytorch" because `_sampling_backend_default` falls back whenever flashinfer is unavailable. The torch fallback sorts the full vocabulary per row on every decode step, so top-k/top-p sampling cost grows linearly with batch size. sgl-kernel-xpu already ships the sampling ops, and they mirror the flashinfer API exactly (`top_k_renorm_prob`, `top_p_renorm_prob`, `top_k_top_p_sampling_from_probs(..., filter_apply_order=...)`, `min_p_sampling_from_probs`), so the existing branch is reused as-is. Adds an "xpu" sampling backend, defaulted on `--device xpu` and guarded so an explicit `--sampling-backend` still wins. Only the default, non-deterministic top-k/top-p/min-p path changes. Greedy and temperature-only requests are untouched, and `--enable-deterministic-inference` still resolves to "pytorch" since the XPU kernels take an RNG generator rather than the per-request sampling seed that batch-invariant sampling requires. --- python/sglang/srt/arg_groups/choices.py | 2 +- python/sglang/srt/arg_groups/platform_hook.py | 6 ++++ python/sglang/srt/layers/sampler.py | 11 ++++--- .../cpu/test_server_args_backend.py | 29 ++++++++++++++++++- 4 files changed, 42 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/arg_groups/choices.py b/python/sglang/srt/arg_groups/choices.py index 1ac5b914cb7e..828a5cf0ceb0 100644 --- a/python/sglang/srt/arg_groups/choices.py +++ b/python/sglang/srt/arg_groups/choices.py @@ -134,7 +134,7 @@ GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"] -SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend"} +SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend", "intel_xpu"} MOE_RUNNER_BACKEND_CHOICES = [ "auto", diff --git a/python/sglang/srt/arg_groups/platform_hook.py b/python/sglang/srt/arg_groups/platform_hook.py index 67007b73427f..7fdffb6c2aef 100644 --- a/python/sglang/srt/arg_groups/platform_hook.py +++ b/python/sglang/srt/arg_groups/platform_hook.py @@ -127,6 +127,12 @@ def handle_symm_mem_device_support(server_args: Any): def handle_xpu_backends(server_args: Any): cfg = resolving_view(server_args) if cfg.device == "xpu": + if cfg.sampling_backend is None: + declare_resolution( + server_args, + "_handle_xpu_backends", + sampling_backend="intel_xpu", + ) # Decode graph is opt-in on XPU: unless the user explicitly set # --cuda-graph-backend-decode (or --cuda-graph-config), keep it # disabled so the default startup doesn't require graph capture. diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index abf59fd1c0a6..e5dcadd541c3 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -31,6 +31,7 @@ is_hip, is_musa, is_npu, + is_xpu, ) if is_cuda(): @@ -43,7 +44,7 @@ top_p_renorm_prob, ) -if is_musa(): +if is_musa() or is_xpu(): from sgl_kernel import ( min_p_sampling_from_probs, top_k_renorm_prob, @@ -73,7 +74,7 @@ SYNC_TOKEN_IDS_ACROSS_TP = get_bool_env_var("SYNC_TOKEN_IDS_ACROSS_TP") SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB") _CUSTOM_SAMPLER_FACTORIES: Dict[str, Callable[[], "Sampler"]] = {} -_BUILT_IN_SAMPLING_BACKENDS = {"flashinfer", "pytorch", "ascend"} +_BUILT_IN_SAMPLING_BACKENDS = {"flashinfer", "pytorch", "ascend", "intel_xpu"} def _trace_e2e_sampler(stage: str, **fields) -> None: @@ -354,9 +355,11 @@ def _sample_from_probs( ) else: backend = get_exec().kernel.sampling_backend - if backend == "flashinfer": + # sgl-kernel-xpu mirrors the flashinfer sampling API, so both share + # this branch. + if backend in ("flashinfer", "intel_xpu"): assert sampling_info.sampling_seed is None, ( - "Sampling seed is not supported for flashinfer backend" + f"Sampling seed is not supported for {backend} backend" ) if sampling_info.need_min_p_sampling: probs = top_k_renorm_prob(probs, sampling_info.top_ks) diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index 0bedb61efa19..a2a3cdc74d13 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -5,8 +5,9 @@ from unittest.mock import patch from sglang.srt.arg_groups.overrides import resolution_result -from sglang.srt.arg_groups.platform_hook import handle_cpu_backends +from sglang.srt.arg_groups.platform_hook import handle_cpu_backends, handle_xpu_backends from sglang.srt.arg_groups.validation_hook import validate_ib_devices +from sglang.srt.model_executor.cuda_graph_config import CudaGraphConfig from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -46,6 +47,32 @@ def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") +class TestServerArgsXPUBackend(CustomTestCase): + def _make_server_args(self, sampling_backend=None): + server_args = ServerArgs.__new__(ServerArgs) + server_args.device = "xpu" + server_args.sampling_backend = sampling_backend + server_args.cuda_graph_config = CudaGraphConfig() + server_args._cuda_graph_config_locked = set() + return server_args + + def test_xpu_defaults_to_xpu_sampling_backend(self): + server_args = self._make_server_args() + + handle_xpu_backends(server_args) + + self.assertEqual( + resolution_result(server_args, "sampling_backend"), "intel_xpu" + ) + + def test_xpu_keeps_explicit_sampling_backend(self): + server_args = self._make_server_args(sampling_backend="pytorch") + + handle_xpu_backends(server_args) + + self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") + + class TestServerArgsIBDeviceValidation(CustomTestCase): def _validate_ib_devices(self, device_str, available_devices=None): available_devices = available_devices or [ From 68fdfbae4012913d96dd28e361211acb02572dea Mon Sep 17 00:00:00 2001 From: KMS07 Date: Thu, 17 Sep 2026 01:24:31 -0700 Subject: [PATCH 02/10] add xpu sampling validation hook and remove tests --- .../sglang/srt/arg_groups/validation_hook.py | 14 ++++++++ .../cpu/test_server_args_backend.py | 35 +++---------------- 2 files changed, 19 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py index 352b5f4c9624..5ba4f105c6cc 100644 --- a/python/sglang/srt/arg_groups/validation_hook.py +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -266,6 +266,8 @@ def check_server_args(server_args: Any): if cfg.enable_quant_communications and cfg.device != "npu": raise ValueError("Communications quantization is only supported for NPU device") + validate_intel_xpu_sampling_backend(cfg.sampling_backend, cfg.device) + # grpc_port is None for HTTP-only launches, so the == comparison is # already False there; no explicit None check needed. if not (cfg.smg_grpc_mode or cfg.grpc_mode) and cfg.grpc_port == cfg.port: @@ -384,6 +386,18 @@ def check_load_publish_args(server_args: Any): raise ValueError(reason) +def validate_intel_xpu_sampling_backend( + sampling_backend: Optional[str], device: str +) -> None: + # sampler.py binds the intel_xpu kernels only under is_xpu(), so on another + # device the backend either aliases to flashinfer's names or NameErrors on + # the first non-greedy decode. + if sampling_backend == "intel_xpu" and device != "xpu": + raise ValueError( + f"--sampling-backend intel_xpu requires --device xpu, got --device {device}" + ) + + def validate_ib_devices(device_str: Optional[str]) -> Optional[str]: """ Validate IB devices before passing to mooncake. diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index a2a3cdc74d13..efdb0d64b9a0 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -5,9 +5,11 @@ from unittest.mock import patch from sglang.srt.arg_groups.overrides import resolution_result -from sglang.srt.arg_groups.platform_hook import handle_cpu_backends, handle_xpu_backends -from sglang.srt.arg_groups.validation_hook import validate_ib_devices -from sglang.srt.model_executor.cuda_graph_config import CudaGraphConfig +from sglang.srt.arg_groups.platform_hook import handle_cpu_backends +from sglang.srt.arg_groups.validation_hook import ( + validate_ib_devices, + validate_intel_xpu_sampling_backend, +) from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -46,33 +48,6 @@ def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): ) self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") - -class TestServerArgsXPUBackend(CustomTestCase): - def _make_server_args(self, sampling_backend=None): - server_args = ServerArgs.__new__(ServerArgs) - server_args.device = "xpu" - server_args.sampling_backend = sampling_backend - server_args.cuda_graph_config = CudaGraphConfig() - server_args._cuda_graph_config_locked = set() - return server_args - - def test_xpu_defaults_to_xpu_sampling_backend(self): - server_args = self._make_server_args() - - handle_xpu_backends(server_args) - - self.assertEqual( - resolution_result(server_args, "sampling_backend"), "intel_xpu" - ) - - def test_xpu_keeps_explicit_sampling_backend(self): - server_args = self._make_server_args(sampling_backend="pytorch") - - handle_xpu_backends(server_args) - - self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") - - class TestServerArgsIBDeviceValidation(CustomTestCase): def _validate_ib_devices(self, device_str, available_devices=None): available_devices = available_devices or [ From 46a2a4fc96e844f9216622b2c770bf49c5eae1a2 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Thu, 17 Sep 2026 01:32:19 -0700 Subject: [PATCH 03/10] update --- test/registered/cpu/test_server_args_backend.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index efdb0d64b9a0..0bedb61efa19 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -6,10 +6,7 @@ from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.platform_hook import handle_cpu_backends -from sglang.srt.arg_groups.validation_hook import ( - validate_ib_devices, - validate_intel_xpu_sampling_backend, -) +from sglang.srt.arg_groups.validation_hook import validate_ib_devices from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -48,6 +45,7 @@ def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): ) self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") + class TestServerArgsIBDeviceValidation(CustomTestCase): def _validate_ib_devices(self, device_str, available_devices=None): available_devices = available_devices or [ From fa645e163c2b342141fe2b6a474393fc129328dc Mon Sep 17 00:00:00 2001 From: KMS07 Date: Thu, 17 Sep 2026 02:50:05 -0700 Subject: [PATCH 04/10] remove comment --- python/sglang/srt/layers/sampler.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index e5dcadd541c3..b72d6a13de0f 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -355,8 +355,6 @@ def _sample_from_probs( ) else: backend = get_exec().kernel.sampling_backend - # sgl-kernel-xpu mirrors the flashinfer sampling API, so both share - # this branch. if backend in ("flashinfer", "intel_xpu"): assert sampling_info.sampling_seed is None, ( f"Sampling seed is not supported for {backend} backend" From ffd7b306f25814a1f429264a7fc0061ede82ff25 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Thu, 17 Sep 2026 06:37:48 -0700 Subject: [PATCH 05/10] add sampling mask tests --- .../xpu/sampling/test_sampling_mask_xpu.py | 267 ++++++++++++++++++ 1 file changed, 267 insertions(+) create mode 100644 test/registered/xpu/sampling/test_sampling_mask_xpu.py diff --git a/test/registered/xpu/sampling/test_sampling_mask_xpu.py b/test/registered/xpu/sampling/test_sampling_mask_xpu.py new file mode 100644 index 000000000000..aa7e3c1329de --- /dev/null +++ b/test/registered/xpu/sampling/test_sampling_mask_xpu.py @@ -0,0 +1,267 @@ +""" +python3 -m unittest test_sampling_mask_xpu.py +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.layers import sampler as sampler_module +from sglang.srt.layers.logits_processor import LogitsProcessorOutput +from sglang.srt.layers.sampler import Sampler +from sglang.srt.sampling.custom_logit_processor import DisallowedTokensLogitsProcessor +from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo +from sglang.srt.utils import is_xpu +from sglang.test.ci.ci_register import register_xpu_ci +from sglang.test.test_utils import CustomTestCase + +register_xpu_ci(est_time=10, suite="stage-a-test-1-gpu-xpu") + + +@unittest.skipUnless(is_xpu(), "Intel XPU not available") +class TestXpuSamplingMaskCapture(CustomTestCase): + """Sampling-mask capture on the intel_xpu backend. + + The backend draws its token from the fused SYCL joint kernel but rebuilds the + captured support from the separate top_k/top_p renorm kernels, which are + written independently of it. If the two disagree on a cutoff or a tie, + selected_weight comes back 0 for a token that really was sampled -- a + silently wrong number rather than an error. + + These cases mirror TestSamplingMaskCapture in + test/registered/sampling/test_sampling_mask.py, which covers the same + invariants against the CUDA and ROCm kernels; the kernels behind them are a + separate implementation on XPU. + """ + + def setUp(self): + self.sampler = Sampler.__new__(Sampler) + torch.nn.Module.__init__(self.sampler) + + def test_hard_exclusion_replay_in_mixed_batch(self): + for backend in ["pytorch", "intel_xpu"]: + with self.subTest(backend=backend): + logits = ( + torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device="xpu") + .log() + .repeat(2, 1) + ) + original = logits.clone() + info = SamplingBatchInfo( + temperatures=torch.ones(2, 1, device="xpu"), + top_ps=torch.full((2,), 0.9, device="xpu"), + top_ks=torch.full((2,), 3, dtype=torch.int32, device="xpu"), + min_ps=torch.zeros(2, device="xpu"), + is_all_greedy=False, + is_any_greedy=False, + need_top_p_sampling=True, + need_top_k_sampling=True, + need_min_p_sampling=False, + vocab_size=5, + has_custom_logit_processor=True, + custom_params=[{"token_ids": [2]}, None], + custom_logit_processor={ + 0: ( + DisallowedTokensLogitsProcessor(), + torch.tensor([True, False], device="xpu"), + ) + }, + return_sampling_masks=[True, True], + ) + logits = self.sampler._preprocess_logits(logits, info) + with patch( + "sglang.srt.layers.sampler.get_exec", + return_value=SimpleNamespace( + kernel=SimpleNamespace(sampling_backend=backend) + ), + ): + sampled, capture = self.sampler._sample_from_probs( + logits.softmax(-1), + info, + positions=torch.zeros(2, dtype=torch.int64, device="xpu"), + simple_sampling_case=False, + return_sampling_mask=True, + ) + output = LogitsProcessorOutput(next_token_logits=None) + self.sampler._attach_sampling_mask_to_output( + output, info, sampled, capture + ) + support = output.next_token_sampling_mask_idx[0] + self.assertEqual(set(support), {0, 1, 3}) + self.assertIn(int(sampled[0]), support) + expected = original[0, sampled[0]] - original[0, support].logsumexp(0) + self.assertAlmostEqual( + output.next_token_sampling_logprobs[0], expected.item(), places=5 + ) + self.assertIn(2, output.next_token_sampling_mask_idx[1]) + + def test_intel_xpu_joint_cutoff_ties_match_capture(self): + batch_size = 256 + top_k = 2 + top_p = 0.45 + base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu") + probs = base_probs.repeat(batch_size, 1) + + # Derive the threshold-based joint support independently. Both filters + # cut at 0.2, so the tied entries must survive even though this yields + # more support entries than top_k. + sorted_probs = base_probs[0].sort(descending=True).values + top_k_cutoff = sorted_probs[top_k - 1] + mass_before = sorted_probs.cumsum(dim=-1) - sorted_probs + top_p_cutoff = sorted_probs[mass_before <= top_p][-1] + expected_support = (base_probs[0] >= top_k_cutoff) & ( + base_probs[0] >= top_p_cutoff + ) + expected_ids = expected_support.nonzero(as_tuple=True)[0].tolist() + self.assertEqual(expected_ids, [0, 1, 2]) + + sampling_info = SimpleNamespace( + sampling_seed=None, + need_top_k_sampling=True, + need_top_p_sampling=True, + need_min_p_sampling=False, + top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="xpu"), + top_ps=torch.full((batch_size,), top_p, device="xpu"), + min_ps=torch.zeros(batch_size, device="xpu"), + return_sampling_masks=[True] * batch_size, + ) + with patch( + "sglang.srt.layers.sampler.get_exec", + return_value=SimpleNamespace( + kernel=SimpleNamespace(sampling_backend="intel_xpu") + ), + ): + sampled, capture = self.sampler._sample_from_probs( + probs, + sampling_info, + positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), + simple_sampling_case=False, + return_sampling_mask=True, + ) + + self.assertIsNotNone(capture) + self.assertEqual(capture.batch_rows.cpu().tolist(), list(range(batch_size))) + actual_support = capture.weights > 0 + self.assertTrue( + torch.equal(actual_support, expected_support.expand_as(actual_support)) + ) + self.assertGreater(int(actual_support[0].sum().item()), top_k) + self.assertTrue( + bool(actual_support.gather(1, sampled.view(-1, 1)).all().item()) + ) + + def test_intel_xpu_capture_only_materializes_requested_rows(self): + batch_size = 4 + top_k = 2 + top_p = 0.45 + requested_rows = [1, 3] + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu").repeat( + batch_size, 1 + ) + sampling_info = SimpleNamespace( + sampling_seed=None, + need_top_k_sampling=True, + need_top_p_sampling=True, + need_min_p_sampling=False, + top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="xpu"), + top_ps=torch.full((batch_size,), top_p, device="xpu"), + min_ps=torch.zeros(batch_size, device="xpu"), + return_sampling_masks=[False, True, False, True], + ) + top_k_renorm = sampler_module.top_k_renorm_prob + top_p_renorm = sampler_module.top_p_renorm_prob + with ( + patch( + "sglang.srt.layers.sampler.get_exec", + return_value=SimpleNamespace( + kernel=SimpleNamespace(sampling_backend="intel_xpu") + ), + ), + patch( + "sglang.srt.layers.sampler.top_k_renorm_prob", + wraps=top_k_renorm, + ) as top_k_mock, + patch( + "sglang.srt.layers.sampler.top_p_renorm_prob", + wraps=top_p_renorm, + ) as top_p_mock, + ): + sampled, capture = self.sampler._sample_from_probs( + probs, + sampling_info, + positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), + simple_sampling_case=False, + return_sampling_mask=True, + ) + + self.assertIsNotNone(capture) + self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows) + self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5)) + self.assertEqual(tuple(top_k_mock.call_args.args[0].shape), (2, 5)) + self.assertEqual(tuple(top_p_mock.call_args.args[0].shape), (2, 5)) + + output = LogitsProcessorOutput(next_token_logits=None) + self.sampler._attach_sampling_mask_to_output( + output, sampling_info, sampled, capture + ) + self.assertIsNone(output.next_token_sampling_mask_idx[0]) + self.assertEqual(set(output.next_token_sampling_mask_idx[1]), {0, 1, 2}) + self.assertIsNone(output.next_token_sampling_mask_idx[2]) + self.assertEqual(set(output.next_token_sampling_mask_idx[3]), {0, 1, 2}) + self.assertIsNone(output.next_token_sampling_logprobs[0]) + self.assertIsNotNone(output.next_token_sampling_logprobs[1]) + self.assertIsNone(output.next_token_sampling_logprobs[2]) + self.assertIsNotNone(output.next_token_sampling_logprobs[3]) + + def test_pytorch_capture_compacts_requested_rows(self): + batch_size = 4 + requested_rows = [1, 3] + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu").repeat( + batch_size, 1 + ) + sampling_info = SimpleNamespace( + sampling_seed=None, + need_top_k_sampling=True, + need_top_p_sampling=True, + need_min_p_sampling=False, + top_ks=torch.full((batch_size,), 2, dtype=torch.int32, device="xpu"), + top_ps=torch.full((batch_size,), 0.45, device="xpu"), + min_ps=torch.zeros(batch_size, device="xpu"), + return_sampling_masks=[False, True, False, True], + ) + with patch( + "sglang.srt.layers.sampler.get_exec", + return_value=SimpleNamespace( + kernel=SimpleNamespace(sampling_backend="pytorch") + ), + ): + sampled, capture = self.sampler._sample_from_probs( + probs, + sampling_info, + positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), + simple_sampling_case=False, + return_sampling_mask=True, + ) + + self.assertIsNotNone(capture) + self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows) + self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5)) + self.assertEqual(tuple(capture.token_ids.shape), (len(requested_rows), 5)) + + output = LogitsProcessorOutput(next_token_logits=None) + self.sampler._attach_sampling_mask_to_output( + output, sampling_info, sampled, capture + ) + for batch_row in requested_rows: + self.assertIn( + int(sampled[batch_row]), + output.next_token_sampling_mask_idx[batch_row], + ) + self.assertIsNotNone(output.next_token_sampling_logprobs[batch_row]) + self.assertIsNone(output.next_token_sampling_mask_idx[0]) + self.assertIsNone(output.next_token_sampling_mask_idx[2]) + +if __name__ == "__main__": + unittest.main() From 931c6cc405c4e9a69336d8321caf2eab7caccdd5 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Thu, 17 Sep 2026 11:22:18 -0700 Subject: [PATCH 06/10] add xpu to test_sampling_mask.py --- .../registered/sampling/test_sampling_mask.py | 81 +++++++++++++------ 1 file changed, 55 insertions(+), 26 deletions(-) diff --git a/test/registered/sampling/test_sampling_mask.py b/test/registered/sampling/test_sampling_mask.py index 187e805dc912..ecde3eb0b23a 100644 --- a/test/registered/sampling/test_sampling_mask.py +++ b/test/registered/sampling/test_sampling_mask.py @@ -21,8 +21,18 @@ Qwen3ThinkingBudgetLogitProcessor, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo -from sglang.srt.utils import is_hip, kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.srt.utils import ( + get_device, + get_device_module, + is_hip, + is_xpu, + kill_process_tree, +) +from sglang.test.ci.ci_register import ( + register_amd_ci, + register_cuda_ci, + register_xpu_ci, +) from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -33,6 +43,20 @@ register_cuda_ci(est_time=250, stage="base-b", runner_config="2-gpu-large") register_amd_ci(est_time=320, suite="stage-b-test-1-gpu-small-amd") +# XPU runs the capture tests only; the server-backed classes below are skipped. +register_xpu_ci(est_time=15, suite="stage-a-test-1-gpu-xpu") + +_DEVICE = get_device() +# sgl-kernel-xpu mirrors the flashinfer sampling API, so both names select the +# same fused-joint-sampler branch in sampler.py; only the label differs. +_FUSED_BACKEND = "intel_xpu" if _DEVICE == "xpu" else "flashinfer" + +# The server-backed classes are backend-independent, so running them on XPU adds +# no kernel coverage, and /generate wedges intermittently there (reproduces under +# --sampling-backend pytorch, so it is not the fused XPU sampler). +_skip_server_on_xpu = unittest.skipIf( + is_xpu(), "server-backed sampling-mask coverage is not validated on XPU" +) _MAX_NEW_TOKENS = 4 _TOP_P = 0.99 @@ -85,10 +109,10 @@ def _sample( need_top_k_sampling=True, need_top_p_sampling=top_p < 1.0, need_min_p_sampling=min_p > 0.0, - top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="cuda"), - top_ps=torch.full((batch_size,), top_p, device="cuda"), - min_ps=torch.full((batch_size,), min_p, device="cuda"), - sampling_mask_batch_indices=torch.tensor(requested_rows, device="cuda"), + top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device=_DEVICE), + top_ps=torch.full((batch_size,), top_p, device=_DEVICE), + min_ps=torch.full((batch_size,), min_p, device=_DEVICE), + sampling_mask_batch_indices=torch.tensor(requested_rows, device=_DEVICE), ) with patch.object( sampler_module, @@ -100,7 +124,7 @@ def _sample( return self.sampler._sample_from_probs( probs, sampling_info, - positions=torch.zeros(batch_size, dtype=torch.int64, device="cuda"), + positions=torch.zeros(batch_size, dtype=torch.int64, device=_DEVICE), simple_sampling_case=False, ) @@ -122,10 +146,10 @@ def _materialize(self, sampled, capture, requested_rows): return output def test_min_p_capture_matches_filtered_support_and_logprob(self): - backends = ("pytorch",) if is_hip() else ("pytorch", "flashinfer") + backends = ("pytorch",) if is_hip() else ("pytorch", _FUSED_BACKEND) for backend in backends: with self.subTest(backend=backend): - probs = torch.tensor([[0.4, 0.3, 0.2, 0.1]], device="cuda") + probs = torch.tensor([[0.4, 0.3, 0.2, 0.1]], device=_DEVICE) sampled, capture = self._sample( probs, backend, top_k=3, top_p=1.0, min_p=0.6 ) @@ -139,20 +163,20 @@ def test_min_p_capture_matches_filtered_support_and_logprob(self): ) def test_hard_exclusion_replay_in_mixed_batch(self): - backends = ["pytorch"] if is_hip() else ["pytorch", "flashinfer"] + backends = ["pytorch"] if is_hip() else ["pytorch", _FUSED_BACKEND] for backend in backends: with self.subTest(backend=backend): logits = ( - torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device="cuda") + torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device=_DEVICE) .log() .repeat(2, 1) ) original = logits.clone() info = SamplingBatchInfo( - temperatures=torch.ones(2, 1, device="cuda"), - top_ps=torch.full((2,), 0.9, device="cuda"), - top_ks=torch.full((2,), 3, dtype=torch.int32, device="cuda"), - min_ps=torch.zeros(2, device="cuda"), + temperatures=torch.ones(2, 1, device=_DEVICE), + top_ps=torch.full((2,), 0.9, device=_DEVICE), + top_ks=torch.full((2,), 3, dtype=torch.int32, device=_DEVICE), + min_ps=torch.zeros(2, device=_DEVICE), is_all_greedy=False, is_any_greedy=False, need_top_p_sampling=True, @@ -164,11 +188,11 @@ def test_hard_exclusion_replay_in_mixed_batch(self): custom_logit_processor={ 0: ( DisallowedTokensLogitsProcessor(), - torch.tensor([True, False], device="cuda"), + torch.tensor([True, False], device=_DEVICE), ) }, return_sampling_masks=[True, True], - sampling_mask_batch_indices=torch.tensor([0, 1], device="cuda"), + sampling_mask_batch_indices=torch.tensor([0, 1], device=_DEVICE), ) logits = self.sampler._preprocess_logits(logits, info) with patch( @@ -180,7 +204,7 @@ def test_hard_exclusion_replay_in_mixed_batch(self): sampled, capture = self.sampler._sample_from_probs( logits.softmax(-1), info, - positions=torch.zeros(2, dtype=torch.int64, device="cuda"), + positions=torch.zeros(2, dtype=torch.int64, device=_DEVICE), simple_sampling_case=False, ) output = self._materialize(sampled, capture, requested_rows=[0, 1]) @@ -198,7 +222,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): batch_size = 256 top_k = 2 top_p = 0.45 - base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda") + base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE) probs = base_probs.repeat(batch_size, 1) # Derive the threshold-based joint support independently. Both filters @@ -214,7 +238,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): expected_ids = expected_support.nonzero(as_tuple=True)[0].tolist() self.assertEqual(expected_ids, [0, 1, 2]) - sampled, capture = self._sample(probs, "flashinfer", top_k=top_k, top_p=top_p) + sampled, capture = self._sample(probs, _FUSED_BACKEND, top_k=top_k, top_p=top_p) self.assertIsNotNone(capture) self.assertEqual(capture.batch_rows.cpu().tolist(), list(range(batch_size))) @@ -231,7 +255,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): def test_flashinfer_capture_only_materializes_requested_rows(self): batch_size = 4 requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat( + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE).repeat( batch_size, 1 ) with ( @@ -247,7 +271,7 @@ def test_flashinfer_capture_only_materializes_requested_rows(self): ) as top_p_mock, ): sampled, capture = self._sample( - probs, "flashinfer", requested_rows=requested_rows + probs, _FUSED_BACKEND, requested_rows=requested_rows ) self.assertIsNotNone(capture) @@ -269,7 +293,7 @@ def test_flashinfer_capture_only_materializes_requested_rows(self): def test_pytorch_capture_compacts_requested_rows(self): batch_size = 4 requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat( + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE).repeat( batch_size, 1 ) sampled, capture = self._sample(probs, "pytorch", requested_rows=requested_rows) @@ -362,6 +386,7 @@ def _generate_sampling_masks(self, sampling_params): return self._assert_sampling_masks(output_ids, meta_info) +@_skip_server_on_xpu class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): _sampling_backend = "flashinfer" @@ -563,15 +588,17 @@ def test_synced_token_logprob_is_recomputed_from_capture(self): def test_greedy_device_output_survives_async_copy(self): from sglang.srt.managers.utils import GenerationBatchResult - tokens = torch.tensor([3, 4, 5], device="cuda") + tokens = torch.tensor([3, 4, 5], device=_DEVICE) output = LogitsProcessorOutput( next_token_logits=None, sampling_mask_output=self.sampler._build_greedy_sampling_mask_output( - torch.tensor([0, 2], device="cuda"), tokens + torch.tensor([0, 2], device=_DEVICE), tokens ), ) result = GenerationBatchResult( - logits_output=output, next_token_ids=tokens, copy_done=torch.cuda.Event() + logits_output=output, + next_token_ids=tokens, + copy_done=get_device_module().Event(), ) result.copy_to_cpu(return_logprob=False) result.copy_done.synchronize() @@ -615,6 +642,7 @@ def test_overflow_never_materializes_a_partial_mask(self): self.assertEqual(output.next_token_sampling_logprobs, [None]) +@_skip_server_on_xpu class TestSamplingMaskDeterministic(SamplingMaskTestMixin, CustomTestCase): @classmethod def setUpClass(cls): @@ -651,6 +679,7 @@ def setUpClass(cls): @unittest.skipIf(is_hip(), "The AMD sampling-mask CI suite provides only one GPU.") +@_skip_server_on_xpu class TestDistributedSamplingMask(CustomTestCase): def _check_parallel_config(self, *, tp_size, pp_size): process = None From 2968fe7e128e3e7bbc0a51582476e09ee67137db Mon Sep 17 00:00:00 2001 From: Kotthagattu Meher Sai Date: Thu, 17 Sep 2026 23:55:22 +0530 Subject: [PATCH 07/10] Delete test/registered/xpu/sampling/test_sampling_mask_xpu.py --- .../xpu/sampling/test_sampling_mask_xpu.py | 267 ------------------ 1 file changed, 267 deletions(-) delete mode 100644 test/registered/xpu/sampling/test_sampling_mask_xpu.py diff --git a/test/registered/xpu/sampling/test_sampling_mask_xpu.py b/test/registered/xpu/sampling/test_sampling_mask_xpu.py deleted file mode 100644 index aa7e3c1329de..000000000000 --- a/test/registered/xpu/sampling/test_sampling_mask_xpu.py +++ /dev/null @@ -1,267 +0,0 @@ -""" -python3 -m unittest test_sampling_mask_xpu.py -""" - -import unittest -from types import SimpleNamespace -from unittest.mock import patch - -import torch - -from sglang.srt.layers import sampler as sampler_module -from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.sampler import Sampler -from sglang.srt.sampling.custom_logit_processor import DisallowedTokensLogitsProcessor -from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo -from sglang.srt.utils import is_xpu -from sglang.test.ci.ci_register import register_xpu_ci -from sglang.test.test_utils import CustomTestCase - -register_xpu_ci(est_time=10, suite="stage-a-test-1-gpu-xpu") - - -@unittest.skipUnless(is_xpu(), "Intel XPU not available") -class TestXpuSamplingMaskCapture(CustomTestCase): - """Sampling-mask capture on the intel_xpu backend. - - The backend draws its token from the fused SYCL joint kernel but rebuilds the - captured support from the separate top_k/top_p renorm kernels, which are - written independently of it. If the two disagree on a cutoff or a tie, - selected_weight comes back 0 for a token that really was sampled -- a - silently wrong number rather than an error. - - These cases mirror TestSamplingMaskCapture in - test/registered/sampling/test_sampling_mask.py, which covers the same - invariants against the CUDA and ROCm kernels; the kernels behind them are a - separate implementation on XPU. - """ - - def setUp(self): - self.sampler = Sampler.__new__(Sampler) - torch.nn.Module.__init__(self.sampler) - - def test_hard_exclusion_replay_in_mixed_batch(self): - for backend in ["pytorch", "intel_xpu"]: - with self.subTest(backend=backend): - logits = ( - torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device="xpu") - .log() - .repeat(2, 1) - ) - original = logits.clone() - info = SamplingBatchInfo( - temperatures=torch.ones(2, 1, device="xpu"), - top_ps=torch.full((2,), 0.9, device="xpu"), - top_ks=torch.full((2,), 3, dtype=torch.int32, device="xpu"), - min_ps=torch.zeros(2, device="xpu"), - is_all_greedy=False, - is_any_greedy=False, - need_top_p_sampling=True, - need_top_k_sampling=True, - need_min_p_sampling=False, - vocab_size=5, - has_custom_logit_processor=True, - custom_params=[{"token_ids": [2]}, None], - custom_logit_processor={ - 0: ( - DisallowedTokensLogitsProcessor(), - torch.tensor([True, False], device="xpu"), - ) - }, - return_sampling_masks=[True, True], - ) - logits = self.sampler._preprocess_logits(logits, info) - with patch( - "sglang.srt.layers.sampler.get_exec", - return_value=SimpleNamespace( - kernel=SimpleNamespace(sampling_backend=backend) - ), - ): - sampled, capture = self.sampler._sample_from_probs( - logits.softmax(-1), - info, - positions=torch.zeros(2, dtype=torch.int64, device="xpu"), - simple_sampling_case=False, - return_sampling_mask=True, - ) - output = LogitsProcessorOutput(next_token_logits=None) - self.sampler._attach_sampling_mask_to_output( - output, info, sampled, capture - ) - support = output.next_token_sampling_mask_idx[0] - self.assertEqual(set(support), {0, 1, 3}) - self.assertIn(int(sampled[0]), support) - expected = original[0, sampled[0]] - original[0, support].logsumexp(0) - self.assertAlmostEqual( - output.next_token_sampling_logprobs[0], expected.item(), places=5 - ) - self.assertIn(2, output.next_token_sampling_mask_idx[1]) - - def test_intel_xpu_joint_cutoff_ties_match_capture(self): - batch_size = 256 - top_k = 2 - top_p = 0.45 - base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu") - probs = base_probs.repeat(batch_size, 1) - - # Derive the threshold-based joint support independently. Both filters - # cut at 0.2, so the tied entries must survive even though this yields - # more support entries than top_k. - sorted_probs = base_probs[0].sort(descending=True).values - top_k_cutoff = sorted_probs[top_k - 1] - mass_before = sorted_probs.cumsum(dim=-1) - sorted_probs - top_p_cutoff = sorted_probs[mass_before <= top_p][-1] - expected_support = (base_probs[0] >= top_k_cutoff) & ( - base_probs[0] >= top_p_cutoff - ) - expected_ids = expected_support.nonzero(as_tuple=True)[0].tolist() - self.assertEqual(expected_ids, [0, 1, 2]) - - sampling_info = SimpleNamespace( - sampling_seed=None, - need_top_k_sampling=True, - need_top_p_sampling=True, - need_min_p_sampling=False, - top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="xpu"), - top_ps=torch.full((batch_size,), top_p, device="xpu"), - min_ps=torch.zeros(batch_size, device="xpu"), - return_sampling_masks=[True] * batch_size, - ) - with patch( - "sglang.srt.layers.sampler.get_exec", - return_value=SimpleNamespace( - kernel=SimpleNamespace(sampling_backend="intel_xpu") - ), - ): - sampled, capture = self.sampler._sample_from_probs( - probs, - sampling_info, - positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), - simple_sampling_case=False, - return_sampling_mask=True, - ) - - self.assertIsNotNone(capture) - self.assertEqual(capture.batch_rows.cpu().tolist(), list(range(batch_size))) - actual_support = capture.weights > 0 - self.assertTrue( - torch.equal(actual_support, expected_support.expand_as(actual_support)) - ) - self.assertGreater(int(actual_support[0].sum().item()), top_k) - self.assertTrue( - bool(actual_support.gather(1, sampled.view(-1, 1)).all().item()) - ) - - def test_intel_xpu_capture_only_materializes_requested_rows(self): - batch_size = 4 - top_k = 2 - top_p = 0.45 - requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu").repeat( - batch_size, 1 - ) - sampling_info = SimpleNamespace( - sampling_seed=None, - need_top_k_sampling=True, - need_top_p_sampling=True, - need_min_p_sampling=False, - top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="xpu"), - top_ps=torch.full((batch_size,), top_p, device="xpu"), - min_ps=torch.zeros(batch_size, device="xpu"), - return_sampling_masks=[False, True, False, True], - ) - top_k_renorm = sampler_module.top_k_renorm_prob - top_p_renorm = sampler_module.top_p_renorm_prob - with ( - patch( - "sglang.srt.layers.sampler.get_exec", - return_value=SimpleNamespace( - kernel=SimpleNamespace(sampling_backend="intel_xpu") - ), - ), - patch( - "sglang.srt.layers.sampler.top_k_renorm_prob", - wraps=top_k_renorm, - ) as top_k_mock, - patch( - "sglang.srt.layers.sampler.top_p_renorm_prob", - wraps=top_p_renorm, - ) as top_p_mock, - ): - sampled, capture = self.sampler._sample_from_probs( - probs, - sampling_info, - positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), - simple_sampling_case=False, - return_sampling_mask=True, - ) - - self.assertIsNotNone(capture) - self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows) - self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5)) - self.assertEqual(tuple(top_k_mock.call_args.args[0].shape), (2, 5)) - self.assertEqual(tuple(top_p_mock.call_args.args[0].shape), (2, 5)) - - output = LogitsProcessorOutput(next_token_logits=None) - self.sampler._attach_sampling_mask_to_output( - output, sampling_info, sampled, capture - ) - self.assertIsNone(output.next_token_sampling_mask_idx[0]) - self.assertEqual(set(output.next_token_sampling_mask_idx[1]), {0, 1, 2}) - self.assertIsNone(output.next_token_sampling_mask_idx[2]) - self.assertEqual(set(output.next_token_sampling_mask_idx[3]), {0, 1, 2}) - self.assertIsNone(output.next_token_sampling_logprobs[0]) - self.assertIsNotNone(output.next_token_sampling_logprobs[1]) - self.assertIsNone(output.next_token_sampling_logprobs[2]) - self.assertIsNotNone(output.next_token_sampling_logprobs[3]) - - def test_pytorch_capture_compacts_requested_rows(self): - batch_size = 4 - requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="xpu").repeat( - batch_size, 1 - ) - sampling_info = SimpleNamespace( - sampling_seed=None, - need_top_k_sampling=True, - need_top_p_sampling=True, - need_min_p_sampling=False, - top_ks=torch.full((batch_size,), 2, dtype=torch.int32, device="xpu"), - top_ps=torch.full((batch_size,), 0.45, device="xpu"), - min_ps=torch.zeros(batch_size, device="xpu"), - return_sampling_masks=[False, True, False, True], - ) - with patch( - "sglang.srt.layers.sampler.get_exec", - return_value=SimpleNamespace( - kernel=SimpleNamespace(sampling_backend="pytorch") - ), - ): - sampled, capture = self.sampler._sample_from_probs( - probs, - sampling_info, - positions=torch.zeros(batch_size, dtype=torch.int64, device="xpu"), - simple_sampling_case=False, - return_sampling_mask=True, - ) - - self.assertIsNotNone(capture) - self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows) - self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5)) - self.assertEqual(tuple(capture.token_ids.shape), (len(requested_rows), 5)) - - output = LogitsProcessorOutput(next_token_logits=None) - self.sampler._attach_sampling_mask_to_output( - output, sampling_info, sampled, capture - ) - for batch_row in requested_rows: - self.assertIn( - int(sampled[batch_row]), - output.next_token_sampling_mask_idx[batch_row], - ) - self.assertIsNotNone(output.next_token_sampling_logprobs[batch_row]) - self.assertIsNone(output.next_token_sampling_mask_idx[0]) - self.assertIsNone(output.next_token_sampling_mask_idx[2]) - -if __name__ == "__main__": - unittest.main() From c6d6cdaedbeee9de9da62d25092d023b781b1dd3 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Fri, 18 Sep 2026 01:53:26 -0700 Subject: [PATCH 08/10] Revert XPU registration of the sampling-mask tests The fused-vs-renorm tie agreement was verified locally on XPU, but keep this PR scoped to the sampling backend itself; the test wiring lands separately. --- .../registered/sampling/test_sampling_mask.py | 81 ++++++------------- 1 file changed, 26 insertions(+), 55 deletions(-) diff --git a/test/registered/sampling/test_sampling_mask.py b/test/registered/sampling/test_sampling_mask.py index ecde3eb0b23a..187e805dc912 100644 --- a/test/registered/sampling/test_sampling_mask.py +++ b/test/registered/sampling/test_sampling_mask.py @@ -21,18 +21,8 @@ Qwen3ThinkingBudgetLogitProcessor, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo -from sglang.srt.utils import ( - get_device, - get_device_module, - is_hip, - is_xpu, - kill_process_tree, -) -from sglang.test.ci.ci_register import ( - register_amd_ci, - register_cuda_ci, - register_xpu_ci, -) +from sglang.srt.utils import is_hip, kill_process_tree +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -43,20 +33,6 @@ register_cuda_ci(est_time=250, stage="base-b", runner_config="2-gpu-large") register_amd_ci(est_time=320, suite="stage-b-test-1-gpu-small-amd") -# XPU runs the capture tests only; the server-backed classes below are skipped. -register_xpu_ci(est_time=15, suite="stage-a-test-1-gpu-xpu") - -_DEVICE = get_device() -# sgl-kernel-xpu mirrors the flashinfer sampling API, so both names select the -# same fused-joint-sampler branch in sampler.py; only the label differs. -_FUSED_BACKEND = "intel_xpu" if _DEVICE == "xpu" else "flashinfer" - -# The server-backed classes are backend-independent, so running them on XPU adds -# no kernel coverage, and /generate wedges intermittently there (reproduces under -# --sampling-backend pytorch, so it is not the fused XPU sampler). -_skip_server_on_xpu = unittest.skipIf( - is_xpu(), "server-backed sampling-mask coverage is not validated on XPU" -) _MAX_NEW_TOKENS = 4 _TOP_P = 0.99 @@ -109,10 +85,10 @@ def _sample( need_top_k_sampling=True, need_top_p_sampling=top_p < 1.0, need_min_p_sampling=min_p > 0.0, - top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device=_DEVICE), - top_ps=torch.full((batch_size,), top_p, device=_DEVICE), - min_ps=torch.full((batch_size,), min_p, device=_DEVICE), - sampling_mask_batch_indices=torch.tensor(requested_rows, device=_DEVICE), + top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="cuda"), + top_ps=torch.full((batch_size,), top_p, device="cuda"), + min_ps=torch.full((batch_size,), min_p, device="cuda"), + sampling_mask_batch_indices=torch.tensor(requested_rows, device="cuda"), ) with patch.object( sampler_module, @@ -124,7 +100,7 @@ def _sample( return self.sampler._sample_from_probs( probs, sampling_info, - positions=torch.zeros(batch_size, dtype=torch.int64, device=_DEVICE), + positions=torch.zeros(batch_size, dtype=torch.int64, device="cuda"), simple_sampling_case=False, ) @@ -146,10 +122,10 @@ def _materialize(self, sampled, capture, requested_rows): return output def test_min_p_capture_matches_filtered_support_and_logprob(self): - backends = ("pytorch",) if is_hip() else ("pytorch", _FUSED_BACKEND) + backends = ("pytorch",) if is_hip() else ("pytorch", "flashinfer") for backend in backends: with self.subTest(backend=backend): - probs = torch.tensor([[0.4, 0.3, 0.2, 0.1]], device=_DEVICE) + probs = torch.tensor([[0.4, 0.3, 0.2, 0.1]], device="cuda") sampled, capture = self._sample( probs, backend, top_k=3, top_p=1.0, min_p=0.6 ) @@ -163,20 +139,20 @@ def test_min_p_capture_matches_filtered_support_and_logprob(self): ) def test_hard_exclusion_replay_in_mixed_batch(self): - backends = ["pytorch"] if is_hip() else ["pytorch", _FUSED_BACKEND] + backends = ["pytorch"] if is_hip() else ["pytorch", "flashinfer"] for backend in backends: with self.subTest(backend=backend): logits = ( - torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device=_DEVICE) + torch.tensor([[0.3, 0.2, 0.5, 0.15, 0.1]], device="cuda") .log() .repeat(2, 1) ) original = logits.clone() info = SamplingBatchInfo( - temperatures=torch.ones(2, 1, device=_DEVICE), - top_ps=torch.full((2,), 0.9, device=_DEVICE), - top_ks=torch.full((2,), 3, dtype=torch.int32, device=_DEVICE), - min_ps=torch.zeros(2, device=_DEVICE), + temperatures=torch.ones(2, 1, device="cuda"), + top_ps=torch.full((2,), 0.9, device="cuda"), + top_ks=torch.full((2,), 3, dtype=torch.int32, device="cuda"), + min_ps=torch.zeros(2, device="cuda"), is_all_greedy=False, is_any_greedy=False, need_top_p_sampling=True, @@ -188,11 +164,11 @@ def test_hard_exclusion_replay_in_mixed_batch(self): custom_logit_processor={ 0: ( DisallowedTokensLogitsProcessor(), - torch.tensor([True, False], device=_DEVICE), + torch.tensor([True, False], device="cuda"), ) }, return_sampling_masks=[True, True], - sampling_mask_batch_indices=torch.tensor([0, 1], device=_DEVICE), + sampling_mask_batch_indices=torch.tensor([0, 1], device="cuda"), ) logits = self.sampler._preprocess_logits(logits, info) with patch( @@ -204,7 +180,7 @@ def test_hard_exclusion_replay_in_mixed_batch(self): sampled, capture = self.sampler._sample_from_probs( logits.softmax(-1), info, - positions=torch.zeros(2, dtype=torch.int64, device=_DEVICE), + positions=torch.zeros(2, dtype=torch.int64, device="cuda"), simple_sampling_case=False, ) output = self._materialize(sampled, capture, requested_rows=[0, 1]) @@ -222,7 +198,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): batch_size = 256 top_k = 2 top_p = 0.45 - base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE) + base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda") probs = base_probs.repeat(batch_size, 1) # Derive the threshold-based joint support independently. Both filters @@ -238,7 +214,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): expected_ids = expected_support.nonzero(as_tuple=True)[0].tolist() self.assertEqual(expected_ids, [0, 1, 2]) - sampled, capture = self._sample(probs, _FUSED_BACKEND, top_k=top_k, top_p=top_p) + sampled, capture = self._sample(probs, "flashinfer", top_k=top_k, top_p=top_p) self.assertIsNotNone(capture) self.assertEqual(capture.batch_rows.cpu().tolist(), list(range(batch_size))) @@ -255,7 +231,7 @@ def test_flashinfer_joint_cutoff_ties_match_capture(self): def test_flashinfer_capture_only_materializes_requested_rows(self): batch_size = 4 requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE).repeat( + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat( batch_size, 1 ) with ( @@ -271,7 +247,7 @@ def test_flashinfer_capture_only_materializes_requested_rows(self): ) as top_p_mock, ): sampled, capture = self._sample( - probs, _FUSED_BACKEND, requested_rows=requested_rows + probs, "flashinfer", requested_rows=requested_rows ) self.assertIsNotNone(capture) @@ -293,7 +269,7 @@ def test_flashinfer_capture_only_materializes_requested_rows(self): def test_pytorch_capture_compacts_requested_rows(self): batch_size = 4 requested_rows = [1, 3] - probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device=_DEVICE).repeat( + probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat( batch_size, 1 ) sampled, capture = self._sample(probs, "pytorch", requested_rows=requested_rows) @@ -386,7 +362,6 @@ def _generate_sampling_masks(self, sampling_params): return self._assert_sampling_masks(output_ids, meta_info) -@_skip_server_on_xpu class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): _sampling_backend = "flashinfer" @@ -588,17 +563,15 @@ def test_synced_token_logprob_is_recomputed_from_capture(self): def test_greedy_device_output_survives_async_copy(self): from sglang.srt.managers.utils import GenerationBatchResult - tokens = torch.tensor([3, 4, 5], device=_DEVICE) + tokens = torch.tensor([3, 4, 5], device="cuda") output = LogitsProcessorOutput( next_token_logits=None, sampling_mask_output=self.sampler._build_greedy_sampling_mask_output( - torch.tensor([0, 2], device=_DEVICE), tokens + torch.tensor([0, 2], device="cuda"), tokens ), ) result = GenerationBatchResult( - logits_output=output, - next_token_ids=tokens, - copy_done=get_device_module().Event(), + logits_output=output, next_token_ids=tokens, copy_done=torch.cuda.Event() ) result.copy_to_cpu(return_logprob=False) result.copy_done.synchronize() @@ -642,7 +615,6 @@ def test_overflow_never_materializes_a_partial_mask(self): self.assertEqual(output.next_token_sampling_logprobs, [None]) -@_skip_server_on_xpu class TestSamplingMaskDeterministic(SamplingMaskTestMixin, CustomTestCase): @classmethod def setUpClass(cls): @@ -679,7 +651,6 @@ def setUpClass(cls): @unittest.skipIf(is_hip(), "The AMD sampling-mask CI suite provides only one GPU.") -@_skip_server_on_xpu class TestDistributedSamplingMask(CustomTestCase): def _check_parallel_config(self, *, tp_size, pp_size): process = None From 3ce2c5ac7d1f9ff2c5f8ab3d5f6f0762dc881474 Mon Sep 17 00:00:00 2001 From: KMS07 Date: Fri, 18 Sep 2026 03:11:57 -0700 Subject: [PATCH 09/10] validation function name change --- python/sglang/srt/arg_groups/validation_hook.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py index 5ba4f105c6cc..5504e7bd9c58 100644 --- a/python/sglang/srt/arg_groups/validation_hook.py +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -266,7 +266,7 @@ def check_server_args(server_args: Any): if cfg.enable_quant_communications and cfg.device != "npu": raise ValueError("Communications quantization is only supported for NPU device") - validate_intel_xpu_sampling_backend(cfg.sampling_backend, cfg.device) + validate_device_sampling_backend(cfg.sampling_backend, cfg.device) # grpc_port is None for HTTP-only launches, so the == comparison is # already False there; no explicit None check needed. @@ -386,7 +386,7 @@ def check_load_publish_args(server_args: Any): raise ValueError(reason) -def validate_intel_xpu_sampling_backend( +def validate_device_sampling_backend( sampling_backend: Optional[str], device: str ) -> None: # sampler.py binds the intel_xpu kernels only under is_xpu(), so on another From a3980fead6ecc4ac2c84ec0e1a41874846c21e1e Mon Sep 17 00:00:00 2001 From: KMS07 Date: Fri, 18 Sep 2026 05:24:04 -0700 Subject: [PATCH 10/10] empty commit