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,8 +1068,13 @@ 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.
# 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)
@@ -1104,6 +1129,12 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
assert page_size == 1, (
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
)
else:
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"
@@ -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,6 +1437,12 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
page_size = get_token_to_kv_pool().page_size
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)
@@ -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,6 +460,9 @@ class DeepseekSparseAttnBackend(
"Disabling fused DSA top-k for IndexShare under PD disaggregation."
)
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"
@@ -94,6 +94,7 @@ class TestModelOverridableWhitelist(CustomTestCase):
"kv_cache_dtype",
"dsa_prefill_backend",
"dsa_decode_backend",
"dsa_topk_backend",
"prefill_attention_backend",
"decode_attention_backend",
"flashinfer_allreduce_fusion_backend",
+636
View File
@@ -0,0 +1,636 @@
"""XPU unit tests for the DSA (Dynamic Sparse Attention) indexer.
Mirrors test/registered/kernels/test_dsa_indexer.py for XPU, covering:
- Indexer creation and basic forward pass (extend + decode modes)
- rotate_activation (Hadamard transform, PyTorch-native fallback on XPU)
- FP8 act_quant dispatch on XPU
- topk selection (torch.topk fallback on XPU, TOPK_V2 disabled)
- HybridAttnBackend: init_forward_metadata + get_indexer_metadata routing
- RotaryEmbedding.forward_xpu with 2D k_rope (DSA indexer single-head key)
NOTE: A full end-to-end GLM5.1 integration test requires the reduced
GlmMoeDsaForCausalLM model which is not publicly available, so it is
not included here. These tests use synthetic tensors and mock runners,
matching the style of the CUDA counterpart.
"""
import unittest
from typing import List, Tuple
from unittest.mock import patch
import torch
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_xpu_ci
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.dsa.dsa_indexer import (
BaseIndexerMetadata,
Indexer,
rotate_activation,
)
from sglang.srt.layers.attention.dsa_backend import (
DeepseekSparseAttnBackend,
)
from sglang.srt.layers.layernorm import LayerNorm
from sglang.srt.layers.linear import LinearBase
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase
register_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu")
# Configuration matching GLM5.1 index head dimensions on XPU
DEFAULT_CONFIG = {
"device": "xpu",
"dtype": torch.bfloat16,
"kv_cache_dtype": torch.float8_e4m3fn,
"context_len": 2048,
"max_bs": 64,
"hidden_size": 5120,
"index_n_heads": 32,
"index_head_dim": 128,
"rope_head_dim": 64,
"index_topk": 64,
"q_lora_rank": 1536,
"kv_lora_rank": 512,
"qk_rope_head_dim": 64,
"qk_nope_head_dim": 128,
"max_position_embeddings": 163840,
"rope_theta": 10000.0,
"layer_id": 0,
"page_size": 128, # XPU uses page_size=128
}
class MockIndexerMetadata(BaseIndexerMetadata):
"""Minimal mock of BaseIndexerMetadata for XPU testing."""
def __init__(self, batch_size, seq_lens, device="xpu"):
self.batch_size = batch_size
self.seq_lens = seq_lens
self.device = device
def get_seqlens_int32(self) -> torch.Tensor:
return torch.tensor(self.seq_lens, dtype=torch.int32, device=self.device)
def get_page_table_64(self) -> torch.Tensor:
max_seq_len = max(self.seq_lens)
num_blocks = (max_seq_len + 63) // 64
page_table = torch.zeros(
(self.batch_size, num_blocks), dtype=torch.int32, device=self.device
)
for i in range(self.batch_size):
n = (self.seq_lens[i] + 63) // 64
page_table[i, :n] = torch.arange(n, device=self.device)
return page_table
def get_page_table_1(self) -> torch.Tensor:
max_seq_len = max(self.seq_lens)
page_table = torch.zeros(
(self.batch_size, max_seq_len), dtype=torch.int32, device=self.device
)
for i in range(self.batch_size):
n = self.seq_lens[i]
page_table[i, :n] = torch.arange(n, device=self.device)
return page_table
def get_seqlens_expanded(self) -> torch.Tensor:
result = []
for seq_len in self.seq_lens:
result.extend(range(1, seq_len + 1))
return torch.tensor(result, dtype=torch.int32, device=self.device)
def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]:
ks_list, ke_list = [], []
k_offset = 0
for seq_len in self.seq_lens:
ks = torch.full((seq_len,), k_offset, dtype=torch.int32, device=self.device)
ke = torch.arange(
k_offset + 1,
k_offset + seq_len + 1,
dtype=torch.int32,
device=self.device,
)
ks_list.append(ks)
ke_list.append(ke)
k_offset += seq_len
return torch.cat(ks_list), torch.cat(ke_list)
def get_indexer_seq_len_cpu(self) -> torch.Tensor:
return torch.tensor(self.seq_lens, dtype=torch.int32, device="cpu")
def get_indexer_seq_len(self) -> torch.Tensor:
return torch.tensor(self.seq_lens, dtype=torch.int32, device=self.device)
def get_dsa_extend_len_cpu(self) -> List[int]:
return list(self.seq_lens)
def get_token_to_batch_idx(self) -> torch.Tensor:
result = []
for batch_idx, seq_len in enumerate(self.seq_lens):
result.extend([batch_idx] * seq_len)
return torch.tensor(result, dtype=torch.int32, device=self.device)
def topk_transform(self, logits, topk, **kwargs):
return torch.topk(logits, k=topk, dim=-1).indices
class MockModelRunner:
def __init__(self, config=None):
cfg = {**DEFAULT_CONFIG, **(config or {})}
self.device = cfg["device"]
self.config = cfg
self.dtype = cfg["dtype"]
self.kv_cache_dtype = cfg["kv_cache_dtype"]
self.is_hybrid_swa = False
self.is_draft_worker = False
hf_config = type(
"HfConfig",
(),
{
"architectures": ["GlmMoeDsaForCausalLM"],
"index_topk": cfg["index_topk"],
"index_head_dim": cfg["index_head_dim"],
"index_n_heads": cfg["index_n_heads"],
},
)()
self.model_config = type(
"ModelConfig",
(),
{
"context_len": cfg["context_len"],
"is_multimodal": False,
"attention_arch": AttentionArch.MLA,
"num_attention_heads": 128,
"kv_lora_rank": cfg["kv_lora_rank"],
"qk_rope_head_dim": cfg["qk_rope_head_dim"],
"qk_nope_head_dim": cfg["qk_nope_head_dim"],
"hf_config": hf_config,
},
)()
self.sliding_window_size = None
self.page_size = cfg["page_size"]
max_batch_size = cfg["max_bs"]
max_context_len = cfg["context_len"]
self.req_to_token_pool = type(
"TokenPool",
(),
{
"size": max_batch_size,
"req_to_token": torch.zeros(
max_batch_size,
max_context_len,
dtype=torch.int32,
device=self.device,
),
},
)()
self.token_to_kv_pool = DSATokenToKVPool(
size=max_batch_size * max_context_len,
page_size=cfg["page_size"],
dtype=cfg["kv_cache_dtype"],
kv_lora_rank=cfg["kv_lora_rank"],
qk_rope_head_dim=cfg["qk_rope_head_dim"],
layer_num=1,
device=self.device,
index_head_dim=cfg["index_head_dim"],
enable_memory_saver=False,
kv_cache_dim=cfg["kv_lora_rank"] + cfg["qk_rope_head_dim"],
)
# XPU-specific DSA backend settings (mirrors server_args.py XPU section)
self.server_args = type(
"ServerArgs",
(),
{
"kv_cache_dtype": "auto",
"speculative_eagle_topk": None,
"speculative_num_draft_tokens": 0,
"enable_deterministic_inference": False,
"dsa_prefill_backend": "intel_xpu",
"dsa_decode_backend": "intel_xpu",
"dsa_topk_backend": "torch", # XPU uses torch.topk fallback
"dsa_paged_mqa_logits_backend": "auto",
"disaggregation_mode": "null",
"enable_two_batch_overlap": False,
},
)()
self.hisparse_coordinator = None
@unittest.skipIf(not torch.xpu.is_available(), "XPU is required")
class TestDSAIndexerXPU(CustomTestCase):
"""Tests for the DSA indexer on XPU, mirroring test_dsa_indexer.py."""
@classmethod
def setUpClass(cls):
server_args = ServerArgs(model_path="dummy")
server_args.enable_dp_attention = False
server_args.dsa_prefill_backend = "intel_xpu"
server_args.dsa_decode_backend = "intel_xpu"
server_args.dsa_topk_backend = "torch"
# Disable CUDA-only JIT topk-v2 (TileLang requires CUDA_HOME)
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
set_global_server_args_for_scheduler(server_args)
def setUp(self):
self.batch_size = 2
self.seq_len = 128
self.config = DEFAULT_CONFIG.copy()
self.device = "xpu"
self.dtype = torch.bfloat16
def _init_model_runner(self, config_override=None):
cfg = {**self.config, **(config_override or {})}
self.model_runner = MockModelRunner(cfg)
self.backend = DeepseekSparseAttnBackend(self.model_runner)
def _create_indexer(self, **kwargs):
params = {
"hidden_size": self.config["hidden_size"],
"index_n_heads": self.config["index_n_heads"],
"index_head_dim": self.config["index_head_dim"],
"rope_head_dim": self.config["rope_head_dim"],
"index_topk": self.config["index_topk"],
"q_lora_rank": self.config["q_lora_rank"],
"max_position_embeddings": self.config["max_position_embeddings"],
"rope_theta": self.config["rope_theta"],
"layer_id": self.config["layer_id"],
"scale_fmt": "ue8m0",
"block_size": 128,
"quant_config": None,
# GLM5.1 has indexer_rope_interleave=True → is_neox_style=False.
# The XPU sgl_kernel.rotary_embedding 3D+neox path returns 4D output,
# so use is_neox_style=False to match the real model config.
"is_neox_style": False,
}
params.update(kwargs)
torch.set_default_dtype(self.dtype)
with torch.device(self.device):
indexer = Indexer(**params)
indexer = indexer.to(device=self.device)
for name, module in indexer.named_modules():
if isinstance(module, LinearBase) and not isinstance(module, LayerNorm):
if "weights_proj" not in name:
module.to(dtype=self.dtype)
return indexer
def _create_forward_batch(self, mode, batch_size=None, seq_len=None):
batch_size = batch_size or self.batch_size
seq_len = seq_len or self.seq_len
if mode == ForwardMode.EXTEND:
forward_batch = ForwardBatch(
batch_size=batch_size,
input_ids=torch.randint(
0, 100, (batch_size, seq_len), device=self.device
),
out_cache_loc=torch.arange(batch_size * seq_len, device=self.device),
seq_lens_sum=batch_size * seq_len,
forward_mode=mode,
req_pool_indices=torch.arange(batch_size, device=self.device),
seq_lens=torch.tensor([seq_len] * batch_size, device=self.device),
seq_lens_cpu=torch.tensor([seq_len] * batch_size, device="cpu"),
extend_prefix_lens=torch.zeros(
batch_size, device=self.device, dtype=torch.int32
),
extend_prefix_lens_cpu=torch.zeros(
batch_size, device="cpu", dtype=torch.int32
),
extend_seq_lens=torch.tensor(
[seq_len] * batch_size, device=self.device
),
extend_seq_lens_cpu=torch.tensor([seq_len] * batch_size, device="cpu"),
)
else: # DECODE
total_len = seq_len + 1
forward_batch = ForwardBatch(
batch_size=batch_size,
input_ids=torch.randint(0, 100, (batch_size, 1), device=self.device),
out_cache_loc=torch.arange(
batch_size * seq_len, batch_size * total_len, device=self.device
),
seq_lens_sum=batch_size * total_len,
forward_mode=mode,
req_pool_indices=torch.arange(batch_size, device=self.device),
seq_lens=torch.tensor([total_len] * batch_size, device=self.device),
seq_lens_cpu=torch.tensor([total_len] * batch_size, device="cpu"),
)
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
set_forward_context(ForwardContext(attn_backend=self.backend))
page_size = self.model_runner.page_size
for i in range(batch_size):
for j in range(seq_len + (0 if mode == ForwardMode.EXTEND else 1)):
self.model_runner.req_to_token_pool.req_to_token[i, j] = (
i * seq_len + j + page_size
)
return forward_batch
def _verify_topk_output(self, topk_indices, batch_size, q_len, topk):
self.assertIsNotNone(topk_indices)
self.assertEqual(topk_indices.device.type, "xpu")
self.assertEqual(len(topk_indices.shape), 2)
self.assertEqual(topk_indices.shape[0], batch_size * q_len)
self.assertGreaterEqual(topk_indices.shape[1], topk)
# ------------------------------------------------------------------
# Test: indexer creation
# ------------------------------------------------------------------
def test_indexer_basic_creation(self):
"""Test basic Indexer instantiation on XPU."""
self._init_model_runner()
indexer = self._create_indexer()
self.assertEqual(indexer.hidden_size, self.config["hidden_size"])
self.assertEqual(indexer.n_heads, self.config["index_n_heads"])
self.assertEqual(indexer.head_dim, self.config["index_head_dim"])
self.assertEqual(indexer.rope_head_dim, self.config["rope_head_dim"])
self.assertEqual(indexer.index_topk, self.config["index_topk"])
# ------------------------------------------------------------------
# Test: rotate_activation (Hadamard, XPU uses PyTorch-native fallback)
# ------------------------------------------------------------------
def test_rotate_activation_power_of_two(self):
"""rotate_activation should work for power-of-2 sizes on XPU (PyTorch fallback)."""
for hidden_size in [64, 128, 256]:
x = torch.randn(16, hidden_size, dtype=torch.bfloat16, device=self.device)
out = rotate_activation(x)
self.assertEqual(out.shape, x.shape)
self.assertEqual(out.dtype, torch.bfloat16)
self.assertEqual(out.device.type, "xpu")
def test_rotate_activation_invalid_size(self):
"""rotate_activation should raise for non-power-of-2 sizes."""
x = torch.randn(16, 129, dtype=torch.bfloat16, device=self.device)
with self.assertRaises(AssertionError):
rotate_activation(x)
# ------------------------------------------------------------------
# Test: indexer forward — extend mode
# ------------------------------------------------------------------
@patch("sglang.kernels.ops.attention.dsa.triton_kernel.act_quant")
def test_forward_extend_mode(self, mock_act_quant):
"""Indexer forward in EXTEND mode calls sgl_kernel.fp8_mqa_logits on XPU."""
def _mock_quant(x, block_size=128, scale_fmt=None, *args, **kwargs):
# Match real act_quant output: scale shape is (*x.shape[:-1], n_groups)
n_groups = x.shape[-1] // block_size
scale_shape = x.shape[:-1] + (n_groups,)
return x.to(torch.float8_e4m3fn), torch.ones(
scale_shape, dtype=torch.float32, device=x.device
)
mock_act_quant.side_effect = _mock_quant
self._init_model_runner()
indexer = self._create_indexer()
forward_batch = self._create_forward_batch(ForwardMode.EXTEND)
total_tokens = self.batch_size * self.seq_len
hidden_states = torch.randn(
total_tokens,
self.config["hidden_size"],
dtype=self.dtype,
device=self.device,
)
q_lora = torch.randn(
total_tokens,
self.config["q_lora_rank"],
dtype=self.dtype,
device=self.device,
)
positions = torch.arange(total_tokens, device=self.device)
with patch.object(
self.backend,
"get_indexer_metadata",
return_value=MockIndexerMetadata(
self.batch_size, [self.seq_len] * self.batch_size, device=self.device
),
):
topk_indices = indexer(
x=hidden_states,
q_lora=q_lora,
positions=positions,
forward_batch=forward_batch,
layer_id=self.config["layer_id"],
)
self._verify_topk_output(
topk_indices, self.batch_size, self.seq_len, self.config["index_topk"]
)
# ------------------------------------------------------------------
# Test: indexer forward — decode mode
# ------------------------------------------------------------------
@patch("sglang.kernels.ops.attention.dsa.triton_kernel.act_quant")
def test_forward_decode_mode(self, mock_act_quant):
"""Indexer forward in DECODE mode calls sgl_kernel.fp8_paged_mqa_logits on XPU."""
def _mock_quant(x, block_size=128, scale_fmt=None, *args, **kwargs):
# Match real act_quant output: scale shape is (*x.shape[:-1], n_groups)
n_groups = x.shape[-1] // block_size
scale_shape = x.shape[:-1] + (n_groups,)
return x.to(torch.float8_e4m3fn), torch.ones(
scale_shape, dtype=torch.float32, device=x.device
)
mock_act_quant.side_effect = _mock_quant
self._init_model_runner()
indexer = self._create_indexer()
forward_batch = self._create_forward_batch(ForwardMode.DECODE)
hidden_states = torch.randn(
self.batch_size,
self.config["hidden_size"],
dtype=self.dtype,
device=self.device,
)
q_lora = torch.randn(
self.batch_size,
self.config["q_lora_rank"],
dtype=self.dtype,
device=self.device,
)
positions = torch.arange(self.batch_size, device=self.device)
with patch.object(
self.backend,
"get_indexer_metadata",
return_value=MockIndexerMetadata(
self.batch_size,
[self.seq_len + 1] * self.batch_size,
device=self.device,
),
):
topk_indices = indexer(
x=hidden_states,
q_lora=q_lora,
positions=positions,
forward_batch=forward_batch,
layer_id=self.config["layer_id"],
)
self._verify_topk_output(
topk_indices, self.batch_size, 1, self.config["index_topk"]
)
# ------------------------------------------------------------------
# Test: skip logits when seq_len <= index_topk
# ------------------------------------------------------------------
def test_skip_logits_short_sequence(self):
"""Indexer returns dense topk when seq_len <= index_topk (no FP8 scoring).
When all KV positions fit within index_topk, the indexer skips the
expensive FP8 MQA logit computation (EXTEND mode only) and returns
sequential dense indices covering all KV positions.
"""
short_seq = self.config["index_topk"] // 2 # 32 < index_topk=64
self._init_model_runner()
indexer = self._create_indexer()
# EXTEND mode is required: _should_skip_logits_computation only triggers
# for extend (prefill) batches, not decode.
forward_batch = self._create_forward_batch(
ForwardMode.EXTEND, seq_len=short_seq
)
total_tokens = self.batch_size * short_seq
hidden_states = torch.randn(
total_tokens,
self.config["hidden_size"],
dtype=self.dtype,
device=self.device,
)
q_lora = torch.randn(
total_tokens,
self.config["q_lora_rank"],
dtype=self.dtype,
device=self.device,
)
positions = torch.arange(total_tokens, device=self.device)
# seq_len (32) < index_topk (64): skip FP8 scoring, use dense fallback.
# Returns sequential topk indices, NOT None.
with patch.object(
self.backend,
"get_indexer_metadata",
return_value=MockIndexerMetadata(
self.batch_size,
[short_seq] * self.batch_size,
device=self.device,
),
):
topk_indices = indexer(
x=hidden_states,
q_lora=q_lora,
positions=positions,
forward_batch=forward_batch,
layer_id=self.config["layer_id"],
)
# Dense fallback: indices are returned (not None), all within [0, index_topk)
self.assertIsNotNone(topk_indices)
self.assertEqual(topk_indices.device.type, "xpu")
self.assertGreaterEqual(topk_indices.shape[-1], self.config["index_topk"])
# ------------------------------------------------------------------
# Test: RotaryEmbedding.forward_xpu with 2D k_rope (DSA indexer path)
# ------------------------------------------------------------------
def test_rotary_embedding_2d_key(self):
"""forward_xpu must handle 2D (N, head_size) k_rope from the DSA indexer.
The DSA indexer creates a RotaryEmbedding with head_size=rope_head_dim
(64) and a single KV head. Its k_rope tensor is 2D (N, 64), not 3D.
The XPU fallback path (sgl_kernel.rotary_embedding) requires 3D input;
forward_xpu must unsqueeze/squeeze transparently.
"""
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
rope_head_dim = self.config["rope_head_dim"] # 64
num_tokens = 8
max_position = self.config["max_position_embeddings"]
rope = RotaryEmbedding(
head_size=rope_head_dim,
rotary_dim=rope_head_dim,
max_position_embeddings=max_position,
base=self.config["rope_theta"],
is_neox_style=False, # GLM5.1 has indexer_rope_interleave=True → is_neox=False
dtype=self.dtype,
).to(self.device)
positions = torch.arange(num_tokens, device=self.device)
# 2D query and key — this is what the DSA indexer passes
query_2d = torch.randn(
num_tokens, rope_head_dim, dtype=self.dtype, device=self.device
)
key_2d = torch.randn(
num_tokens, rope_head_dim, dtype=self.dtype, device=self.device
)
q_out, k_out = rope.forward_xpu(positions, query_2d, key_2d, rope_head_dim)
self.assertEqual(q_out.shape, query_2d.shape)
self.assertEqual(k_out.shape, key_2d.shape)
self.assertEqual(q_out.device.type, "xpu")
# ------------------------------------------------------------------
# Test: XPU uses a single DeepseekSparseAttnBackend for both modes
# ------------------------------------------------------------------
def test_unified_dsa_backend_both_modes(self):
"""On XPU, DeepseekSparseAttnBackend handles both prefill and decode.
Verifies that after _init_model_runner the server_args use the same
"intel_xpu" impl for both dsa_prefill_backend and dsa_decode_backend,
matching the unified flash_mla_prefill + flash_mla_decode path in
sgl-kernel-xpu (no HybridAttnBackend required).
"""
self._init_model_runner()
sa = self.model_runner.server_args
self.assertEqual(sa.dsa_prefill_backend, "intel_xpu")
self.assertEqual(sa.dsa_decode_backend, "intel_xpu")
# The backend created for both forward modes should be DeepseekSparseAttnBackend
backend = self.backend
self.assertIsInstance(backend, DeepseekSparseAttnBackend)
# Its prefill impl should also be "intel_xpu"
self.assertEqual(
backend.dsa_prefill_impl,
"intel_xpu",
"Expected prefill impl 'intel_xpu'; got {}".format(
backend.dsa_prefill_impl
),
)
if __name__ == "__main__":
unittest.main()