diff --git a/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py index 6a7165b7d..f0c7a99ec 100644 --- a/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py +++ b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py @@ -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 diff --git a/python/sglang/srt/arg_groups/memory_hook.py b/python/sglang/srt/arg_groups/memory_hook.py index af06d3bd7..437a4d8bb 100644 --- a/python/sglang/srt/arg_groups/memory_hook.py +++ b/python/sglang/srt/arg_groups/memory_hook.py @@ -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) diff --git a/python/sglang/srt/arg_groups/model_hook.py b/python/sglang/srt/arg_groups/model_hook.py index f11033d6a..00b7234b3 100644 --- a/python/sglang/srt/arg_groups/model_hook.py +++ b/python/sglang/srt/arg_groups/model_hook.py @@ -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." diff --git a/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py b/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py index 47e5793b5..237ede965 100644 --- a/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py +++ b/python/sglang/srt/arg_groups/model_overrides/deepseek_v2.py @@ -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: diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 6d4e463aa..070c67ddb 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -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 diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index d7cc363d4..b398df93f 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -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( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index a7e834743..b53130894 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -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. diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index a2c35aeae..bb403461b 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -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, diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index f3161b35c..778ab7e9b 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -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): diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index ebc8cc3d2..55c46a742 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -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() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e056a7b4e..e6c003fe6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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" diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 5387fbb4c..7618dc2dc 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -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", diff --git a/test/registered/xpu/test_dsa_indexer_xpu.py b/test/registered/xpu/test_dsa_indexer_xpu.py new file mode 100644 index 000000000..292545e3b --- /dev/null +++ b/test/registered/xpu/test_dsa_indexer_xpu.py @@ -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()