diff --git a/python/sglang/srt/arg_groups/attention_hook.py b/python/sglang/srt/arg_groups/attention_hook.py index 14a5e3a12..f772f05e4 100644 --- a/python/sglang/srt/arg_groups/attention_hook.py +++ b/python/sglang/srt/arg_groups/attention_hook.py @@ -27,7 +27,6 @@ from sglang.srt.arg_groups.overrides import ( resolved_view, resolving_view, run_post_process_pass, - use_mla_backend, ) from sglang.srt.connector import ConnectorType from sglang.srt.environ import envs @@ -202,12 +201,7 @@ def handle_attention_backend_compatibility(server_args: Any): # Other platforms backends run_post_process_pass(server_args, _attention_backend_platform_fallbacks) - prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args)) - if use_mla_backend(server_args) and prefill_backend == "intel_xpu": - raise ValueError( - "intel_xpu backend is only supported on decode for MLA models, please set --decode-attention-backend to intel_xpu and do not set --attention-backend or --prefill-attention-backend to intel_xpu for prefill instead use triton." - ) - + # XPU platforms backends run_post_process_pass(server_args, _intel_xpu_page_constraint) # Dual chunk flash attention backend diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index ad944e23e..c02e823af 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1416,14 +1416,13 @@ def _attention_backend_platform_fallbacks(view: Any) -> dict: @register_post_process def _intel_xpu_page_constraint(view: Any) -> dict: - _, decode_backend = attention_backends_of(view) - if decode_backend == "intel_xpu": + prefill_backend, decode_backend = attention_backends_of(view) + if "intel_xpu" in (prefill_backend, decode_backend): + supported_page_sizes = [64, 128] + msg = "Intel XPU attention backend" if use_mla_backend(view): - supported_page_sizes = [16, 32, 64, 128] - msg = "Intel XPU attention backend for MLA Decode" - else: - supported_page_sizes = [64, 128] - msg = "Intel XPU attention backend" + supported_page_sizes.extend([16, 32]) + msg = msg + " for MLA" if view.page_size not in supported_page_sizes: logger.warning( f"{msg} only supports page_sizes of {supported_page_sizes}, changing page_size from {view.page_size} to 128." diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index 900d0fbd7..5573bc477 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -25,7 +25,13 @@ if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner -from sgl_kernel import flash_mla_decode, flash_mla_get_workspace_size, merge_state_v2 +from sgl_kernel import ( + flash_mla_decode, + flash_mla_decode_get_workspace_size, + flash_mla_prefill, + flash_mla_prefill_get_workspace_size, + merge_state_v2, +) from sgl_kernel.flash_attn import flash_attn_varlen_func, flash_attn_with_kvcache @@ -36,7 +42,6 @@ class XPUAttentionBackend(AttentionBackend): - Prefill and Decode disaggregation, currently only chunked prefill is supported - Speculative Decoding support - XPU Graph support, see https://github.com/pytorch/pytorch/issues/162143 - - MLA Prefill support """ def __init__( @@ -417,19 +422,34 @@ class XPUAttentionBackend(AttentionBackend): ) if self.use_mla: - workspace_size = flash_mla_get_workspace_size( + workspace_kwargs = dict( max_seq_len=self.max_context_len, num_batches=batch_size, num_heads=self.num_local_heads, page_size=self.page_size, num_kv_splits=-1, ) + + workspace_decode_size = flash_mla_decode_get_workspace_size( + **workspace_kwargs + ) if ( - not hasattr(self, "workspace") - or self.workspace.numel() < workspace_size + not hasattr(self, "workspace_decode") + or self.workspace_decode.numel() < workspace_decode_size ): - self.workspace = torch.empty( - workspace_size, device=self.device, dtype=torch.uint8 + self.workspace_decode = torch.empty( + workspace_decode_size, device=self.device, dtype=torch.uint8 + ) + + workspace_prefill_size = flash_mla_prefill_get_workspace_size( + **workspace_kwargs + ) + if ( + not hasattr(self, "workspace_prefill") + or self.workspace_prefill.numel() < workspace_prefill_size + ): + self.workspace_prefill = torch.empty( + workspace_prefill_size, device=self.device, dtype=torch.uint8 ) # Convert the page table to a strided format which is needed by FA3 API. @@ -678,6 +698,9 @@ class XPUAttentionBackend(AttentionBackend): forward_batch.attn_attend_prefix_cache is not None and not forward_batch.forward_mode.is_target_verify() ): + q = q.contiguous() + k = k.contiguous() + v = v.contiguous() # Do multi-head attention with chunked prefix cache if forward_batch.attn_attend_prefix_cache: assert not get_schedule().disable_chunked_prefix_cache @@ -722,21 +745,20 @@ class XPUAttentionBackend(AttentionBackend): return output, lse return output else: + assert not use_cascade_attn, ( + "Cascade attention is not supported with MLA" + ) + assert causal, "Non-causal MLA prefill is not supported" + # flash_mla_prefill has no softcap argument, unlike the + # flash_attn_with_kvcache call it replaced. No MLA model sets a + # logit cap today; fail loudly rather than ignore one silently. + assert not layer.logit_cap, "MLA prefill does not support logit_cap" + # Do absorbed multi-latent attention kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to( q.dtype ) - k_rope = kv_cache[:, :, layer.v_head_dim :] - c_kv = kv_cache[:, :, : layer.v_head_dim] - k_rope_cache = k_rope.view( - -1, - self.page_size, - layer.tp_k_head_num, - layer.head_dim - layer.v_head_dim, - ) - c_kv_cache = c_kv.view( - -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim - ) + if q_rope is not None: q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) q_rope = q_rope.view( @@ -747,55 +769,20 @@ class XPUAttentionBackend(AttentionBackend): q_nope = q_all[:, :, : layer.v_head_dim] q_rope = q_all[:, :, layer.v_head_dim :] - result = flash_attn_with_kvcache( - q=q_rope, - k_cache=k_rope_cache, - v_cache=c_kv_cache, - qv=q_nope, - page_table=page_table, - cache_seqlens=cache_seqlens, + o = flash_mla_prefill( + q_nope=q_nope, + q_pe=q_rope, + kv_c_and_k_pe_cache=kv_cache.view( + -1, self.page_size, layer.head_dim + ), cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k_new=None, + seq_lens_k=cache_seqlens, max_seqlen_q=max_seqlen_q, - softmax_scale=layer.scaling, + page_table=page_table, + workspace=self.workspace_prefill, + sm_scale=layer.scaling, causal=False if use_cascade_attn else causal, - softcap=layer.logit_cap, - 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 - o_expand, softmax_lse_expand, *rest_expand = ( - flash_attn_with_kvcache( - q=q_rope, - k_cache=k_rope_cache, - v_cache=c_kv_cache, - qv=q_nope, - page_table=self.forward_metadata_spec_decode_expand.page_table, - cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, - cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, - cu_seqlens_k_new=None, - max_seqlen_q=self.forward_metadata_spec_decode_expand.max_seq_len_q, - softmax_scale=layer.scaling, - causal=False, - window_size=window_size, - softcap=layer.logit_cap, - k_descale=k_descale, - v_descale=v_descale, - return_softmax_lse=True, - num_splits=self.num_splits, - ) - ) - o, _ = merge_state_v2_wrapper( - o, - softmax_lse.T.contiguous(), - o_expand, - softmax_lse_expand.T.contiguous(), - ) - else: - o = result out = o.view(-1, layer.tp_q_head_num * layer.v_head_dim) return out @@ -1114,7 +1101,7 @@ class XPUAttentionBackend(AttentionBackend): kv_cache.view(-1, self.page_size, layer.head_dim), metadata.cache_seqlens_int32, metadata.page_table, - self.workspace, + self.workspace_decode, 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 diff --git a/test/registered/xpu/test_intel_xpu_backend.py b/test/registered/xpu/test_intel_xpu_backend.py index dc884ae9e..3f4af63a7 100644 --- a/test/registered/xpu/test_intel_xpu_backend.py +++ b/test/registered/xpu/test_intel_xpu_backend.py @@ -72,14 +72,14 @@ class TestIntelXPUBackend(CustomTestCase): [ "--json-model-override-args", '{"num_hidden_layers": 4}', - "--decode-attention-backend", + "--attention-backend", "intel_xpu", "--moe-runner-backend", "triton", # FP8 is not yet supported in sgl-kernel ], min_throughput=32, ) - def test_mla_decode_attention_backend(self): + def test_mla_models_with_intel_xpu_attention_backend(self): return DEFAULT_MODEL_NAME_FOR_TEST_FP8_WITH_MOE