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 20d58409ac4e..0628d7eb7152 100644 --- a/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py +++ b/python/sglang/kernels/ops/attention/dsa/index_buf_accessor.py @@ -6,9 +6,10 @@ from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.layers.attention.dsa.utils import 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 @@ -308,6 +309,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/hardware_backend/xpu/kernels/dsa/act_quant.py b/python/sglang/srt/hardware_backend/xpu/kernels/dsa/act_quant.py new file mode 100644 index 000000000000..b3d2becff6f6 --- /dev/null +++ b/python/sglang/srt/hardware_backend/xpu/kernels/dsa/act_quant.py @@ -0,0 +1,49 @@ +"""Torch-native per-group FP8 activation quantization for XPU.""" + +from typing import Optional, Tuple + +import torch + + +def act_quant( + x: torch.Tensor, block_size: int = 128, scale_fmt: Optional[str] = None +) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-group FP8 activation quantization using PyTorch ops. + + For each group of `block_size` columns, computes the abs-max, derives a + per-group scale, and quantizes to float8_e4m3fn. + + Args: + x: Input tensor (contiguous, last dim divisible by block_size). + block_size: Number of columns per quantization group. + scale_fmt: If not None, round scales to nearest power of 2. + + Returns: + (y_fp8, scale) where y_fp8 has dtype float8_e4m3fn and scale is float32. + """ + assert x.is_contiguous(), "Input tensor must be contiguous" + N = x.size(-1) + assert ( + N % block_size == 0 + ), f"Last dim must be divisible by block_size ({block_size})" + + FP8_MAX = 448.0 + x_flat = x.view(-1, N).float() + M = x_flat.size(0) + n_groups = N // block_size + + # Reshape to (M, n_groups, block_size) for per-group quantization + x_grouped = x_flat.view(M, n_groups, block_size) + amax = x_grouped.abs().amax(dim=-1).clamp(min=1e-4) # (M, n_groups) + + if scale_fmt is not None: + # Round scale to power of 2 + scale = torch.exp2(torch.log2(amax / FP8_MAX).ceil()) + else: + scale = amax / FP8_MAX + + # Quantize + y = (x_grouped / scale.unsqueeze(-1)).clamp(-FP8_MAX, FP8_MAX) + y = y.view(M, N).to(torch.float8_e4m3fn).view(*x.shape[:-1], N) + scale = scale.view(*x.shape[:-1], n_groups) + return y, scale diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 196639f18146..f8fe7247f3ec 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -98,6 +98,10 @@ 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 @@ -337,7 +341,6 @@ def topk_transform( 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: @@ -894,6 +897,11 @@ def _get_topk_paged( 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 @@ -985,6 +993,17 @@ def _get_topk_paged( 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, @@ -1052,7 +1071,10 @@ def _get_mqa_logits_budget_bytes(self, device_index: int) -> int: if cached_budget is not None: return cached_budget - total_mem = torch.cuda.get_device_properties(device_index).total_memory + if _is_xpu: + total_mem = torch.xpu.get_device_properties(device_index).total_memory + else: + total_mem = torch.cuda.get_device_properties(device_index).total_memory total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION) mem_fraction_static = get_server_args().mem_fraction_static @@ -1073,10 +1095,15 @@ def _get_mqa_logits_budget_bytes(self, device_index: int) -> int: 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 @@ -1125,13 +1152,19 @@ def _get_topk_ragged( page_size == 1 ), 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 ( - forward_batch.seq_lens_cpu is not None - and forward_batch.extend_seq_lens_cpu is not None - ) + assert len(weights.shape) == 3 + assert ( + forward_batch.seq_lens_cpu is not None + and forward_batch.extend_seq_lens_cpu is not None + ) weights = weights.squeeze(-1) if _is_hip and not _use_aiter_preshuffle: @@ -1206,6 +1239,15 @@ def _get_topk_ragged( 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] @@ -1262,6 +1304,15 @@ def _get_topk_ragged( 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] @@ -1406,7 +1457,13 @@ def _get_topk_ragged_with_cp( 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 = [] @@ -1558,7 +1615,13 @@ def forward_indexer( from sglang.kernels.ops.attention.dsa.tilelang_kernel import fp8_index 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) @@ -1735,6 +1798,10 @@ def forward_cuda( ) -> Optional[torch.Tensor]: if _is_hip: from sglang.kernels.ops.attention.dsa.tilelang_kernel import act_quant + elif _is_xpu: + from sglang.srt.hardware_backend.xpu.kernels.dsa.act_quant import ( + act_quant, + ) elif not _is_npu: from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant @@ -1968,7 +2035,7 @@ def forward_cuda( 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 2606b608157f..b07d2de01213 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -68,6 +68,7 @@ is_gfx95_supported, is_hip, is_sm100_supported, + is_xpu, print_warning_once, ) @@ -106,6 +107,7 @@ def _all_gather_dsa_trtllm_fp8_kv( _is_hip = is_hip() +_is_xpu = is_xpu() if _is_hip: from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices @@ -327,7 +329,13 @@ def topk_transform( _DSA_IMPL_T: TypeAlias = Literal[ - "flashmla_sparse", "flashmla_sparse_q8", "flashmla_kv", "fa3", "tilelang", "trtllm" + "flashmla_sparse", + "flashmla_sparse_q8", + "flashmla_kv", + "fa3", + "tilelang", + "trtllm", + "intel_xpu", ] @@ -448,7 +456,10 @@ def __init__( "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 @@ -1030,7 +1041,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): 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, @@ -2115,6 +2126,15 @@ def forward_extend( 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." @@ -2274,6 +2294,15 @@ def forward_decode( 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 {self.dsa_decode_impl = }" @@ -2603,7 +2632,24 @@ def _forward_standard_mha( skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), ) - # Use FA3 for SM90 (Hopper/H200) + # Use FA3 for SM90 (Hopper/H200) / XPU + if _is_xpu: + from sgl_kernel.flash_attn import ( + flash_attn_varlen_func as xpu_flash_attn_varlen_func, + ) + + return xpu_flash_attn_varlen_func( + q=q, + k=k, + v=v, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=metadata.max_seq_len_q, + max_seqlen_k=max_seqlen_k, + softmax_scale=layer.scaling, + causal=causal, + ) + return flash_attn_varlen_func( q=q, k=k, @@ -2634,6 +2680,123 @@ def _forward_tilelang( 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, @@ -3001,10 +3164,11 @@ def set_dsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): device_sm = get_device_sm() # Requirements: H200/B200, short sequences, supported dtype, fits in chunk + # XPU uses flash_mla_prefill for all prefill lengths; MHA_ONE_SHOT not needed. self.use_mha = ( ( device_sm == 90 or (device_sm >= 100 and device_sm < 110) - ) # SM90/SM100 only + ) # SM90/SM100 only (not XPU) and max_kv_len <= envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get() # Short enough for MHA and self.token_to_kv_pool.dtype in [torch.bfloat16, torch.float8_e4m3fn] diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index 8deac2d77fd6..c64be8c79ee2 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -70,6 +70,14 @@ def init_forward_metadata_out_graph( ): backend = self._select_backend(forward_batch.forward_mode) backend.init_forward_metadata_out_graph(forward_batch, in_capture=in_capture) + # If the selected backend is not the decode backend, also initialize + # the decode backend so that get_indexer_metadata works for prefill + # (e.g., Triton prefill + DSA decode hybrid, where only the DSA decode + # backend manages the DSA index K-cache). + if backend is not self.decode_backend: + self.decode_backend.init_forward_metadata_out_graph( + forward_batch, in_capture=in_capture + ) def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): backend = self._select_backend(forward_batch.forward_mode) @@ -78,6 +86,8 @@ def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch): backend = self._select_backend(forward_batch.forward_mode) backend.init_forward_metadata(forward_batch) + if backend is not self.decode_backend: + self.decode_backend.init_forward_metadata(forward_batch) def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens) @@ -152,6 +162,13 @@ def forward_extend( def get_indexer_metadata( self, layer_id: int, forward_batch: ForwardBatch ) -> Optional[BaseIndexerMetadata]: + # The DSA indexer (K-cache storage + logit computation) always uses the + # decode backend because the decode backend (DeepseekSparseAttnBackend) + # maintains the DSA index K-cache. The prefill attention backend (e.g. + # Triton) handles the main attention but does not manage the DSA index. + meta = self.decode_backend.get_indexer_metadata(layer_id, forward_batch) + if meta is not None: + return meta backend = self._select_backend(forward_batch.forward_mode) return backend.get_indexer_metadata(layer_id, forward_batch) diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 2f68e4965111..c20ab28c025c 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -470,8 +470,16 @@ def forward_xpu( ) return query, key else: - # Use fallback kernel of 'rotary_embedding' - 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.unsqueeze(1) + if k_2d: + key = key.unsqueeze(1) + q_out, k_out = torch.ops.sgl_kernel.rotary_embedding( positions, query, key, @@ -479,6 +487,11 @@ def forward_xpu( self.cos_sin_cache, self.is_neox_style, ) + if q_2d: + q_out = q_out.squeeze(1) + if k_2d: + k_out = k_out.squeeze(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 0599ec7a1a06..4a76da02c735 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -79,6 +79,7 @@ is_float4_e2m1fn_x2, is_hip, is_npu, + is_xpu, next_power_of_2, ) from sglang.srt.utils.async_probe import maybe_detect_oob @@ -3703,6 +3704,11 @@ def __init__( 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._create_index_buffers() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index f2cedff017b6..95e2bdd2d000 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3993,6 +3993,11 @@ def _handle_gpu_memory_settings(self, gpu_mem): reserved_mem = max(reserved_mem, 10 * 1024) # Reserve headroom for DeepEP all-to-all buffers on top of the floor. reserved_mem += self.reserve_for_deepep_a2a_mb() + # 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 is_xpu(): + reserved_mem += 2 * 1024 self.mem_fraction_static = ( round((gpu_mem - reserved_mem) / gpu_mem, 3) @@ -4342,6 +4347,42 @@ def _handle_model_specific_adjustments(self): ) self._set_default_dsa_backends(major) + elif is_xpu(): + self.page_size = 128 + logger.warning("Setting page size to 128 for DeepSeek DSA on XPU.") + if self.kv_cache_dtype == "auto": + self.kv_cache_dtype = "bfloat16" + logger.warning( + f"Setting KV cache dtype to {self.kv_cache_dtype} for DeepSeek DSA on XPU." + ) + if self.dsa_prefill_backend is None: + self.dsa_prefill_backend = "intel_xpu" + if self.dsa_decode_backend is None: + self.dsa_decode_backend = "intel_xpu" + self.dsa_topk_backend = "torch" + # Disable fused topk (requires sgl-kernel ops not available on XPU) + import os + + if os.environ.get("SGLANG_DSA_FUSE_TOPK", "0") != "0": + 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) + logger.warning( + f"Set DSA backends for XPU: prefill={self.dsa_prefill_backend}, decode={self.dsa_decode_backend}." + ) + # Use dsa for both prefill and decode so DeepseekSparseAttnBackend + # handles the full forward pass (KV cache store, DSA indexer, and + # MLA attention) without needing a HybridAttnBackend. + if self.decode_attention_backend is None: + self.decode_attention_backend = "dsa" + # Prefill now uses flash_mla_prefill via the intel_xpu dsa path; + # no longer needs a separate Triton backend. + if self.prefill_attention_backend is None: + self.prefill_attention_backend = "dsa" + if self.enable_prefill_cp: assert ( self.disaggregation_mode != "decode" 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 000000000000..d54f96cb172a --- /dev/null +++ b/test/registered/xpu/test_dsa_indexer_xpu.py @@ -0,0 +1,714 @@ +"""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 MagicMock, 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 + + 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", + }, + )() + 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.srt.hardware_backend.xpu.kernels.dsa.act_quant.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.srt.hardware_backend.xpu.kernels.dsa.act_quant.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 + ), + ) + + # ------------------------------------------------------------------ + # Test: HybridAttnBackend routes indexer metadata to decode backend + # ------------------------------------------------------------------ + + def test_hybrid_backend_indexer_metadata_routing(self): + """HybridAttnBackend.get_indexer_metadata delegates to the decode backend. + + Tests the generic HybridAttnBackend routing: when a hybrid setup uses + DSA for decode and another backend (e.g. Triton) for prefill, the DSA + decode backend must return indexer metadata even for prefill batches. + Note: on XPU, GLM5.1 no longer uses HybridAttnBackend (both prefill + and decode use DeepseekSparseAttnBackend directly), but this routing + logic remains valid for other hybrid configurations. + """ + from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend + from sglang.srt.layers.attention.triton_backend import TritonAttnBackend + + self._init_model_runner() + + # Build a minimal HybridAttnBackend: decode=DSA, prefill=Triton mock + mock_triton = MagicMock(spec=TritonAttnBackend) + mock_triton.get_indexer_metadata.return_value = None # Triton has no indexer + + hybrid = HybridAttnBackend.__new__(HybridAttnBackend) + hybrid.decode_backend = self.backend + hybrid.prefill_backend = mock_triton + + # DSA backend should provide metadata for a decode batch + forward_batch = MagicMock() + forward_batch.forward_mode = ForwardMode.DECODE + forward_batch.batch_size = self.batch_size + forward_batch.seq_lens = torch.tensor( + [self.seq_len + 1] * self.batch_size, device=self.device + ) + + # Patch DSA backend's own get_indexer_metadata to return a mock result + sentinel = object() + with patch.object(self.backend, "get_indexer_metadata", return_value=sentinel): + result = hybrid.get_indexer_metadata( + layer_id=0, forward_batch=forward_batch + ) + + self.assertIs(result, sentinel) + # Triton backend's get_indexer_metadata should NOT have been called + mock_triton.get_indexer_metadata.assert_not_called() + + # ------------------------------------------------------------------ + # Test: init_forward_metadata calls both backends for prefill + # ------------------------------------------------------------------ + + def test_hybrid_backend_prefill_initializes_decode_backend(self): + """HybridAttnBackend.init_forward_metadata calls decode_backend for prefill. + + Tests the generic HybridAttnBackend routing: in a hybrid setup, the + DSA decode backend must be initialized during prefill so it can manage + the K-cache. Note: on XPU, GLM5.1 no longer uses HybridAttnBackend + (both prefill and decode go through DeepseekSparseAttnBackend directly), + but this routing logic remains valid for other hybrid configurations. + """ + from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend + from sglang.srt.layers.attention.triton_backend import TritonAttnBackend + + self._init_model_runner() + + mock_triton = MagicMock(spec=TritonAttnBackend) + mock_dsa = MagicMock(spec=DeepseekSparseAttnBackend) + + hybrid = HybridAttnBackend.__new__(HybridAttnBackend) + hybrid.decode_backend = mock_dsa + hybrid.prefill_backend = mock_triton + hybrid._select_backend = lambda mode: mock_triton # Always pick triton + + forward_batch = MagicMock() + forward_batch.forward_mode = ForwardMode.EXTEND + + hybrid.init_forward_metadata(forward_batch) + + # Both backends must be initialized + mock_triton.init_forward_metadata.assert_called_once_with(forward_batch) + mock_dsa.init_forward_metadata.assert_called_once_with(forward_batch) + + +if __name__ == "__main__": + unittest.main()