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:
co-authored by
Copilot
Ma Mingfei
parent
39a80354aa
commit
c4e52a1051
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user