From e06058ed624fce8bb1fdb57dc37be32e9498e8f7 Mon Sep 17 00:00:00 2001 From: Chunan Zeng Date: Wed, 27 May 2026 14:32:44 -0700 Subject: [PATCH] [Kernel] Import flash_mla kernels from sglang kernel for deepseek v4 (#26499) --- python/sglang/srt/layers/attention/deepseek_v4_backend.py | 6 +++--- .../srt/layers/attention/deepseek_v4_backend_hip_radix.py | 4 ++-- python/sglang/srt/layers/attention/hip_flash_mla.py | 2 +- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 883e38d28..e11c2d3a6 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -58,7 +58,7 @@ from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import ceil_align if TYPE_CHECKING: - from flash_mla.flash_mla_interface import FlashMLASchedMeta + from sgl_kernel.flash_mla import FlashMLASchedMeta from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner @@ -82,7 +82,7 @@ def _pad_last_dim(x: T, multiples_of: int = PAGE_INDEX_ALIGNED_SIZE) -> T: def _create_flashmla_metadata(): - import flash_mla + import sgl_kernel.flash_mla as flash_mla return flash_mla.get_mla_metadata()[0] @@ -1045,7 +1045,7 @@ class DeepseekV4AttnBackend( extra_indices.shape[-1] % 64 == 0 ), f"{extra_indices.shape=}'s last dimension is not aligned to 64" - import flash_mla + import sgl_kernel.flash_mla as flash_mla o = flash_mla.flash_mla_with_kvcache( q=q, diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 3e0ee41ab..ec59114d9 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -55,7 +55,7 @@ from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import ceil_align if TYPE_CHECKING: - from flash_mla.flash_mla_interface import FlashMLASchedMeta + from sgl_kernel.flash_mla import FlashMLASchedMeta from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.model_runner import ModelRunner @@ -83,7 +83,7 @@ def _create_flashmla_metadata(): if is_hip(): return None - import flash_mla + import sgl_kernel.flash_mla as flash_mla return flash_mla.get_mla_metadata()[0] diff --git a/python/sglang/srt/layers/attention/hip_flash_mla.py b/python/sglang/srt/layers/attention/hip_flash_mla.py index ae6da641f..c26705c0c 100644 --- a/python/sglang/srt/layers/attention/hip_flash_mla.py +++ b/python/sglang/srt/layers/attention/hip_flash_mla.py @@ -14,7 +14,7 @@ def flash_mla_with_kvcache_entrypoint(backend: str, **kwargs): backend = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "tilelang") else: - import flash_mla + import sgl_kernel.flash_mla as flash_mla if backend == "comparison": pack_ref, pack_fast_via_tester = flash_mla_with_kvcache_entrypoint(