XPU: Enable GLM5.1 (GlmMoeDsaForCausalLM) DSA Attention (#24959)

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
Xia Weiwen
2026-09-07 09:24:45 +08:00
committed by GitHub
co-authored by Copilot Ma Mingfei
parent 39a80354aa
commit c4e52a1051
13 changed files with 934 additions and 23 deletions
@@ -9,9 +9,10 @@ from sglang.srt.layers.attention.dsa.utils import (
INDEXER_K_CACHE_PRESHUFFLE_TILE,
aiter_can_use_preshuffle_paged_mqa,
)
from sglang.srt.utils import get_bool_env_var, is_hip
from sglang.srt.utils import get_bool_env_var, is_hip, is_xpu
_is_hip = is_hip()
_is_xpu = is_xpu()
_is_fp8_fnuz = is_fp8_fnuz()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
# aiter cp_gather kernel with preshuffle=True is only valid when the indexer
@@ -311,6 +312,11 @@ def _set_k_and_s_triton(
assert page_size % 16 == 0, (
f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}"
)
elif _is_xpu:
assert page_size in (
64,
128,
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
else:
assert page_size == 64
@@ -17,6 +17,7 @@ from sglang.srt.arg_groups.overrides import (
)
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
from sglang.srt.runtime_context import get_platform
logger = logging.getLogger(__name__)
@@ -279,6 +280,11 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
reserved_mem = max(reserved_mem, 10 * 1024)
# Reserve headroom for DeepEP all-to-all buffers on top of the floor.
reserved_mem += reserve_for_deepep_a2a_mb(server_args)
# XPU: oneDNN allocates scratch space for matmul when
# M is not a power-of-2-aligned value (e.g. M=2100). Reserve extra
# headroom so non-aligned prefill lengths don't hit OOM.
if get_platform().is_xpu:
reserved_mem += 2 * 1024
mem_fraction_static = (
round((gpu_mem - reserved_mem) / gpu_mem, 3)
@@ -244,6 +244,21 @@ def handle_model_specific_adjustments(server_args: Any):
run_post_process_pass(server_args, _dsa_kv_cache_dtype_default)
run_post_process_pass(server_args, _dsa_split_backend_resolution)
elif get_platform().is_xpu:
run_post_process_pass(server_args, _dsa_kv_cache_dtype_default)
run_post_process_pass(server_args, _dsa_split_backend_resolution)
# Disable fused topk (requires sgl-kernel ops not available on XPU)
if (
envs.SGLANG_DSA_FUSE_TOPK.is_set()
and envs.SGLANG_DSA_FUSE_TOPK.get()
):
logger.warning(
"Disabling fused topk for DeepSeek DSA on XPU (SGLANG_DSA_FUSE_TOPK=0). Not supported yet."
)
envs.SGLANG_DSA_FUSE_TOPK.set(False)
# Disable CUDA-JIT topk-v2 (TileLang/TVM-based, requires CUDA)
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
if cfg.enable_prefill_cp:
assert cfg.disaggregation_mode != "decode", (
"CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp."
@@ -143,6 +143,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
else:
overrides["page_size"] = 64
logger.warning("Setting page size to 64 for DeepSeek DSA.")
elif get_platform().is_xpu:
overrides["page_size"] = 128
logger.warning("Setting page size to 128 for DeepSeek DSA on XPU.")
else:
# DeepSeek V3/R1/V3.1
if get_platform().is_sm100:
+27 -2
View File
@@ -611,7 +611,14 @@ def _dsa_kv_cache_dtype_default(view: Any) -> dict:
return {}
if not is_deepseek_dsa(hf_config):
return {}
if get_platform().is_npu or get_platform().is_xpu:
if get_platform().is_npu:
return {}
if get_platform().is_xpu:
if view.kv_cache_dtype == "auto":
logger.warning(
"Setting KV cache dtype to bfloat16 for DeepSeek DSA on XPU."
)
return {"kv_cache_dtype": "bfloat16"}
return {}
import torch
@@ -689,8 +696,26 @@ def _dsa_split_backend_resolution(view: Any) -> dict:
return {}
if not is_deepseek_dsa(hf_config):
return {}
if get_platform().is_npu or get_platform().is_xpu:
if get_platform().is_npu:
return {}
if get_platform().is_xpu:
declared: Dict[str, Any] = {}
if view.dsa_prefill_backend is None:
declared["dsa_prefill_backend"] = "intel_xpu"
if view.dsa_decode_backend is None:
declared["dsa_decode_backend"] = "intel_xpu"
# sgl-kernel topk ops (the default) are CUDA-only; fall back to the
# torch-native topk implementation on XPU, unless the user already
# picked a different backend explicitly (e.g. "flashinfer").
if view.dsa_topk_backend == "sgl-kernel":
declared["dsa_topk_backend"] = "torch"
logger.warning(
"Set DSA backends for XPU: prefill=%s, decode=%s, topk=%s.",
declared.get("dsa_prefill_backend", view.dsa_prefill_backend),
declared.get("dsa_decode_backend", view.dsa_decode_backend),
declared.get("dsa_topk_backend", view.dsa_topk_backend),
)
return declared
import torch
@@ -70,6 +70,7 @@ from sglang.srt.layers.cp.utils import is_cp_v2_active
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_platform,
get_spec,
@@ -565,13 +566,10 @@ class DeepseekV4AttnBackend(
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
)
self.enable_deepseek_v4_fp4_indexer: bool = (
model_runner.server_args.enable_deepseek_v4_fp4_indexer
)
kernel = get_exec().kernel
self.enable_deepseek_v4_fp4_indexer = kernel.enable_deepseek_v4_fp4_indexer
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
self.dsv4_prefill_backend: str = getattr(
model_runner.server_args, "dsv4_prefill_backend", "auto"
)
self.dsv4_prefill_backend = getattr(kernel, "dsv4_prefill_backend", "auto")
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
if not get_platform().is_sm90:
raise ValueError(
@@ -52,6 +52,7 @@ from sglang.srt.utils import (
add_prefix,
ceil_align,
get_bool_env_var,
get_device_module,
is_cuda,
is_gfx95_supported,
is_hip,
@@ -104,6 +105,10 @@ if _is_cuda:
except ImportError as e:
deep_gemm = e
if _is_xpu:
from sgl_kernel import fp8_mqa_logits as sgl_fp8_mqa_logits
from sgl_kernel import fp8_paged_mqa_logits as sgl_fp8_paged_mqa_logits
if _use_aiter:
from aiter.ops.cache import indexer_k_quant_and_cache
@@ -187,7 +192,6 @@ def _broadcast_indexer_topk_from_rank0(
def rotate_activation(x: torch.Tensor) -> torch.Tensor:
# from sgl_kernel import hadamard_transform
if _is_hip:
from fast_hadamard_transform import hadamard_transform
elif _is_xpu:
@@ -815,6 +819,11 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
assert page_size == 1, (
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
)
elif _is_xpu:
assert page_size in (
64,
128,
), f"XPU DSA only supports page_size 64 or 128, got {page_size}"
else:
assert page_size == 64, "only support page size 64"
# NOTE(dark): this support extend/decode/decode+graph
@@ -960,6 +969,17 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
preshuffle=_use_aiter_preshuffle,
kv_block_size=block_kv,
)
elif _is_xpu:
logits = sgl_fp8_paged_mqa_logits(
q_fp8[:q_offset],
kv_cache_fp8,
weights[:q_offset],
seqlens_32_2d,
block_tables,
None,
max_seq_len,
clean_logits=False,
)
elif use_cute_dsl:
logits = cutedsl_paged_mqa_logits(
q_fp8,
@@ -1027,7 +1047,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
if cached_budget is not None:
return cached_budget
total_mem = torch.cuda.get_device_properties(device_index).total_memory
total_mem = get_device_module().get_device_properties(device_index).total_memory
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
mem_fraction_static = get_schedule().mem_fraction_static
@@ -1048,10 +1068,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
return static_budget
# Match the original free-memory guard: logits_bytes * 2 > free_mem.
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
# capped by the workload-independent serving-memory headroom.
free_mem, _ = torch.cuda.mem_get_info(device_index)
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
# Synchronizes the host; cache the result capped by serving-memory headroom.
if _is_xpu:
# On XPU, use total_mem budget as the free-memory estimate;
# dynamic free-memory query is not supported the same way as CUDA.
# TODO Use torch.xpu.mem_get_info() when available (planned end of 2026).
budget_bytes = static_budget
else:
free_mem, _ = torch.cuda.mem_get_info(device_index)
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
budget_bytes = max(1, budget_bytes)
self._mqa_logits_budget_bytes[device_index] = budget_bytes
@@ -1105,7 +1130,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
)
else:
assert page_size == 64, "only support page size 64"
if _is_xpu:
assert page_size in (
64,
128,
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
else:
assert page_size == 64, "only support page size 64"
assert len(weights.shape) == 3
assert (
@@ -1186,6 +1217,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
ke,
clean_logits=False,
)
elif _is_xpu:
logits = sgl_fp8_mqa_logits(
q_fp8[:q_offset],
kv_fp8,
weights[:q_offset],
ks,
ke,
clean_logits=False,
)
else:
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
q_fp8[:q_offset], weights[:q_offset]
@@ -1242,6 +1282,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
ke[start:end],
clean_logits=False,
)
elif _is_xpu:
logits_chunk = sgl_fp8_mqa_logits(
q_fp8[start:end],
kv_fp8,
weights[start:end],
ks[start:end],
ke[start:end],
clean_logits=False,
)
else:
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
q_fp8[start:end], weights[start:end]
@@ -1388,7 +1437,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
page_size = get_token_to_kv_pool().page_size
assert page_size == 64, "only support page size 64"
if _is_xpu:
assert page_size in (
64,
128,
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
else:
assert page_size == 64, "only support page size 64"
assert len(weights.shape) == 3
weights = weights.squeeze(-1)
k_fp8_list = []
@@ -1871,7 +1926,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
else:
weights = self._get_logits_head_gate(x_for_gate, q_scale)
if _is_cuda or _is_hip:
if _is_cuda or _is_hip or _is_xpu:
# In piecewise/breakable CUDA graph, any access to seq_lens_cpu
# creates a Dynamo shape guard. These graph modes never have empty
# batches.
@@ -101,6 +101,7 @@ from sglang.srt.utils import (
is_cuda,
is_gfx95_supported,
is_hip,
is_xpu,
print_warning_once,
)
@@ -158,6 +159,7 @@ def materialize_full_kv_cp(
_is_hip = is_hip()
_is_xpu = is_xpu()
if _is_hip:
from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices
@@ -176,6 +178,11 @@ if _is_hip:
print(
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
)
elif _is_xpu:
from sgl_kernel.flash_attn import (
flash_attn_varlen_func,
flash_attn_with_kvcache,
)
else:
from sglang.kernels.ops.attention.flash_attention import (
flash_attn_varlen_func,
@@ -313,6 +320,7 @@ _DSA_IMPL_T: TypeAlias = Literal[
"fa3",
"tilelang",
"trtllm",
"intel_xpu",
]
@@ -452,7 +460,10 @@ class DeepseekSparseAttnBackend(
"Disabling fused DSA top-k for IndexShare under PD disaggregation."
)
self.device_capability = torch.cuda.get_device_capability()
if _is_xpu:
self.device_capability = (0, 0)
else:
self.device_capability = torch.cuda.get_device_capability()
self.device_sm_major = self.device_capability[0]
self.kv_cache_dtype = model_runner.kv_cache_dtype
@@ -523,7 +534,7 @@ class DeepseekSparseAttnBackend(
self._q8kv8_born_q_buf: Optional[torch.Tensor] = None
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
self._q8kv8_born_q_sentinel: Optional[torch.Tensor] = None
self._q8kv8_born_q_tbo = model_runner.server_args.enable_two_batch_overlap
self._q8kv8_born_q_tbo = get_exec().overlap.enable_two_batch_overlap
from sglang.kernels.ops.attention.flash_mla_sm120 import (
_validate_flashinfer_sparse_mla_backend,
@@ -1139,7 +1150,7 @@ class DeepseekSparseAttnBackend(
cache_seqlens=dsa_cache_seqlens_int32,
seq_len_q=1,
)
if use_flashmla_kv
if use_flashmla_kv and not _is_xpu
else None
),
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
@@ -2309,6 +2320,15 @@ class DeepseekSparseAttnBackend(
page_table_1=page_table_1,
layer=layer,
)
elif dsa_impl == "intel_xpu":
return self._forward_intel_xpu_dense_prefill(
q_nope=q_nope,
q_rope=q_rope,
kv_cache=kv_cache,
v_head_dim=layer.v_head_dim,
sm_scale=layer.scaling,
metadata=metadata,
)
else:
raise ValueError(
f"Unsupported {dsa_impl = } for forward_extend. Consider using an other attention backend."
@@ -2492,6 +2512,15 @@ class DeepseekSparseAttnBackend(
metadata=metadata,
bs=forward_batch.batch_size,
)
elif self.dsa_decode_impl == "intel_xpu":
return self._forward_intel_xpu_sparse_decode(
q_nope=q_nope,
q_rope=q_rope,
kv_cache=kv_cache,
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
else:
assert False, f"Unsupported {dsa_impl = }"
@@ -3115,6 +3144,123 @@ class DeepseekSparseAttnBackend(
d_v=v_head_dim,
)
def _forward_intel_xpu_sparse_decode(
self,
q_nope: torch.Tensor,
q_rope: torch.Tensor,
kv_cache: torch.Tensor,
page_table_1: torch.Tensor,
sm_scale: float,
v_head_dim: int,
) -> torch.Tensor:
"""Sparse decode for XPU using flash_mla_decode with gathered KV.
Gathers the sparse KV tokens selected by the indexer into a contiguous
buffer organized as virtual pages, then runs flash_mla_decode on it.
"""
from sgl_kernel import flash_mla_decode, flash_mla_get_workspace_size
B = q_nope.shape[0]
TOPK = page_table_1.shape[1]
D_ckv = kv_cache.shape[-1]
GATHER_PAGE_SIZE = 16
assert TOPK % GATHER_PAGE_SIZE == 0, (
f"TOPK {TOPK} must be a multiple of GATHER_PAGE_SIZE {GATHER_PAGE_SIZE}"
)
NUM_PAGES = TOPK // GATHER_PAGE_SIZE
# Count valid tokens per batch (non -1 entries)
valid_counts = (page_table_1 >= 0).sum(dim=1).to(torch.int32)
# Gather KV tokens: replace -1 with 0 for safe indexing, zero-fill after
safe_indices = page_table_1.clamp(min=0)
# kv_cache is [total_tokens, D_ckv], index with flat indices
gathered_kv = kv_cache[safe_indices.view(-1)].view(B, TOPK, D_ckv)
# Zero out invalid entries
invalid_mask = page_table_1 < 0 # [B, TOPK]
gathered_kv[invalid_mask] = 0
# Reshape to pages: [B * NUM_PAGES, GATHER_PAGE_SIZE, D_ckv]
gathered_kv_paged = gathered_kv.view(B * NUM_PAGES, GATHER_PAGE_SIZE, D_ckv)
# Identity page table: batch i → pages [i*NUM_PAGES, ..., (i+1)*NUM_PAGES-1]
identity_page_table = torch.arange(
B * NUM_PAGES, device=q_nope.device, dtype=torch.int32
).view(B, NUM_PAGES)
# Workspace
ws_size = flash_mla_get_workspace_size(
TOPK, B, q_nope.shape[1], GATHER_PAGE_SIZE
)
if self.workspace_buffer is None:
self.workspace_buffer = torch.empty(
ws_size, device=q_nope.device, dtype=torch.uint8
)
elif self.workspace_buffer.numel() < ws_size:
self.workspace_buffer.resize_(ws_size)
o = flash_mla_decode(
q_nope,
q_rope,
gathered_kv_paged,
valid_counts,
identity_page_table,
self.workspace_buffer,
sm_scale,
)
return o
def _forward_intel_xpu_dense_prefill(
self,
q_nope: torch.Tensor,
q_rope: torch.Tensor,
kv_cache: torch.Tensor,
v_head_dim: int,
sm_scale: float,
metadata: DSAMetadata,
) -> torch.Tensor:
"""Dense prefill for XPU using flash_mla_prefill.
Runs full causal MLA attention over all KV positions (no sparse top-K
selection). This is the XPU equivalent of the CUDA flashmla_kv /
flashmla_sparse prefill paths.
"""
from sgl_kernel import flash_mla_prefill, flash_mla_prefill_get_workspace_size
D_ckv = kv_cache.shape[-1]
# kv_cache: (N_tokens, 1, D_ckv) from MLATokenToKVPool → (N_pages, page_size, D_ckv)
kv_paged = kv_cache.view(-1, self.real_page_size, D_ckv)
block_table = metadata.real_page_table # (B, max_pages), page-indexed int32
seq_lens_k = metadata.cache_seqlens_int32 # (B,) total KV lengths
cu_seqlens_q = metadata.cu_seqlens_q # (B+1,) cumulative Q lengths
max_seqlen_q = metadata.max_seq_len_q
max_seq_len_k = int(seq_lens_k.max().item())
ws_size = flash_mla_prefill_get_workspace_size(
max_seq_len_k, seq_lens_k.shape[0]
)
if self.workspace_buffer is None:
self.workspace_buffer = torch.empty(
ws_size, device=q_nope.device, dtype=torch.uint8
)
elif self.workspace_buffer.numel() < ws_size:
self.workspace_buffer.resize_(ws_size)
return flash_mla_prefill(
q_nope,
q_rope,
kv_paged,
cu_seqlens_q,
seq_lens_k,
max_seqlen_q,
block_table,
self.workspace_buffer,
sm_scale,
causal=True,
num_kv_splits=1,
)
def _forward_aiter(
self,
q_all: torch.Tensor,
@@ -471,9 +471,17 @@ class RotaryEmbedding(BaseFusedOp):
)
return query, key
else:
# Use fallback kernel of 'rotary_embedding'
self._match_cos_sin_cache_dtype(query)
return torch.ops.sgl_kernel.rotary_embedding(
# Use fallback kernel of 'rotary_embedding'.
# The kernel requires 3D tensors (batch, num_heads, head_size);
# add a num_heads=1 dim for 2D tensors (e.g. DSA indexer k_rope).
q_2d = query.dim() == 2
k_2d = key.dim() == 2
if q_2d:
query = query.view(query.shape[0], -1, self.head_size)
if k_2d:
key = key.view(key.shape[0], -1, self.head_size)
q_out, k_out = torch.ops.sgl_kernel.rotary_embedding(
positions,
query,
key,
@@ -481,6 +489,11 @@ class RotaryEmbedding(BaseFusedOp):
self.cos_sin_cache,
self.is_neox_style,
)
if q_2d:
q_out = q_out.view(q_out.shape[0], -1)
if k_2d:
k_out = k_out.view(k_out.shape[0], -1)
return q_out, k_out
class LinearScalingRotaryEmbedding(RotaryEmbedding):
@@ -80,6 +80,7 @@ from sglang.srt.utils import (
is_float4_e2m1fn_x2,
is_hip,
is_npu,
is_xpu,
next_power_of_2,
)
from sglang.srt.utils.async_probe import (
@@ -4722,6 +4723,11 @@ class DSATokenToKVPool(MLATokenToKVPool):
assert self.page_size == 1, (
f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
)
elif is_xpu():
assert self.page_size in (
64,
128,
), f"XPU DSA requires page_size 64 or 128, got {self.page_size}"
else:
assert self.page_size == 64
self.index_key_cache = self._create_index_key_cache()
+1
View File
@@ -1843,6 +1843,7 @@ class ServerArgs:
Arg(
help="DSA indexer top-k backend for the target model. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
choices=["sgl-kernel", "torch", "flashinfer"],
resolvable=True,
),
NS("exec.kernel"),
] = "sgl-kernel"