diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index e8c6206f5..8cbb6895e 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -99,6 +99,15 @@ class XPUAttentionBackend(AttentionBackend): self.has_swa = ( self.sliding_window_size is not None and self.sliding_window_size > -1 ) + + # If num_splits == 0, the kernel uses a heuristic to automatically + # determine the number of splits. Split-KV reduces across a + # non-deterministic number of partitions, so we pin num_splits to 1 + # when deterministic inference is enabled to keep attention reduction + # order fixed. This mirrors the flash-attention (fa3) backend. + self.num_splits = ( + 1 if model_runner.server_args.enable_deterministic_inference else 0 + ) self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder def init_forward_metadata(self, forward_batch: ForwardBatch): @@ -553,6 +562,10 @@ class XPUAttentionBackend(AttentionBackend): # Use Flash Attention for prefill if not self.use_mla: # Do multi-head attention + # The MLA branch passes num_splits explicitly per call site, since the + # chunked-prefix varlen kernels there keep their own default. + kwargs["num_splits"] = self.num_splits + key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) key_cache = key_cache.view( -1, self.page_size, layer.tp_k_head_num, layer.head_dim @@ -716,6 +729,7 @@ class XPUAttentionBackend(AttentionBackend): k_descale=k_descale, v_descale=v_descale, return_softmax_lse=use_cascade_attn, + num_splits=self.num_splits, ) if use_cascade_attn: o, softmax_lse, *rest = result @@ -737,6 +751,7 @@ class XPUAttentionBackend(AttentionBackend): k_descale=k_descale, v_descale=v_descale, return_softmax_lse=True, + num_splits=self.num_splits, ) ) o, _ = merge_state_v2_wrapper( @@ -843,6 +858,11 @@ class XPUAttentionBackend(AttentionBackend): if not self.use_mla: # Do multi-head attention + # Only the MHA kernels below take num_splits. The MLA path calls + # flash_mla_decode, whose own num_kv_splits already defaults to 1 + # (no split-KV), so it needs no deterministic override here. + kwargs["num_splits"] = self.num_splits + key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) key_cache = key_cache.view( -1, self.page_size, layer.tp_k_head_num, layer.head_dim @@ -987,6 +1007,12 @@ class XPUAttentionBackend(AttentionBackend): metadata.page_table, self.workspace, layer.scaling, + # flash_mla_decode's heuristic only kicks in when num_kv_splits + # < 1, and it derives the split count from batch * num_heads and + # seq_len_kv, which is not batch-invariant. Pin it to 1 (the + # kernel's current default) so the reduction order stays fixed + # regardless of upstream default changes. + num_kv_splits=1, ) out = o.view(-1, layer.tp_q_head_num * layer.v_head_dim) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index cdf2e82b3..7b584d561 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -241,6 +241,7 @@ DETERMINISTIC_ATTENTION_BACKEND_CHOICES = [ "fa3", "fa4", "flashinfer", + "intel_xpu", "triton", ] diff --git a/test/registered/attention/test_deterministic.py b/test/registered/attention/test_deterministic.py index b7cdc7814..4b0d46809 100644 --- a/test/registered/attention/test_deterministic.py +++ b/test/registered/attention/test_deterministic.py @@ -9,17 +9,29 @@ test into unit tests so that's easily reproducible in CI. import unittest -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.srt.utils import is_xpu +from sglang.test.ci.ci_register import ( + register_amd_ci, + register_cuda_ci, + register_xpu_ci, +) from sglang.test.test_deterministic_utils import ( COMMON_SERVER_ARGS, TestDeterministicBase, ) -from sglang.test.test_utils import is_in_amd_ci +from sglang.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN, + is_in_amd_ci, +) register_cuda_ci(est_time=207, stage="base-b", runner_config="1-gpu-large") register_amd_ci(est_time=278, suite="stage-b-test-1-gpu-small-amd") +register_xpu_ci(est_time=207, suite="stage-b-test-1-gpu-xpu") + +_is_xpu = is_xpu() +@unittest.skipIf(_is_xpu, "CUDA runner only") @unittest.skipIf(is_in_amd_ci(), "Skip for AMD CI.") class TestFlashinferDeterministic(TestDeterministicBase): # Test with flashinfer attention backend @@ -35,6 +47,7 @@ class TestFlashinferDeterministic(TestDeterministicBase): return args +@unittest.skipIf(_is_xpu, "CUDA runner only") @unittest.skipIf(is_in_amd_ci(), "Skip for AMD CI.") class TestFa3Deterministic(TestDeterministicBase): # Test with fa3 attention backend @@ -50,6 +63,7 @@ class TestFa3Deterministic(TestDeterministicBase): return args +@unittest.skipIf(_is_xpu, "CUDA/AMD runner only") class TestTritonDeterministic(TestDeterministicBase): # Test with triton attention backend @classmethod @@ -64,5 +78,29 @@ class TestTritonDeterministic(TestDeterministicBase): return args +@unittest.skipUnless(_is_xpu, "XPU runner only") +class TestIntelXPUDeterministic(TestDeterministicBase): + # Test with intel_xpu attention backend using smaller model to avoid OOM + @classmethod + def get_model(cls): + # Use smaller model for XPU to avoid OOM + return DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN + + @classmethod + def get_server_args(cls): + args = COMMON_SERVER_ARGS + args.extend( + [ + "--attention-backend", + "intel_xpu", + "--device", + "xpu", + "--mem-fraction-static", + "0.80", + ] + ) + return args + + if __name__ == "__main__": unittest.main()