From 5ebb16005d2804bc3f87c689e3333a3734d80bf7 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sun, 13 Sep 2026 15:43:19 -0700 Subject: [PATCH] [DeepSeek-V4.1] Bump FlashMLA to the fork's rebase head (v4.1 kernels) (#39171) Co-authored-by: Chunan Zeng --- .../sglang/kernels/aot/cmake/flashmla.cmake | 100 ++++++++++-------- .../kernels/aot/csrc/flashmla_extension.cc | 34 +++++- .../test/kits/basic_decode_correctness_kit.py | 5 +- 3 files changed, 91 insertions(+), 48 deletions(-) diff --git a/python/sglang/kernels/aot/cmake/flashmla.cmake b/python/sglang/kernels/aot/cmake/flashmla.cmake index 5b4decca3..ae506d7a2 100644 --- a/python/sglang/kernels/aot/cmake/flashmla.cmake +++ b/python/sglang/kernels/aot/cmake/flashmla.cmake @@ -1,9 +1,9 @@ # flash_mla -# sm90 dense decode HEAD_DIM_K=512 support (sgl-project/FlashMLA#9, merged). +# DeepSeek v4.1 kernels merged into the SGLang fork (sgl-project/FlashMLA@3e18517). FetchContent_Declare( repo-flashmla - URL https://${GITHUB_ARTIFACTORY}/sgl-project/FlashMLA/archive/c1dee569a494b184811a08171a690ece21420262.tar.gz - URL_HASH SHA256=77d3f1714b5903dc8f7a99fbc3a5d9a2b886e449f994b6c4b1d2feeacb467b1e + URL https://${GITHUB_ARTIFACTORY}/sgl-project/FlashMLA/archive/3e18517fb055a6c9608eef5a1f1347fb1a047bbd.tar.gz + URL_HASH SHA256=ab2af4657683a1bbaa707a2781a5e2c2792eba574ff36905f1832598e67fea79 ) FetchContent_Populate(repo-flashmla) @@ -42,23 +42,16 @@ if(${CUDA_VERSION} VERSION_GREATER 12.8) set(FLASHMLA_ENABLE_SM100 ON) endif() if(${CUDA_VERSION} VERSION_GREATER_EQUAL "13.0") - # Patch FlashMLA sources for SM103a support. - # These patches are only needed (and only valid) with CUDA 13+. - - # Patch utils.h: widen IS_SM100 to cover the full SM100 family. - # Newer FlashMLA versions use csrc/utils.h. - set(FLASHMLA_UTILS_FILE "${repo-flashmla_SOURCE_DIR}/csrc/utils.h") - file(READ "${FLASHMLA_UTILS_FILE}" FLASHMLA_UTILS_CONTENT) - string(REPLACE - "#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 1000) -#define IS_SM100 1" - "#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) && (__CUDA_ARCH__ < 1100) -#define IS_SM100 1" - FLASHMLA_UTILS_CONTENT "${FLASHMLA_UTILS_CONTENT}") - file(WRITE "${FLASHMLA_UTILS_FILE}" "${FLASHMLA_UTILS_CONTENT}") - message(STATUS "Patched utils.h for SM103a support") + # B300 is sm_103; sm_100f is family-compatible so it runs there, but ptxas emits + # better SASS for the native arch. FlashMLA's setup.py picks sm_100a + sm_103a + # over sm_100f for the same reason. The cutlass patch below is what lets a TU + # compile at __CUDA_ARCH__ == 1030 against the pinned cutlass. + list(APPEND FLASHMLA_CUDA_FLAGS + "-gencode=arch=compute_103a,code=sm_103a" + ) # Patch cutlass/arch/config.h: add SM103 architecture defines. + # This patch is only needed (and only valid) with CUDA 13+. # The new block is inserted right before the existing "// SM101 and SM101a" # anchor in the upstream header. set(CUTLASS_CONFIG_FILE "${repo-flashmla_SOURCE_DIR}/csrc/cutlass/include/cutlass/arch/config.h") @@ -99,26 +92,33 @@ set(FlashMLA_SOURCES # Compatibility shim for sgl-kernel torch.ops API. ${repo-flashmla_SOURCE_DIR}/csrc/python_api.cpp + # Attention entry points (the pybind registrations in these files are + # compiled out by FLASH_MLA_LIBTORCH_ONLY). + ${repo-flashmla_SOURCE_DIR}/csrc/api/dense_decode.cpp + ${repo-flashmla_SOURCE_DIR}/csrc/api/sparse_decode.cpp + ${repo-flashmla_SOURCE_DIR}/csrc/api/sparse_prefill.cpp + # Decode metadata/combine kernels. - ${repo-flashmla_SOURCE_DIR}/csrc/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu - ${repo-flashmla_SOURCE_DIR}/csrc/smxx/decode/combine/combine.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/smxx/decode/get_decoding_sched_meta/get_decoding_sched_meta.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/smxx/decode/combine/combine.cu # sm90 dense decode. - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/dense/instantiations/fp16.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/dense/instantiations/bf16.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/dense/instantiations/fp16.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/dense/instantiations/bf16.cu # sm90 sparse decode. - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/sparse_fp8/instantiations/model1_persistent_h64.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/sparse_fp8/instantiations/model1_persistent_h128.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/sparse_fp8/instantiations/v32_persistent_h64.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/decode/sparse_fp8/instantiations/v32_persistent_h128.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v4_persistent_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v4_persistent_h128.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h128.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_no_rope_persistent_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_no_rope_persistent_h128.cu # sm90 sparse prefill. - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/prefill/sparse/fwd.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/prefill/sparse/instantiations/phase1_k512.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/prefill/sparse/instantiations/phase1_k512_topklen.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/prefill/sparse/instantiations/phase1_k576.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm90/prefill/sparse/instantiations/phase1_k576_topklen.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k512_topklen.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k576.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k576_topklen.cu ${repo-flashmla_SOURCE_DIR}/csrc/extension/sm90/dense_fp8/dense_fp8_python_api.cpp ${repo-flashmla_SOURCE_DIR}/csrc/extension/sm90/dense_fp8/flash_fwd_mla_fp8_sm90.cu @@ -128,20 +128,33 @@ set(FlashMLA_SOURCES if(FLASHMLA_ENABLE_SM100) list(APPEND FlashMLA_SOURCES # sm100 dense prefill/bwd. - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/dense/fmha_cutlass_fwd_sm100.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/dense/fmha_cutlass_bwd_sm100.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/dense/fmha_cutlass_fwd_sm100.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/dense/fmha_cutlass_bwd_sm100.cu # sm100 sparse prefill. - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd/head64/instantiations/phase1_k512.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd/head64/instantiations/phase1_k576.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd/head128/instantiations/phase1_k512.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd/head128/instantiations/phase1_k576.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_prefill_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd/head64/instantiations/phase1_h64_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd/head64/instantiations/phase1_h64_k576.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd/head128/instantiations/phase1_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd/head128/instantiations/phase1_k576.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_k512.cu # sm100 sparse decode. - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/decode/head64/instantiations/v32.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/decode/head64/instantiations/model1.cu - ${repo-flashmla_SOURCE_DIR}/csrc/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v32_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v32_h64_no_split.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v32_no_rope_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v32_no_rope_h64_no_split.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v4_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v4_h64_no_split.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v41_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v41_h64_no_split.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v41fp4_h64.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/decode/sparse/head64/instantiations/v41fp4_h64_no_split.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512_splitkv.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512_v41.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512_v41_splitkv.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512_v41fp4.cu + ${repo-flashmla_SOURCE_DIR}/csrc/kernels/sm100/prefill/sparse/fwd_for_small_topk/head128/instantiations/phase1_decode_k512_v41fp4_splitkv.cu ) endif() @@ -167,7 +180,6 @@ endif() target_include_directories(flashmla_ops PRIVATE ${repo-flashmla_SOURCE_DIR}/csrc ${repo-flashmla_SOURCE_DIR}/csrc/kerutils/include - ${repo-flashmla_SOURCE_DIR}/csrc/sm90 ${repo-flashmla_SOURCE_DIR}/csrc/extension/sm90/dense_fp8/ ${repo-flashmla_SOURCE_DIR}/csrc/cutlass/include ${repo-flashmla_SOURCE_DIR}/csrc/cutlass/tools/util/include @@ -178,4 +190,6 @@ target_link_libraries(flashmla_ops PRIVATE ${TORCH_LIBRARIES} c10 cuda) install(TARGETS flashmla_ops LIBRARY DESTINATION "sgl_kernel") -target_compile_definitions(flashmla_ops PRIVATE) +# FlashMLA's csrc/api/*.cpp register their ops through pybind by default; sgl-kernel +# links them as plain libtorch C++ and registers the ops itself in flashmla_extension.cc. +target_compile_definitions(flashmla_ops PRIVATE FLASH_MLA_LIBTORCH_ONLY) diff --git a/python/sglang/kernels/aot/csrc/flashmla_extension.cc b/python/sglang/kernels/aot/csrc/flashmla_extension.cc index b9f2fe003..c4f155036 100644 --- a/python/sglang/kernels/aot/csrc/flashmla_extension.cc +++ b/python/sglang/kernels/aot/csrc/flashmla_extension.cc @@ -16,11 +16,36 @@ limitations under the License. #include #include -#include "api/dense_decode.h" -#include "api/sparse_decode.h" -#include "api/sparse_fwd.h" #include "sgl_kernel_ops.h" +// FlashMLA exposes these two entry points only from csrc/api/*.cpp (its own pybind layer lives +// behind FLASH_MLA_LIBTORCH_ONLY), so declare them here the way csrc/python_api.cpp does. +std::tuple, std::optional> dense_attn_decode_interface( + at::Tensor& q, + const at::Tensor& kcache, + const int head_size_v, + const at::Tensor& seqlens_k, + const at::Tensor& block_table, + const float softmax_scale, + bool is_causal, + std::optional& tile_scheduler_metadata, + std::optional& num_splits); + +std::tuple, std::optional> sparse_attn_decode_interface( + const at::Tensor& q, + const at::Tensor& kv, + const at::Tensor& indices, + const std::optional& topk_length, + const std::optional& attn_sink, + std::optional& tile_scheduler_metadata, + std::optional& num_splits, + const std::optional& extra_kv, + const std::optional& extra_indices, + const std::optional& extra_topk_length, + int d_v, + float sm_scale, + const std::optional& kv_format); + static std::tuple, std::optional> sgl_sparse_decode_fwd( const at::Tensor& q, const at::Tensor& kv, @@ -46,7 +71,8 @@ static std::tuple, std::option extra_indices, extra_topk_length, static_cast(d_v), - static_cast(sm_scale)); + static_cast(sm_scale), + std::nullopt); } static std::tuple, std::optional> sgl_dense_decode_fwd( diff --git a/python/sglang/test/kits/basic_decode_correctness_kit.py b/python/sglang/test/kits/basic_decode_correctness_kit.py index 3255d6a17..29b94893f 100644 --- a/python/sglang/test/kits/basic_decode_correctness_kit.py +++ b/python/sglang/test/kits/basic_decode_correctness_kit.py @@ -58,8 +58,11 @@ class BasicDecodeCorrectnessMixin: # Language-agnostic gibberish detector. Healthy English output is # >90% printable ASCII; multilingual token salad / Unicode noise # from broken weight load drops well below 50%. + # Q/A framing, as in the probes above: a base model continues a bare + # instruction with arbitrary text whose language is not pinned down, + # so the ratio would measure the continuation, not output health. out = self._decode_generate( - "Write a single sentence about a sunny day in the park.", + "Q: Write a single sentence about a sunny day in the park.\nA:", self.sanity_max_new_tokens_long, ) printable = sum(1 for c in out if 32 <= ord(c) < 127 or c in "\n\t")