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,
|
INDEXER_K_CACHE_PRESHUFFLE_TILE,
|
||||||
aiter_can_use_preshuffle_paged_mqa,
|
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_hip = is_hip()
|
||||||
|
_is_xpu = is_xpu()
|
||||||
_is_fp8_fnuz = is_fp8_fnuz()
|
_is_fp8_fnuz = is_fp8_fnuz()
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_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
|
# 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, (
|
assert page_size % 16 == 0, (
|
||||||
f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}"
|
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:
|
else:
|
||||||
assert page_size == 64
|
assert page_size == 64
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
|
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
|
||||||
|
from sglang.srt.runtime_context import get_platform
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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)
|
reserved_mem = max(reserved_mem, 10 * 1024)
|
||||||
# Reserve headroom for DeepEP all-to-all buffers on top of the floor.
|
# Reserve headroom for DeepEP all-to-all buffers on top of the floor.
|
||||||
reserved_mem += reserve_for_deepep_a2a_mb(server_args)
|
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 = (
|
mem_fraction_static = (
|
||||||
round((gpu_mem - reserved_mem) / gpu_mem, 3)
|
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_kv_cache_dtype_default)
|
||||||
run_post_process_pass(server_args, _dsa_split_backend_resolution)
|
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:
|
if cfg.enable_prefill_cp:
|
||||||
assert cfg.disaggregation_mode != "decode", (
|
assert cfg.disaggregation_mode != "decode", (
|
||||||
"CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp."
|
"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:
|
else:
|
||||||
overrides["page_size"] = 64
|
overrides["page_size"] = 64
|
||||||
logger.warning("Setting page size to 64 for DeepSeek DSA.")
|
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:
|
else:
|
||||||
# DeepSeek V3/R1/V3.1
|
# DeepSeek V3/R1/V3.1
|
||||||
if get_platform().is_sm100:
|
if get_platform().is_sm100:
|
||||||
|
|||||||
@@ -611,7 +611,14 @@ def _dsa_kv_cache_dtype_default(view: Any) -> dict:
|
|||||||
return {}
|
return {}
|
||||||
if not is_deepseek_dsa(hf_config):
|
if not is_deepseek_dsa(hf_config):
|
||||||
return {}
|
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 {}
|
return {}
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -689,8 +696,26 @@ def _dsa_split_backend_resolution(view: Any) -> dict:
|
|||||||
return {}
|
return {}
|
||||||
if not is_deepseek_dsa(hf_config):
|
if not is_deepseek_dsa(hf_config):
|
||||||
return {}
|
return {}
|
||||||
if get_platform().is_npu or get_platform().is_xpu:
|
if get_platform().is_npu:
|
||||||
return {}
|
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
|
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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_platform,
|
get_platform,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -565,13 +566,10 @@ class DeepseekV4AttnBackend(
|
|||||||
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
|
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
|
||||||
)
|
)
|
||||||
|
|
||||||
self.enable_deepseek_v4_fp4_indexer: bool = (
|
kernel = get_exec().kernel
|
||||||
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
self.enable_deepseek_v4_fp4_indexer = kernel.enable_deepseek_v4_fp4_indexer
|
||||||
)
|
|
||||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
|
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
|
||||||
self.dsv4_prefill_backend: str = getattr(
|
self.dsv4_prefill_backend = getattr(kernel, "dsv4_prefill_backend", "auto")
|
||||||
model_runner.server_args, "dsv4_prefill_backend", "auto"
|
|
||||||
)
|
|
||||||
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
|
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
|
||||||
if not get_platform().is_sm90:
|
if not get_platform().is_sm90:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -52,6 +52,7 @@ from sglang.srt.utils import (
|
|||||||
add_prefix,
|
add_prefix,
|
||||||
ceil_align,
|
ceil_align,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
|
get_device_module,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
@@ -104,6 +105,10 @@ if _is_cuda:
|
|||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
deep_gemm = 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:
|
if _use_aiter:
|
||||||
from aiter.ops.cache import indexer_k_quant_and_cache
|
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:
|
def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
||||||
# from sgl_kernel import hadamard_transform
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
from fast_hadamard_transform import hadamard_transform
|
from fast_hadamard_transform import hadamard_transform
|
||||||
elif _is_xpu:
|
elif _is_xpu:
|
||||||
@@ -815,6 +819,11 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
assert page_size == 1, (
|
assert page_size == 1, (
|
||||||
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
|
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:
|
else:
|
||||||
assert page_size == 64, "only support page size 64"
|
assert page_size == 64, "only support page size 64"
|
||||||
# NOTE(dark): this support extend/decode/decode+graph
|
# NOTE(dark): this support extend/decode/decode+graph
|
||||||
@@ -960,6 +969,17 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
preshuffle=_use_aiter_preshuffle,
|
preshuffle=_use_aiter_preshuffle,
|
||||||
kv_block_size=block_kv,
|
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:
|
elif use_cute_dsl:
|
||||||
logits = cutedsl_paged_mqa_logits(
|
logits = cutedsl_paged_mqa_logits(
|
||||||
q_fp8,
|
q_fp8,
|
||||||
@@ -1027,7 +1047,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
if cached_budget is not None:
|
if cached_budget is not None:
|
||||||
return cached_budget
|
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)
|
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
|
||||||
mem_fraction_static = get_schedule().mem_fraction_static
|
mem_fraction_static = get_schedule().mem_fraction_static
|
||||||
@@ -1048,8 +1068,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
return static_budget
|
return static_budget
|
||||||
|
|
||||||
# Match the original free-memory guard: logits_bytes * 2 > free_mem.
|
# Match the original free-memory guard: logits_bytes * 2 > free_mem.
|
||||||
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
|
# Synchronizes the host; cache the result capped by serving-memory headroom.
|
||||||
# capped by the workload-independent 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)
|
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
||||||
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
|
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
|
||||||
|
|
||||||
@@ -1104,6 +1129,12 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
assert page_size == 1, (
|
assert page_size == 1, (
|
||||||
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
|
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:
|
else:
|
||||||
assert page_size == 64, "only support page size 64"
|
assert page_size == 64, "only support page size 64"
|
||||||
|
|
||||||
@@ -1186,6 +1217,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
ke,
|
ke,
|
||||||
clean_logits=False,
|
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:
|
else:
|
||||||
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
||||||
q_fp8[:q_offset], weights[:q_offset]
|
q_fp8[:q_offset], weights[:q_offset]
|
||||||
@@ -1242,6 +1282,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
ke[start:end],
|
ke[start:end],
|
||||||
clean_logits=False,
|
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:
|
else:
|
||||||
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
||||||
q_fp8[start:end], weights[start:end]
|
q_fp8[start:end], weights[start:end]
|
||||||
@@ -1388,6 +1437,12 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
|
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
|
||||||
|
|
||||||
page_size = get_token_to_kv_pool().page_size
|
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 page_size == 64, "only support page size 64"
|
||||||
assert len(weights.shape) == 3
|
assert len(weights.shape) == 3
|
||||||
weights = weights.squeeze(-1)
|
weights = weights.squeeze(-1)
|
||||||
@@ -1871,7 +1926,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
else:
|
else:
|
||||||
weights = self._get_logits_head_gate(x_for_gate, q_scale)
|
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
|
# In piecewise/breakable CUDA graph, any access to seq_lens_cpu
|
||||||
# creates a Dynamo shape guard. These graph modes never have empty
|
# creates a Dynamo shape guard. These graph modes never have empty
|
||||||
# batches.
|
# batches.
|
||||||
|
|||||||
@@ -101,6 +101,7 @@ from sglang.srt.utils import (
|
|||||||
is_cuda,
|
is_cuda,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
|
is_xpu,
|
||||||
print_warning_once,
|
print_warning_once,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -158,6 +159,7 @@ def materialize_full_kv_cp(
|
|||||||
|
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices
|
from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices
|
||||||
@@ -176,6 +178,11 @@ if _is_hip:
|
|||||||
print(
|
print(
|
||||||
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
"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:
|
else:
|
||||||
from sglang.kernels.ops.attention.flash_attention import (
|
from sglang.kernels.ops.attention.flash_attention import (
|
||||||
flash_attn_varlen_func,
|
flash_attn_varlen_func,
|
||||||
@@ -313,6 +320,7 @@ _DSA_IMPL_T: TypeAlias = Literal[
|
|||||||
"fa3",
|
"fa3",
|
||||||
"tilelang",
|
"tilelang",
|
||||||
"trtllm",
|
"trtllm",
|
||||||
|
"intel_xpu",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@@ -452,6 +460,9 @@ class DeepseekSparseAttnBackend(
|
|||||||
"Disabling fused DSA top-k for IndexShare under PD disaggregation."
|
"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_capability = torch.cuda.get_device_capability()
|
||||||
self.device_sm_major = self.device_capability[0]
|
self.device_sm_major = self.device_capability[0]
|
||||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
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_buf: Optional[torch.Tensor] = None
|
||||||
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
|
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
|
||||||
self._q8kv8_born_q_sentinel: Optional[torch.Tensor] = 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 (
|
from sglang.kernels.ops.attention.flash_mla_sm120 import (
|
||||||
_validate_flashinfer_sparse_mla_backend,
|
_validate_flashinfer_sparse_mla_backend,
|
||||||
@@ -1139,7 +1150,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
cache_seqlens=dsa_cache_seqlens_int32,
|
cache_seqlens=dsa_cache_seqlens_int32,
|
||||||
seq_len_q=1,
|
seq_len_q=1,
|
||||||
)
|
)
|
||||||
if use_flashmla_kv
|
if use_flashmla_kv and not _is_xpu
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
||||||
@@ -2309,6 +2320,15 @@ class DeepseekSparseAttnBackend(
|
|||||||
page_table_1=page_table_1,
|
page_table_1=page_table_1,
|
||||||
layer=layer,
|
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:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported {dsa_impl = } for forward_extend. Consider using an other attention backend."
|
f"Unsupported {dsa_impl = } for forward_extend. Consider using an other attention backend."
|
||||||
@@ -2492,6 +2512,15 @@ class DeepseekSparseAttnBackend(
|
|||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
bs=forward_batch.batch_size,
|
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:
|
else:
|
||||||
assert False, f"Unsupported {dsa_impl = }"
|
assert False, f"Unsupported {dsa_impl = }"
|
||||||
@@ -3115,6 +3144,123 @@ class DeepseekSparseAttnBackend(
|
|||||||
d_v=v_head_dim,
|
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(
|
def _forward_aiter(
|
||||||
self,
|
self,
|
||||||
q_all: torch.Tensor,
|
q_all: torch.Tensor,
|
||||||
|
|||||||
@@ -471,9 +471,17 @@ class RotaryEmbedding(BaseFusedOp):
|
|||||||
)
|
)
|
||||||
return query, key
|
return query, key
|
||||||
else:
|
else:
|
||||||
# Use fallback kernel of 'rotary_embedding'
|
|
||||||
self._match_cos_sin_cache_dtype(query)
|
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,
|
positions,
|
||||||
query,
|
query,
|
||||||
key,
|
key,
|
||||||
@@ -481,6 +489,11 @@ class RotaryEmbedding(BaseFusedOp):
|
|||||||
self.cos_sin_cache,
|
self.cos_sin_cache,
|
||||||
self.is_neox_style,
|
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):
|
class LinearScalingRotaryEmbedding(RotaryEmbedding):
|
||||||
|
|||||||
@@ -80,6 +80,7 @@ from sglang.srt.utils import (
|
|||||||
is_float4_e2m1fn_x2,
|
is_float4_e2m1fn_x2,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
|
is_xpu,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.async_probe import (
|
from sglang.srt.utils.async_probe import (
|
||||||
@@ -4722,6 +4723,11 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
|||||||
assert self.page_size == 1, (
|
assert self.page_size == 1, (
|
||||||
f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
|
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:
|
else:
|
||||||
assert self.page_size == 64
|
assert self.page_size == 64
|
||||||
self.index_key_cache = self._create_index_key_cache()
|
self.index_key_cache = self._create_index_key_cache()
|
||||||
|
|||||||
@@ -1843,6 +1843,7 @@ class ServerArgs:
|
|||||||
Arg(
|
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.",
|
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"],
|
choices=["sgl-kernel", "torch", "flashinfer"],
|
||||||
|
resolvable=True,
|
||||||
),
|
),
|
||||||
NS("exec.kernel"),
|
NS("exec.kernel"),
|
||||||
] = "sgl-kernel"
|
] = "sgl-kernel"
|
||||||
|
|||||||
@@ -94,6 +94,7 @@ class TestModelOverridableWhitelist(CustomTestCase):
|
|||||||
"kv_cache_dtype",
|
"kv_cache_dtype",
|
||||||
"dsa_prefill_backend",
|
"dsa_prefill_backend",
|
||||||
"dsa_decode_backend",
|
"dsa_decode_backend",
|
||||||
|
"dsa_topk_backend",
|
||||||
"prefill_attention_backend",
|
"prefill_attention_backend",
|
||||||
"decode_attention_backend",
|
"decode_attention_backend",
|
||||||
"flashinfer_allreduce_fusion_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