[Intel GPU] Add MLA support to Intel XPU Attention backend for Prefill (#35866)
Signed-off-by: P V R K Jyothendra Varma <polisetty.v.r.k.jyothendra.varma@intel.com>
This commit is contained in:
@@ -27,7 +27,6 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
resolved_view,
|
resolved_view,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
run_post_process_pass,
|
run_post_process_pass,
|
||||||
use_mla_backend,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.connector import ConnectorType
|
from sglang.srt.connector import ConnectorType
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -202,12 +201,7 @@ def handle_attention_backend_compatibility(server_args: Any):
|
|||||||
# Other platforms backends
|
# Other platforms backends
|
||||||
run_post_process_pass(server_args, _attention_backend_platform_fallbacks)
|
run_post_process_pass(server_args, _attention_backend_platform_fallbacks)
|
||||||
|
|
||||||
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
|
# XPU platforms backends
|
||||||
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."
|
|
||||||
)
|
|
||||||
|
|
||||||
run_post_process_pass(server_args, _intel_xpu_page_constraint)
|
run_post_process_pass(server_args, _intel_xpu_page_constraint)
|
||||||
|
|
||||||
# Dual chunk flash attention backend
|
# Dual chunk flash attention backend
|
||||||
|
|||||||
@@ -1416,14 +1416,13 @@ def _attention_backend_platform_fallbacks(view: Any) -> dict:
|
|||||||
|
|
||||||
@register_post_process
|
@register_post_process
|
||||||
def _intel_xpu_page_constraint(view: Any) -> dict:
|
def _intel_xpu_page_constraint(view: Any) -> dict:
|
||||||
_, decode_backend = attention_backends_of(view)
|
prefill_backend, decode_backend = attention_backends_of(view)
|
||||||
if decode_backend == "intel_xpu":
|
if "intel_xpu" in (prefill_backend, decode_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]
|
supported_page_sizes = [64, 128]
|
||||||
msg = "Intel XPU attention backend"
|
msg = "Intel XPU attention backend"
|
||||||
|
if use_mla_backend(view):
|
||||||
|
supported_page_sizes.extend([16, 32])
|
||||||
|
msg = msg + " for MLA"
|
||||||
if view.page_size not in supported_page_sizes:
|
if view.page_size not in supported_page_sizes:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"{msg} only supports page_sizes of {supported_page_sizes}, changing page_size from {view.page_size} to 128."
|
f"{msg} only supports page_sizes of {supported_page_sizes}, changing page_size from {view.page_size} to 128."
|
||||||
|
|||||||
@@ -25,7 +25,13 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
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
|
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
|
- Prefill and Decode disaggregation, currently only chunked prefill is supported
|
||||||
- Speculative Decoding support
|
- Speculative Decoding support
|
||||||
- XPU Graph support, see https://github.com/pytorch/pytorch/issues/162143
|
- XPU Graph support, see https://github.com/pytorch/pytorch/issues/162143
|
||||||
- MLA Prefill support
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -417,19 +422,34 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
workspace_size = flash_mla_get_workspace_size(
|
workspace_kwargs = dict(
|
||||||
max_seq_len=self.max_context_len,
|
max_seq_len=self.max_context_len,
|
||||||
num_batches=batch_size,
|
num_batches=batch_size,
|
||||||
num_heads=self.num_local_heads,
|
num_heads=self.num_local_heads,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
num_kv_splits=-1,
|
num_kv_splits=-1,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
workspace_decode_size = flash_mla_decode_get_workspace_size(
|
||||||
|
**workspace_kwargs
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
not hasattr(self, "workspace")
|
not hasattr(self, "workspace_decode")
|
||||||
or self.workspace.numel() < workspace_size
|
or self.workspace_decode.numel() < workspace_decode_size
|
||||||
):
|
):
|
||||||
self.workspace = torch.empty(
|
self.workspace_decode = torch.empty(
|
||||||
workspace_size, device=self.device, dtype=torch.uint8
|
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.
|
# 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
|
forward_batch.attn_attend_prefix_cache is not None
|
||||||
and not forward_batch.forward_mode.is_target_verify()
|
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
|
# Do multi-head attention with chunked prefix cache
|
||||||
if forward_batch.attn_attend_prefix_cache:
|
if forward_batch.attn_attend_prefix_cache:
|
||||||
assert not get_schedule().disable_chunked_prefix_cache
|
assert not get_schedule().disable_chunked_prefix_cache
|
||||||
@@ -722,21 +745,20 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
return output, lse
|
return output, lse
|
||||||
return output
|
return output
|
||||||
else:
|
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
|
# Do absorbed multi-latent attention
|
||||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||||
q.dtype
|
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:
|
if q_rope is not None:
|
||||||
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||||
q_rope = q_rope.view(
|
q_rope = q_rope.view(
|
||||||
@@ -747,55 +769,20 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
q_nope = q_all[:, :, : layer.v_head_dim]
|
q_nope = q_all[:, :, : layer.v_head_dim]
|
||||||
q_rope = q_all[:, :, layer.v_head_dim :]
|
q_rope = q_all[:, :, layer.v_head_dim :]
|
||||||
|
|
||||||
result = flash_attn_with_kvcache(
|
o = flash_mla_prefill(
|
||||||
q=q_rope,
|
q_nope=q_nope,
|
||||||
k_cache=k_rope_cache,
|
q_pe=q_rope,
|
||||||
v_cache=c_kv_cache,
|
kv_c_and_k_pe_cache=kv_cache.view(
|
||||||
qv=q_nope,
|
-1, self.page_size, layer.head_dim
|
||||||
page_table=page_table,
|
),
|
||||||
cache_seqlens=cache_seqlens,
|
|
||||||
cu_seqlens_q=cu_seqlens_q,
|
cu_seqlens_q=cu_seqlens_q,
|
||||||
cu_seqlens_k_new=None,
|
seq_lens_k=cache_seqlens,
|
||||||
max_seqlen_q=max_seqlen_q,
|
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,
|
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)
|
out = o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
return out
|
return out
|
||||||
@@ -1114,7 +1101,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
kv_cache.view(-1, self.page_size, layer.head_dim),
|
kv_cache.view(-1, self.page_size, layer.head_dim),
|
||||||
metadata.cache_seqlens_int32,
|
metadata.cache_seqlens_int32,
|
||||||
metadata.page_table,
|
metadata.page_table,
|
||||||
self.workspace,
|
self.workspace_decode,
|
||||||
layer.scaling,
|
layer.scaling,
|
||||||
# flash_mla_decode's heuristic only kicks in when num_kv_splits
|
# flash_mla_decode's heuristic only kicks in when num_kv_splits
|
||||||
# < 1, and it derives the split count from batch * num_heads and
|
# < 1, and it derives the split count from batch * num_heads and
|
||||||
|
|||||||
@@ -72,14 +72,14 @@ class TestIntelXPUBackend(CustomTestCase):
|
|||||||
[
|
[
|
||||||
"--json-model-override-args",
|
"--json-model-override-args",
|
||||||
'{"num_hidden_layers": 4}',
|
'{"num_hidden_layers": 4}',
|
||||||
"--decode-attention-backend",
|
"--attention-backend",
|
||||||
"intel_xpu",
|
"intel_xpu",
|
||||||
"--moe-runner-backend",
|
"--moe-runner-backend",
|
||||||
"triton", # FP8 is not yet supported in sgl-kernel
|
"triton", # FP8 is not yet supported in sgl-kernel
|
||||||
],
|
],
|
||||||
min_throughput=32,
|
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
|
return DEFAULT_MODEL_NAME_FOR_TEST_FP8_WITH_MOE
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user