XPU: Enable GLM5.1 (GlmMoeDsaForCausalLM) DSA Attention (#24959)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
co-authored by
Copilot
Ma Mingfei
parent
39a80354aa
commit
c4e52a1051
@@ -9,9 +9,10 @@ from sglang.srt.layers.attention.dsa.utils import (
|
||||
INDEXER_K_CACHE_PRESHUFFLE_TILE,
|
||||
aiter_can_use_preshuffle_paged_mqa,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip, is_xpu
|
||||
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
_is_fp8_fnuz = is_fp8_fnuz()
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
# aiter cp_gather kernel with preshuffle=True is only valid when the indexer
|
||||
@@ -311,6 +312,11 @@ def _set_k_and_s_triton(
|
||||
assert page_size % 16 == 0, (
|
||||
f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}"
|
||||
)
|
||||
elif _is_xpu:
|
||||
assert page_size in (
|
||||
64,
|
||||
128,
|
||||
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
|
||||
else:
|
||||
assert page_size == 64
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ from sglang.srt.arg_groups.overrides import (
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
|
||||
from sglang.srt.runtime_context import get_platform
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -279,6 +280,11 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
|
||||
reserved_mem = max(reserved_mem, 10 * 1024)
|
||||
# Reserve headroom for DeepEP all-to-all buffers on top of the floor.
|
||||
reserved_mem += reserve_for_deepep_a2a_mb(server_args)
|
||||
# XPU: oneDNN allocates scratch space for matmul when
|
||||
# M is not a power-of-2-aligned value (e.g. M=2100). Reserve extra
|
||||
# headroom so non-aligned prefill lengths don't hit OOM.
|
||||
if get_platform().is_xpu:
|
||||
reserved_mem += 2 * 1024
|
||||
|
||||
mem_fraction_static = (
|
||||
round((gpu_mem - reserved_mem) / gpu_mem, 3)
|
||||
|
||||
@@ -244,6 +244,21 @@ def handle_model_specific_adjustments(server_args: Any):
|
||||
run_post_process_pass(server_args, _dsa_kv_cache_dtype_default)
|
||||
run_post_process_pass(server_args, _dsa_split_backend_resolution)
|
||||
|
||||
elif get_platform().is_xpu:
|
||||
run_post_process_pass(server_args, _dsa_kv_cache_dtype_default)
|
||||
run_post_process_pass(server_args, _dsa_split_backend_resolution)
|
||||
# Disable fused topk (requires sgl-kernel ops not available on XPU)
|
||||
if (
|
||||
envs.SGLANG_DSA_FUSE_TOPK.is_set()
|
||||
and envs.SGLANG_DSA_FUSE_TOPK.get()
|
||||
):
|
||||
logger.warning(
|
||||
"Disabling fused topk for DeepSeek DSA on XPU (SGLANG_DSA_FUSE_TOPK=0). Not supported yet."
|
||||
)
|
||||
envs.SGLANG_DSA_FUSE_TOPK.set(False)
|
||||
# Disable CUDA-JIT topk-v2 (TileLang/TVM-based, requires CUDA)
|
||||
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
|
||||
|
||||
if cfg.enable_prefill_cp:
|
||||
assert cfg.disaggregation_mode != "decode", (
|
||||
"CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp."
|
||||
|
||||
@@ -143,6 +143,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
else:
|
||||
overrides["page_size"] = 64
|
||||
logger.warning("Setting page size to 64 for DeepSeek DSA.")
|
||||
elif get_platform().is_xpu:
|
||||
overrides["page_size"] = 128
|
||||
logger.warning("Setting page size to 128 for DeepSeek DSA on XPU.")
|
||||
else:
|
||||
# DeepSeek V3/R1/V3.1
|
||||
if get_platform().is_sm100:
|
||||
|
||||
@@ -611,7 +611,14 @@ def _dsa_kv_cache_dtype_default(view: Any) -> dict:
|
||||
return {}
|
||||
if not is_deepseek_dsa(hf_config):
|
||||
return {}
|
||||
if get_platform().is_npu or get_platform().is_xpu:
|
||||
if get_platform().is_npu:
|
||||
return {}
|
||||
if get_platform().is_xpu:
|
||||
if view.kv_cache_dtype == "auto":
|
||||
logger.warning(
|
||||
"Setting KV cache dtype to bfloat16 for DeepSeek DSA on XPU."
|
||||
)
|
||||
return {"kv_cache_dtype": "bfloat16"}
|
||||
return {}
|
||||
|
||||
import torch
|
||||
@@ -689,8 +696,26 @@ def _dsa_split_backend_resolution(view: Any) -> dict:
|
||||
return {}
|
||||
if not is_deepseek_dsa(hf_config):
|
||||
return {}
|
||||
if get_platform().is_npu or get_platform().is_xpu:
|
||||
if get_platform().is_npu:
|
||||
return {}
|
||||
if get_platform().is_xpu:
|
||||
declared: Dict[str, Any] = {}
|
||||
if view.dsa_prefill_backend is None:
|
||||
declared["dsa_prefill_backend"] = "intel_xpu"
|
||||
if view.dsa_decode_backend is None:
|
||||
declared["dsa_decode_backend"] = "intel_xpu"
|
||||
# sgl-kernel topk ops (the default) are CUDA-only; fall back to the
|
||||
# torch-native topk implementation on XPU, unless the user already
|
||||
# picked a different backend explicitly (e.g. "flashinfer").
|
||||
if view.dsa_topk_backend == "sgl-kernel":
|
||||
declared["dsa_topk_backend"] = "torch"
|
||||
logger.warning(
|
||||
"Set DSA backends for XPU: prefill=%s, decode=%s, topk=%s.",
|
||||
declared.get("dsa_prefill_backend", view.dsa_prefill_backend),
|
||||
declared.get("dsa_decode_backend", view.dsa_decode_backend),
|
||||
declared.get("dsa_topk_backend", view.dsa_topk_backend),
|
||||
)
|
||||
return declared
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@@ -70,6 +70,7 @@ from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_platform,
|
||||
get_spec,
|
||||
@@ -565,13 +566,10 @@ class DeepseekV4AttnBackend(
|
||||
model_runner.model_config.hf_text_config, "index_topk", C4_TOPK
|
||||
)
|
||||
|
||||
self.enable_deepseek_v4_fp4_indexer: bool = (
|
||||
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
||||
)
|
||||
kernel = get_exec().kernel
|
||||
self.enable_deepseek_v4_fp4_indexer = kernel.enable_deepseek_v4_fp4_indexer
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
|
||||
self.dsv4_prefill_backend: str = getattr(
|
||||
model_runner.server_args, "dsv4_prefill_backend", "auto"
|
||||
)
|
||||
self.dsv4_prefill_backend = getattr(kernel, "dsv4_prefill_backend", "auto")
|
||||
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
|
||||
if not get_platform().is_sm90:
|
||||
raise ValueError(
|
||||
|
||||
@@ -52,6 +52,7 @@ from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
ceil_align,
|
||||
get_bool_env_var,
|
||||
get_device_module,
|
||||
is_cuda,
|
||||
is_gfx95_supported,
|
||||
is_hip,
|
||||
@@ -104,6 +105,10 @@ if _is_cuda:
|
||||
except ImportError as e:
|
||||
deep_gemm = e
|
||||
|
||||
if _is_xpu:
|
||||
from sgl_kernel import fp8_mqa_logits as sgl_fp8_mqa_logits
|
||||
from sgl_kernel import fp8_paged_mqa_logits as sgl_fp8_paged_mqa_logits
|
||||
|
||||
if _use_aiter:
|
||||
from aiter.ops.cache import indexer_k_quant_and_cache
|
||||
|
||||
@@ -187,7 +192,6 @@ def _broadcast_indexer_topk_from_rank0(
|
||||
|
||||
|
||||
def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
||||
# from sgl_kernel import hadamard_transform
|
||||
if _is_hip:
|
||||
from fast_hadamard_transform import hadamard_transform
|
||||
elif _is_xpu:
|
||||
@@ -815,6 +819,11 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
assert page_size == 1, (
|
||||
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
|
||||
)
|
||||
elif _is_xpu:
|
||||
assert page_size in (
|
||||
64,
|
||||
128,
|
||||
), f"XPU DSA only supports page_size 64 or 128, got {page_size}"
|
||||
else:
|
||||
assert page_size == 64, "only support page size 64"
|
||||
# NOTE(dark): this support extend/decode/decode+graph
|
||||
@@ -960,6 +969,17 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
preshuffle=_use_aiter_preshuffle,
|
||||
kv_block_size=block_kv,
|
||||
)
|
||||
elif _is_xpu:
|
||||
logits = sgl_fp8_paged_mqa_logits(
|
||||
q_fp8[:q_offset],
|
||||
kv_cache_fp8,
|
||||
weights[:q_offset],
|
||||
seqlens_32_2d,
|
||||
block_tables,
|
||||
None,
|
||||
max_seq_len,
|
||||
clean_logits=False,
|
||||
)
|
||||
elif use_cute_dsl:
|
||||
logits = cutedsl_paged_mqa_logits(
|
||||
q_fp8,
|
||||
@@ -1027,7 +1047,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
if cached_budget is not None:
|
||||
return cached_budget
|
||||
|
||||
total_mem = torch.cuda.get_device_properties(device_index).total_memory
|
||||
total_mem = get_device_module().get_device_properties(device_index).total_memory
|
||||
|
||||
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
|
||||
mem_fraction_static = get_schedule().mem_fraction_static
|
||||
@@ -1048,10 +1068,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
return static_budget
|
||||
|
||||
# Match the original free-memory guard: logits_bytes * 2 > free_mem.
|
||||
# torch.cuda.mem_get_info synchronizes the host, so cache the result,
|
||||
# capped by the workload-independent serving-memory headroom.
|
||||
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
||||
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
|
||||
# Synchronizes the host; cache the result capped by serving-memory headroom.
|
||||
if _is_xpu:
|
||||
# On XPU, use total_mem budget as the free-memory estimate;
|
||||
# dynamic free-memory query is not supported the same way as CUDA.
|
||||
# TODO Use torch.xpu.mem_get_info() when available (planned end of 2026).
|
||||
budget_bytes = static_budget
|
||||
else:
|
||||
free_mem, _ = torch.cuda.mem_get_info(device_index)
|
||||
budget_bytes = min(int(free_mem * free_mem_fraction), static_budget)
|
||||
|
||||
budget_bytes = max(1, budget_bytes)
|
||||
self._mqa_logits_budget_bytes[device_index] = budget_bytes
|
||||
@@ -1105,7 +1130,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
f"HIP legacy DSA path requires page_size == 1, got {page_size}"
|
||||
)
|
||||
else:
|
||||
assert page_size == 64, "only support page size 64"
|
||||
if _is_xpu:
|
||||
assert page_size in (
|
||||
64,
|
||||
128,
|
||||
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
|
||||
else:
|
||||
assert page_size == 64, "only support page size 64"
|
||||
|
||||
assert len(weights.shape) == 3
|
||||
assert (
|
||||
@@ -1186,6 +1217,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
ke,
|
||||
clean_logits=False,
|
||||
)
|
||||
elif _is_xpu:
|
||||
logits = sgl_fp8_mqa_logits(
|
||||
q_fp8[:q_offset],
|
||||
kv_fp8,
|
||||
weights[:q_offset],
|
||||
ks,
|
||||
ke,
|
||||
clean_logits=False,
|
||||
)
|
||||
else:
|
||||
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
||||
q_fp8[:q_offset], weights[:q_offset]
|
||||
@@ -1242,6 +1282,15 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
ke[start:end],
|
||||
clean_logits=False,
|
||||
)
|
||||
elif _is_xpu:
|
||||
logits_chunk = sgl_fp8_mqa_logits(
|
||||
q_fp8[start:end],
|
||||
kv_fp8,
|
||||
weights[start:end],
|
||||
ks[start:end],
|
||||
ke[start:end],
|
||||
clean_logits=False,
|
||||
)
|
||||
else:
|
||||
q_padded, w_padded, _ = self._pad_heads_for_deep_gemm(
|
||||
q_fp8[start:end], weights[start:end]
|
||||
@@ -1388,7 +1437,13 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
|
||||
|
||||
page_size = get_token_to_kv_pool().page_size
|
||||
assert page_size == 64, "only support page size 64"
|
||||
if _is_xpu:
|
||||
assert page_size in (
|
||||
64,
|
||||
128,
|
||||
), f"XPU DSA requires page_size 64 or 128, got {page_size}"
|
||||
else:
|
||||
assert page_size == 64, "only support page size 64"
|
||||
assert len(weights.shape) == 3
|
||||
weights = weights.squeeze(-1)
|
||||
k_fp8_list = []
|
||||
@@ -1871,7 +1926,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
||||
else:
|
||||
weights = self._get_logits_head_gate(x_for_gate, q_scale)
|
||||
|
||||
if _is_cuda or _is_hip:
|
||||
if _is_cuda or _is_hip or _is_xpu:
|
||||
# In piecewise/breakable CUDA graph, any access to seq_lens_cpu
|
||||
# creates a Dynamo shape guard. These graph modes never have empty
|
||||
# batches.
|
||||
|
||||
@@ -101,6 +101,7 @@ from sglang.srt.utils import (
|
||||
is_cuda,
|
||||
is_gfx95_supported,
|
||||
is_hip,
|
||||
is_xpu,
|
||||
print_warning_once,
|
||||
)
|
||||
|
||||
@@ -158,6 +159,7 @@ def materialize_full_kv_cp(
|
||||
|
||||
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
if _is_hip:
|
||||
from sglang.kernels.ops.attention.dsa.triton_kernel import get_valid_kv_indices
|
||||
@@ -176,6 +178,11 @@ if _is_hip:
|
||||
print(
|
||||
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
||||
)
|
||||
elif _is_xpu:
|
||||
from sgl_kernel.flash_attn import (
|
||||
flash_attn_varlen_func,
|
||||
flash_attn_with_kvcache,
|
||||
)
|
||||
else:
|
||||
from sglang.kernels.ops.attention.flash_attention import (
|
||||
flash_attn_varlen_func,
|
||||
@@ -313,6 +320,7 @@ _DSA_IMPL_T: TypeAlias = Literal[
|
||||
"fa3",
|
||||
"tilelang",
|
||||
"trtllm",
|
||||
"intel_xpu",
|
||||
]
|
||||
|
||||
|
||||
@@ -452,7 +460,10 @@ class DeepseekSparseAttnBackend(
|
||||
"Disabling fused DSA top-k for IndexShare under PD disaggregation."
|
||||
)
|
||||
|
||||
self.device_capability = torch.cuda.get_device_capability()
|
||||
if _is_xpu:
|
||||
self.device_capability = (0, 0)
|
||||
else:
|
||||
self.device_capability = torch.cuda.get_device_capability()
|
||||
self.device_sm_major = self.device_capability[0]
|
||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||
|
||||
@@ -523,7 +534,7 @@ class DeepseekSparseAttnBackend(
|
||||
self._q8kv8_born_q_buf: Optional[torch.Tensor] = None
|
||||
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
|
||||
self._q8kv8_born_q_sentinel: Optional[torch.Tensor] = None
|
||||
self._q8kv8_born_q_tbo = model_runner.server_args.enable_two_batch_overlap
|
||||
self._q8kv8_born_q_tbo = get_exec().overlap.enable_two_batch_overlap
|
||||
|
||||
from sglang.kernels.ops.attention.flash_mla_sm120 import (
|
||||
_validate_flashinfer_sparse_mla_backend,
|
||||
@@ -1139,7 +1150,7 @@ class DeepseekSparseAttnBackend(
|
||||
cache_seqlens=dsa_cache_seqlens_int32,
|
||||
seq_len_q=1,
|
||||
)
|
||||
if use_flashmla_kv
|
||||
if use_flashmla_kv and not _is_xpu
|
||||
else None
|
||||
),
|
||||
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
||||
@@ -2309,6 +2320,15 @@ class DeepseekSparseAttnBackend(
|
||||
page_table_1=page_table_1,
|
||||
layer=layer,
|
||||
)
|
||||
elif dsa_impl == "intel_xpu":
|
||||
return self._forward_intel_xpu_dense_prefill(
|
||||
q_nope=q_nope,
|
||||
q_rope=q_rope,
|
||||
kv_cache=kv_cache,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
sm_scale=layer.scaling,
|
||||
metadata=metadata,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported {dsa_impl = } for forward_extend. Consider using an other attention backend."
|
||||
@@ -2492,6 +2512,15 @@ class DeepseekSparseAttnBackend(
|
||||
metadata=metadata,
|
||||
bs=forward_batch.batch_size,
|
||||
)
|
||||
elif self.dsa_decode_impl == "intel_xpu":
|
||||
return self._forward_intel_xpu_sparse_decode(
|
||||
q_nope=q_nope,
|
||||
q_rope=q_rope,
|
||||
kv_cache=kv_cache,
|
||||
page_table_1=page_table_1,
|
||||
sm_scale=layer.scaling,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
)
|
||||
|
||||
else:
|
||||
assert False, f"Unsupported {dsa_impl = }"
|
||||
@@ -3115,6 +3144,123 @@ class DeepseekSparseAttnBackend(
|
||||
d_v=v_head_dim,
|
||||
)
|
||||
|
||||
def _forward_intel_xpu_sparse_decode(
|
||||
self,
|
||||
q_nope: torch.Tensor,
|
||||
q_rope: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
page_table_1: torch.Tensor,
|
||||
sm_scale: float,
|
||||
v_head_dim: int,
|
||||
) -> torch.Tensor:
|
||||
"""Sparse decode for XPU using flash_mla_decode with gathered KV.
|
||||
|
||||
Gathers the sparse KV tokens selected by the indexer into a contiguous
|
||||
buffer organized as virtual pages, then runs flash_mla_decode on it.
|
||||
"""
|
||||
from sgl_kernel import flash_mla_decode, flash_mla_get_workspace_size
|
||||
|
||||
B = q_nope.shape[0]
|
||||
TOPK = page_table_1.shape[1]
|
||||
D_ckv = kv_cache.shape[-1]
|
||||
GATHER_PAGE_SIZE = 16
|
||||
assert TOPK % GATHER_PAGE_SIZE == 0, (
|
||||
f"TOPK {TOPK} must be a multiple of GATHER_PAGE_SIZE {GATHER_PAGE_SIZE}"
|
||||
)
|
||||
NUM_PAGES = TOPK // GATHER_PAGE_SIZE
|
||||
|
||||
# Count valid tokens per batch (non -1 entries)
|
||||
valid_counts = (page_table_1 >= 0).sum(dim=1).to(torch.int32)
|
||||
|
||||
# Gather KV tokens: replace -1 with 0 for safe indexing, zero-fill after
|
||||
safe_indices = page_table_1.clamp(min=0)
|
||||
# kv_cache is [total_tokens, D_ckv], index with flat indices
|
||||
gathered_kv = kv_cache[safe_indices.view(-1)].view(B, TOPK, D_ckv)
|
||||
# Zero out invalid entries
|
||||
invalid_mask = page_table_1 < 0 # [B, TOPK]
|
||||
gathered_kv[invalid_mask] = 0
|
||||
|
||||
# Reshape to pages: [B * NUM_PAGES, GATHER_PAGE_SIZE, D_ckv]
|
||||
gathered_kv_paged = gathered_kv.view(B * NUM_PAGES, GATHER_PAGE_SIZE, D_ckv)
|
||||
|
||||
# Identity page table: batch i → pages [i*NUM_PAGES, ..., (i+1)*NUM_PAGES-1]
|
||||
identity_page_table = torch.arange(
|
||||
B * NUM_PAGES, device=q_nope.device, dtype=torch.int32
|
||||
).view(B, NUM_PAGES)
|
||||
|
||||
# Workspace
|
||||
ws_size = flash_mla_get_workspace_size(
|
||||
TOPK, B, q_nope.shape[1], GATHER_PAGE_SIZE
|
||||
)
|
||||
if self.workspace_buffer is None:
|
||||
self.workspace_buffer = torch.empty(
|
||||
ws_size, device=q_nope.device, dtype=torch.uint8
|
||||
)
|
||||
elif self.workspace_buffer.numel() < ws_size:
|
||||
self.workspace_buffer.resize_(ws_size)
|
||||
|
||||
o = flash_mla_decode(
|
||||
q_nope,
|
||||
q_rope,
|
||||
gathered_kv_paged,
|
||||
valid_counts,
|
||||
identity_page_table,
|
||||
self.workspace_buffer,
|
||||
sm_scale,
|
||||
)
|
||||
return o
|
||||
|
||||
def _forward_intel_xpu_dense_prefill(
|
||||
self,
|
||||
q_nope: torch.Tensor,
|
||||
q_rope: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
v_head_dim: int,
|
||||
sm_scale: float,
|
||||
metadata: DSAMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Dense prefill for XPU using flash_mla_prefill.
|
||||
|
||||
Runs full causal MLA attention over all KV positions (no sparse top-K
|
||||
selection). This is the XPU equivalent of the CUDA flashmla_kv /
|
||||
flashmla_sparse prefill paths.
|
||||
"""
|
||||
from sgl_kernel import flash_mla_prefill, flash_mla_prefill_get_workspace_size
|
||||
|
||||
D_ckv = kv_cache.shape[-1]
|
||||
# kv_cache: (N_tokens, 1, D_ckv) from MLATokenToKVPool → (N_pages, page_size, D_ckv)
|
||||
kv_paged = kv_cache.view(-1, self.real_page_size, D_ckv)
|
||||
|
||||
block_table = metadata.real_page_table # (B, max_pages), page-indexed int32
|
||||
seq_lens_k = metadata.cache_seqlens_int32 # (B,) total KV lengths
|
||||
cu_seqlens_q = metadata.cu_seqlens_q # (B+1,) cumulative Q lengths
|
||||
max_seqlen_q = metadata.max_seq_len_q
|
||||
|
||||
max_seq_len_k = int(seq_lens_k.max().item())
|
||||
ws_size = flash_mla_prefill_get_workspace_size(
|
||||
max_seq_len_k, seq_lens_k.shape[0]
|
||||
)
|
||||
if self.workspace_buffer is None:
|
||||
self.workspace_buffer = torch.empty(
|
||||
ws_size, device=q_nope.device, dtype=torch.uint8
|
||||
)
|
||||
elif self.workspace_buffer.numel() < ws_size:
|
||||
self.workspace_buffer.resize_(ws_size)
|
||||
|
||||
return flash_mla_prefill(
|
||||
q_nope,
|
||||
q_rope,
|
||||
kv_paged,
|
||||
cu_seqlens_q,
|
||||
seq_lens_k,
|
||||
max_seqlen_q,
|
||||
block_table,
|
||||
self.workspace_buffer,
|
||||
sm_scale,
|
||||
causal=True,
|
||||
num_kv_splits=1,
|
||||
)
|
||||
|
||||
def _forward_aiter(
|
||||
self,
|
||||
q_all: torch.Tensor,
|
||||
|
||||
@@ -471,9 +471,17 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
)
|
||||
return query, key
|
||||
else:
|
||||
# Use fallback kernel of 'rotary_embedding'
|
||||
self._match_cos_sin_cache_dtype(query)
|
||||
return torch.ops.sgl_kernel.rotary_embedding(
|
||||
# Use fallback kernel of 'rotary_embedding'.
|
||||
# The kernel requires 3D tensors (batch, num_heads, head_size);
|
||||
# add a num_heads=1 dim for 2D tensors (e.g. DSA indexer k_rope).
|
||||
q_2d = query.dim() == 2
|
||||
k_2d = key.dim() == 2
|
||||
if q_2d:
|
||||
query = query.view(query.shape[0], -1, self.head_size)
|
||||
if k_2d:
|
||||
key = key.view(key.shape[0], -1, self.head_size)
|
||||
q_out, k_out = torch.ops.sgl_kernel.rotary_embedding(
|
||||
positions,
|
||||
query,
|
||||
key,
|
||||
@@ -481,6 +489,11 @@ class RotaryEmbedding(BaseFusedOp):
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style,
|
||||
)
|
||||
if q_2d:
|
||||
q_out = q_out.view(q_out.shape[0], -1)
|
||||
if k_2d:
|
||||
k_out = k_out.view(k_out.shape[0], -1)
|
||||
return q_out, k_out
|
||||
|
||||
|
||||
class LinearScalingRotaryEmbedding(RotaryEmbedding):
|
||||
|
||||
@@ -80,6 +80,7 @@ from sglang.srt.utils import (
|
||||
is_float4_e2m1fn_x2,
|
||||
is_hip,
|
||||
is_npu,
|
||||
is_xpu,
|
||||
next_power_of_2,
|
||||
)
|
||||
from sglang.srt.utils.async_probe import (
|
||||
@@ -4722,6 +4723,11 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
assert self.page_size == 1, (
|
||||
f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
|
||||
)
|
||||
elif is_xpu():
|
||||
assert self.page_size in (
|
||||
64,
|
||||
128,
|
||||
), f"XPU DSA requires page_size 64 or 128, got {self.page_size}"
|
||||
else:
|
||||
assert self.page_size == 64
|
||||
self.index_key_cache = self._create_index_key_cache()
|
||||
|
||||
@@ -1843,6 +1843,7 @@ class ServerArgs:
|
||||
Arg(
|
||||
help="DSA indexer top-k backend for the target model. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
|
||||
choices=["sgl-kernel", "torch", "flashinfer"],
|
||||
resolvable=True,
|
||||
),
|
||||
NS("exec.kernel"),
|
||||
] = "sgl-kernel"
|
||||
|
||||
Reference in New Issue
Block a user