CoolFace
Modelpublic

replicate/flash-mla

sourceHugging Facemitupdated 22d agoView on Hugging Face
0likes162downloads
flash_fwd_mla_kernel.h604 linesDownload Raw Back to flash_mla
1#pragma once2 3#include <cute/tensor.hpp>4#include <cutlass/cutlass.h>5#include <cutlass/array.h>6#include <cutlass/numeric_types.h>7 8using namespace cute;9 10#include "named_barrier.h"11#include "utils.h"12#include "softmax.h"13#include "static_switch.h"14#include "flash_mla.h"15 16 17template<typename PrecType, int DIM, int DIM2 = DIM>18constexpr auto getSmemLayoutK() {19    constexpr int headSizeBytes = sizeof(PrecType) * DIM;20    constexpr int headSizeBytes2 = sizeof(PrecType) * DIM2;21 22    if constexpr (headSizeBytes % 128 == 0 && headSizeBytes2 % 128 == 0) {23        return GMMA::Layout_K_SW128_Atom<PrecType>{};24    } else if constexpr (headSizeBytes % 64 == 0 && headSizeBytes2 % 64 == 0) {25        return GMMA::Layout_K_SW64_Atom<PrecType>{};26    } else {27        return GMMA::Layout_K_SW32_Atom<PrecType>{};28    }29}30 31template<int kHeadDim_, int kBlockM_, int kBlockN_, int kNWarps_, typename elem_type=cutlass::bfloat16_t, int kHeadDimV_ = 0>32struct Flash_fwd_kernel_traits_mla {33    using Element = elem_type;34    using ElementAccum = float;35    using index_t = int64_t;36 37    static constexpr int kNWarps = kNWarps_;38    static constexpr int kNThreads = kNWarps * 32;39    static constexpr int kNWarpsS = 4;40    static constexpr int kNThreadsS = kNWarpsS * 32;41 42    static constexpr int kBlockM = kBlockM_;43    static constexpr int kBlockN = kBlockN_;44    static constexpr int kHeadDim = kHeadDim_;45    static_assert(kHeadDim % 32 == 0);46    static constexpr int kHeadDimV = kHeadDimV_ != 0 ? kHeadDimV_ : kHeadDim;47    static_assert(kHeadDimV % 32 == 0);48    static_assert(kHeadDimV <= kHeadDim);49    static constexpr int kBlockKSmem = kHeadDim % 64 == 0 ? 64 : 32;50    static constexpr int kSwizzle = kBlockKSmem == 32 ? 2 : 3;51 52    using TiledMma = decltype(make_tiled_mma(53            cute::GMMA::ss_op_selector<Element, Element, ElementAccum, Shape<Int<kBlockM>, Int<kBlockN>, Int<kHeadDim>>,54                    GMMA::Major::K, GMMA::Major::K>(),55            Layout<Shape<Int<kNWarpsS / 4>, _1, _1>>{}));56 57    static constexpr int AtomLayoutNO = kNThreads / kNThreadsS;58    using TiledMmaO = decltype(make_tiled_mma(59            cute::GMMA::rs_op_selector<Element, Element, ElementAccum, Shape<Int<kBlockM>, Int<kHeadDimV / AtomLayoutNO>, Int<kBlockN>>,60                    GMMA::Major::K, GMMA::Major::MN>(),61            Layout<Shape<Int<kNWarpsS / 4>, Int<AtomLayoutNO>, _1>>{}));62 63    using SmemLayoutQ = decltype(tile_to_shape(64            getSmemLayoutK<Element, kHeadDim>(),65            Shape<Int<kBlockM>, Int<kHeadDim>>{}));66 67    using SmemLayoutK = decltype(tile_to_shape(68            getSmemLayoutK<Element, kHeadDim, kHeadDimV>(),69            Shape<Int<kBlockN>, Int<kHeadDim>>{}));70 71    using SmemLayoutV = decltype(tile_to_shape(72            getSmemLayoutK<Element, kHeadDim, kHeadDimV>(),73            Shape<Int<kBlockN>, Int<kHeadDimV>>{}));74    using SmemLayoutVtransposed = decltype(composition(SmemLayoutV{}, make_layout(Shape<Int<kHeadDimV>, Int<kBlockN>>{}, GenRowMajor{})));75 76    using SmemLayoutP = Layout<Shape<Shape<_2, _2>, Int<kNThreadsS>, _1, Int<kBlockN / 8>>>;77    using SmemLayoutRow = Layout<Shape<_2, Int<kNThreadsS>>, Stride<_1, _2>>;78 79    using SmemLayoutAtomO = decltype(composition(80            Swizzle<kSwizzle, 3, 3>{},81            Layout<Shape<Int<8>, Int<kBlockKSmem>>, Stride<Int<kBlockKSmem>, _1>>{}));82    using SmemLayoutO = decltype(tile_to_shape(83            SmemLayoutAtomO{},84            Shape<Int<kBlockM>, Int<kHeadDimV>>{}));85    using SmemCopyAtomO = Copy_Atom<SM90_U32x4_STSM_N, Element>;86    using SmemCopyAtomOaccum = Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementAccum>;87 88    static constexpr int kGmemElemsPerLoad = sizeof(cute::uint128_t) / sizeof(Element);89    static_assert(kHeadDim % kGmemElemsPerLoad == 0, "kHeadDim must be a multiple of kGmemElemsPerLoad");90    static constexpr int kGmemThreadsPerRow = kBlockKSmem / kGmemElemsPerLoad;91    using Gmem_copy_struct = SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>;92    static constexpr int kNThreadsLoad = kNThreads - kNThreadsS;93    static_assert(kNThreadsLoad % kGmemThreadsPerRow == 0, "kNThreads must be a multiple of kGmemThreadsPerRow");94 95    using GmemLayoutAtom = Layout<96            Shape<Int<kNThreadsLoad / kGmemThreadsPerRow>, Int<kGmemThreadsPerRow>>,97            Stride<Int<kGmemThreadsPerRow>, _1>>;98    using GmemTiledCopy = decltype(make_tiled_copy(99            Copy_Atom<Gmem_copy_struct, Element>{},100            GmemLayoutAtom{},101            Layout<Shape<_1, _8>>{}));  // Val layout, 8 vals per read102 103    using GmemLayoutAtomO = Layout<104            Shape<Int<kNThreadsS / kGmemThreadsPerRow>, Int<kGmemThreadsPerRow>>,105            Stride<Int<kGmemThreadsPerRow>, _1>>;106    using GmemTiledCopyO = decltype(make_tiled_copy(107            Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, Element>{},108            GmemLayoutAtomO{},109            Layout<Shape<_1, _8>>{}));  // Val layout, 8 vals per store110 111    static constexpr int kGmemElemsPerLoadAccum = sizeof(cute::uint128_t) / sizeof(ElementAccum);112    static constexpr int kGmemThreadsPerRowAccum = kBlockKSmem / kGmemElemsPerLoadAccum;113    using GmemLayoutAtomOaccum = Layout<114            Shape<Int<kNThreadsS / kGmemThreadsPerRowAccum>, Int<kGmemThreadsPerRowAccum>>,115            Stride<Int<kGmemThreadsPerRowAccum>, _1>>;116    using GmemTiledCopyOaccum = decltype(make_tiled_copy(117            Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementAccum>{},118            GmemLayoutAtomOaccum{},119            Layout<Shape<_1, _4>>{}));  // Val layout, 4 vals per store120};121 122namespace flash {123 124using namespace cute;125 126template<typename Kernel_traits>127struct SharedStorageMLA {128    union {129        struct {130            cute::array_aligned<typename Kernel_traits::Element, cute::cosize_v<typename Kernel_traits::SmemLayoutQ>> smem_q;131            cute::array_aligned<typename Kernel_traits::Element, cute::cosize_v<typename Kernel_traits::SmemLayoutK> * 2> smem_k;  // Double buffer132            cute::array_aligned<typename Kernel_traits::Element, cute::cosize_v<typename Kernel_traits::SmemLayoutP>> smem_p;133            cute::array_aligned<typename Kernel_traits::ElementAccum, cute::cosize_v<typename Kernel_traits::SmemLayoutRow>> smem_scale;134        };135        struct {136            cute::array_aligned<typename Kernel_traits::ElementAccum, cute::cosize_v<typename Kernel_traits::SmemLayoutRow>> smem_max;137            cute::array_aligned<typename Kernel_traits::ElementAccum, cute::cosize_v<typename Kernel_traits::SmemLayoutRow>> smem_sum;138            cute::array_aligned<typename Kernel_traits::ElementAccum, cute::cosize_v<typename Kernel_traits::SmemLayoutO>> smem_o;139        };140    };141};142 143////////////////////////////////////////////////////////////////////////////////////////////////////144 145template<typename Kernel_traits, bool Split, typename SharedStorage, typename AccO, typename Softmax>146__forceinline__ __device__ void store(const Flash_fwd_mla_params &params, const int bidb, const int bidh, const int m_block, const int n_split_idx,147                                      SharedStorage &shared_storage, AccO tOrO, Softmax softmax) {148    constexpr int kBlockM = Kernel_traits::kBlockM;149    constexpr int kHeadDimV = Kernel_traits::kHeadDimV;150    constexpr int kNThreadsS = Kernel_traits::kNThreadsS;151    using Element = typename Kernel_traits::Element;152    using ElementAccum = typename Kernel_traits::ElementAccum;153    using index_t = typename Kernel_traits::index_t;154 155    const int tidx = threadIdx.x;156 157    typename Kernel_traits::TiledMmaO tiled_mma_o;158    auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);159 160    // Epilogue161 162    const int split_offset = __ldg(params.num_splits_ptr + bidb);163 164    Tensor lse = softmax.template normalize_softmax_lse</*Is_dropout=*/false, Split>(tOrO, params.scale_softmax);165 166    using ElementO = std::conditional_t<!Split, Element, ElementAccum>;167    Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast<ElementO *>(shared_storage.smem_o.data())), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N)168    // Partition sO to match the accumulator partitioning169    using SmemTiledCopyO = std::conditional_t<170            !Split,171            typename Kernel_traits::SmemCopyAtomO,172            typename Kernel_traits::SmemCopyAtomOaccum173    >;174    auto smem_tiled_copy_Oaccum = make_tiled_copy_C(SmemTiledCopyO{}, tiled_mma_o);175    auto smem_thr_copy_Oaccum = smem_tiled_copy_Oaccum.get_thread_slice(tidx);176    Tensor rO = flash::convert_type<ElementO>(tOrO);177    Tensor taccOrOaccum = smem_thr_copy_Oaccum.retile_S(rO);        // ((Atom,AtomNum), MMA_M, MMA_N)178    Tensor taccOsOaccum = smem_thr_copy_Oaccum.partition_D(sOaccum);     // ((Atom,AtomNum),PIPE_M,PIPE_N)179 180    __syncthreads();181 182    cute::copy(smem_tiled_copy_Oaccum, taccOrOaccum, taccOsOaccum);183 184    const index_t row_offset_o = bidb * params.o_batch_stride + m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride;185    const index_t row_offset_oaccum = (((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM) * params.d_v;186    const index_t row_offset_lse = (bidb * params.h + bidh) * params.seqlen_q + m_block * kBlockM;187    const index_t row_offset_lseaccum = ((split_offset + n_split_idx) * params.h + bidh) * params.seqlen_q + m_block * kBlockM;188 189    Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(Split ? params.oaccum_ptr : params.o_ptr) + (Split ? row_offset_oaccum : row_offset_o)),190                                 Shape<Int<kBlockM>, Int<kHeadDimV>>{},191                                 make_stride(Split ? kHeadDimV : params.o_row_stride, _1{}));192    Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + (Split ? row_offset_lseaccum : row_offset_lse)),193                                   Shape<Int<kBlockM>>{}, Stride<_1>{});194 195    using GmemTiledCopyO = std::conditional_t<!Split, typename Kernel_traits::GmemTiledCopyO, typename Kernel_traits::GmemTiledCopyOaccum>;196    GmemTiledCopyO gmem_tiled_copy_Oaccum;197    auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);198    Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum);        // ((Atom,AtomNum),ATOM_M,ATOM_N)199    Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum);200 201    __syncthreads();202 203    if (tidx >= kNThreadsS) { return; }204 205    Tensor tOrOaccum = make_tensor<ElementO>(shape(tOgOaccum));206    cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);207 208    Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{});    // (BLK_M,BLK_K) -> (blk_m,blk_k)209    Tensor taccOcO = thr_mma_o.partition_C(caccO);                           // ((MMA=4, X), MMA_M, MMA_K=1)210    Tensor taccOcO_row = taccOcO(make_coord(0, _, 0), _, 0);211    CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row));                     // MMA_M212    if (get<1>(taccOcO_row(0)) == 0) {213#pragma unroll214        for (int mi = 0; mi < size(lse); ++mi) {215            const int row = get<0>(taccOcO_row(mi));216            if (row < params.seqlen_q - m_block * kBlockM) { gLSEaccum(row) = lse(mi); }217        }218    }219 220    // Construct identity layout for sO221    Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum)));    // (BLK_M,BLK_K) -> (blk_m,blk_k)222    // Repeat the partitioning with identity layouts223    Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO);                           // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)224    Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgOaccum)));225    // Clear_OOB_K must be false since we don't want to write zeros to gmem226    flash::copy</*Is_even_MN=*/false, /*Is_even_K=*/true, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(227            gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, params.seqlen_q - m_block * kBlockM228    );229}230 231template<typename Kernel_traits, bool Is_causal, typename SharedStorage>232__forceinline__ __device__ void compute_attn_1rowblock_splitkv_mla(const Flash_fwd_mla_params &params,233                                                                   const int bidb, const int bidh, const int m_block,234                                                                   const int n_split_idx, const int seqlen_k,235                                                                   const int n_block_min, const int n_block_max, const bool NoSplit,236                                                                   SharedStorage &shared_storage) {237    constexpr int kBlockM = Kernel_traits::kBlockM;238    constexpr int kBlockN = Kernel_traits::kBlockN;239    constexpr int kHeadDim = Kernel_traits::kHeadDim;240    constexpr int kHeadDimV = Kernel_traits::kHeadDimV;241    constexpr int kNThreads = Kernel_traits::kNThreads;242    constexpr int kNThreadsS = Kernel_traits::kNThreadsS;243    static_assert(kNThreads == 256 and kNThreadsS == 128);244    using Element = typename Kernel_traits::Element;245    using index_t = typename Kernel_traits::index_t;246 247    const int tidx = threadIdx.x;248    int n_block = n_block_max - 1;249 250    Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.data()), typename Kernel_traits::SmemLayoutQ{});251    Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.data()), typename Kernel_traits::SmemLayoutK{});252    Tensor sV = make_tensor(make_smem_ptr(shared_storage.smem_k.data()), typename Kernel_traits::SmemLayoutV{});253    Tensor sVt = make_tensor(make_smem_ptr(shared_storage.smem_k.data()), typename Kernel_traits::SmemLayoutVtransposed{});254 255    Tensor sP = make_tensor(make_smem_ptr(shared_storage.smem_p.data()), typename Kernel_traits::SmemLayoutP{});256    Tensor tPsP = sP(_, tidx % kNThreadsS, _, _);257    Tensor sScale_o = make_tensor(make_smem_ptr(shared_storage.smem_scale.data()), typename Kernel_traits::SmemLayoutRow{});258    Tensor tScale_osScale_o = sScale_o(_, tidx % kNThreadsS);259    Tensor sRow_max = make_tensor(make_smem_ptr(shared_storage.smem_max.data()), typename Kernel_traits::SmemLayoutRow{});260    Tensor tRow_maxsRow_max = sRow_max(_, tidx % kNThreadsS);261    Tensor sRow_sum = make_tensor(make_smem_ptr(shared_storage.smem_sum.data()), typename Kernel_traits::SmemLayoutRow{});262    Tensor tRow_sumsRow_sum = sRow_sum(_, tidx % kNThreadsS);263 264    typename Kernel_traits::TiledMmaO tiled_mma_o;265    auto thr_mma_o = tiled_mma_o.get_thread_slice(tidx);266    Tensor tOrVt = thr_mma_o.partition_fragment_B(sVt);                // (MMA, MMA_K,MMA_N)267    Tensor tOrO = partition_fragment_C(tiled_mma_o, Shape<Int<kBlockM>, Int<kHeadDimV>>{});  // ((MMA=4, X), MMA_M, MMA_N=1)268    clear(tOrO);269 270    flash::Softmax<2 * size<1>(tOrO)> softmax;271 272    int warp_group_idx = cutlass::canonical_warp_group_idx();273    if (warp_group_idx == 0) {274        typename Kernel_traits::TiledMma tiled_mma;275        auto thr_mma = tiled_mma.get_thread_slice(tidx);276        Tensor tSrQ = thr_mma.partition_fragment_A(sQ);                           // (MMA,MMA_M,MMA_K)277        Tensor tSrK = thr_mma.partition_fragment_B(sK);                           // (MMA,MMA_N,MMA_K)278 279        if (n_block % 2 == 1) {280            // Double buffer for sK281            constexpr int sK_offset = size(sK);282            tSrK.data() = tSrK.data() + sK_offset / 8;283            tOrVt.data() = tOrVt.data() + sK_offset / 8;284        }285 286        // We need masking on S for the very last block when K and V has length not multiple of kBlockN.287        // We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks.288        // We will have at least 1 "masking" iteration.289        // If not even_N, then seqlen_k might end in the middle of a block. In that case we need to290        // mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1.291        constexpr int n_masking_steps = !Is_causal ? 1 : cute::ceil_div(kBlockM, kBlockN) + 1;292#pragma unroll 1293        for (int masking_step = n_masking_steps; n_block >= n_block_min; --masking_step, --n_block) {294            __syncthreads();295 296            Tensor tSrS = partition_fragment_C(tiled_mma, Shape<Int<kBlockM>, Int<kBlockN>>{});  // ((MMA=4, X), MMA_M, MMA_N=1)297            flash::gemm</*zero_init=*/true, /*wg_wait=*/0>(tiled_mma, tSrQ, tSrK, tSrS);298 299            const bool is_masking_step = masking_step > 0;300            const bool is_first_masking_step = masking_step == n_masking_steps;301 302            if (is_masking_step) {303                Tensor cS = make_identity_tensor(Shape<Int<kBlockM>, Int<kBlockN>>{});304                Tensor tScS = thr_mma.partition_C(cS);305#pragma unroll306                for (int i = 0; i < size(tSrS); ++i) {307                    if constexpr (!Is_causal) {  // Just masking based on col308                        if (int(get<1>(tScS(i))) >= int(seqlen_k - n_block * kBlockN)) tSrS(i) = -INFINITY;309                    } else {310                        // Ensure seqlen_k - 1 - (n_block * kBlockN + col) >= (seqlen_q - 1 - (m_block * kBlockM + row)) / ngroups311                        // col <= seqlen_k - 1 - n_block * kBlockN - (seqlen_q - 1 - (m_block * kBlockM + row)) / ngroups312                        int row = int(get<0>(tScS(i)));313                        int col_limit_right = seqlen_k - 1 - n_block * kBlockN - (params.seqlen_q - 1 - (m_block * kBlockM + row)) / params.ngroups;314                        if (int(get<1>(tScS(i))) > col_limit_right) tSrS(i) = -INFINITY;315                    }316                }317            }318 319            // We have key_padding_mask so we'll need to Check_inf320            Tensor scale_o = is_first_masking_step321                             ? softmax.template softmax</*Is_first=*/true,  /*Check_inf=*/Is_causal>(tSrS, params.scale_softmax_log2)322                             : is_masking_step ?323                               softmax.template softmax</*Is_first=*/false, /*Check_inf=*/Is_causal>(tSrS, params.scale_softmax_log2)324                                               : softmax.template softmax</*Is_first=*/false, /*Check_inf=*//*Is_local=*/false>(tSrS, params.scale_softmax_log2);325 326            Tensor rP = flash::convert_type<Element>(tSrS);327            cute::copy(rP, tPsP);328            cute::copy(scale_o, tScale_osScale_o);329 330            cutlass::arch::NamedBarrier::arrive(kNThreads, static_cast<int>(NamedBarriers::SReady));331 332            flash::rescale_o(tOrO, scale_o);333 334            Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));335            flash::gemm</*zero_init=*/false, /*wg_wait=*/0>(tiled_mma_o, tOrP, tOrVt, tOrO);336 337            // Double buffer for sK338            const int sK_offset = n_block % 2 == 0 ? size(sK) : -size(sK);339            tSrK.data() = tSrK.data() + sK_offset / 8;340            tOrVt.data() = tOrVt.data() + sK_offset / 8;341        }342 343        cute::copy(softmax.row_max, tRow_maxsRow_max);344        cute::copy(softmax.row_sum, tRow_sumsRow_sum);345        cutlass::arch::NamedBarrier::arrive(kNThreads, static_cast<int>(NamedBarriers::SoftmaxReady));346    } else {347        const int *block_table = params.block_table + bidb * params.block_table_batch_stride;348        int cur_block_table = __ldg(&block_table[n_block]);349 350        const index_t row_offset_q = bidb * params.q_batch_stride + m_block * kBlockM * params.q_row_stride + bidh * params.q_head_stride;351        Tensor gQ = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.q_ptr) + row_offset_q),352                                Shape<Int<kBlockM>, Int<kHeadDim>>{},353                                make_stride(params.q_row_stride, _1{}));354        typename Kernel_traits::GmemTiledCopy gmem_tiled_copy_Q;355        auto gmem_thr_copy_Q = gmem_tiled_copy_Q.get_thread_slice(tidx - kNThreadsS);356        Tensor tQgQ = gmem_thr_copy_Q.partition_S(gQ);357        Tensor tQsQ = gmem_thr_copy_Q.partition_D(sQ);358        Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ)));  // (BLK_M,BLK_K) -> (blk_m,blk_k)359        Tensor tQcQ = gmem_thr_copy_Q.partition_S(cQ);  // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)360        Tensor tQpQ = make_tensor<bool>(make_shape(size<2>(tQsQ)));361 362        // We don't need to clear the sQ smem tiles since we'll only write out the valid outputs363        flash::copy</*Is_even_MN=*/false, /*Is_even_K=*/true>(gmem_tiled_copy_Q, tQgQ, tQsQ, tQcQ, tQpQ,364                                                              params.seqlen_q - m_block * kBlockM);365 366        const index_t row_offset_k = (bidh / params.h_h_k_ratio) * params.k_head_stride;367        Tensor gK = make_tensor(make_gmem_ptr(reinterpret_cast<Element *>(params.k_ptr) + row_offset_k),368                                Shape<Int<kBlockN>, Int<kHeadDim>>{},369                                make_stride(params.k_row_stride, _1{}));370        typename Kernel_traits::GmemTiledCopy gmem_tiled_copy_K;371        auto gmem_thr_copy_K = gmem_tiled_copy_K.get_thread_slice(tidx - kNThreadsS);372        Tensor tKgK = gmem_thr_copy_K.partition_S(gK);373        Tensor tKsK = gmem_thr_copy_K.partition_D(sK);374        Tensor cK = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK)));  // (BLK_N,BLK_K) -> (blk_n,blk_k)375        Tensor tKcK = gmem_thr_copy_K.partition_S(cK);  // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)376        Tensor tKpK = make_tensor<bool>(make_shape(size<2>(tKsK)));377 378        if (n_block % 2 == 1) {379            // Double buffer for sK380            constexpr int sK_offset = size(sK);381            tKsK.data() = tKsK.data() + sK_offset;382            tOrVt.data() = tOrVt.data() + sK_offset / 8;383        }384 385        // We need to clear the sK smem tiles because K is V.386        const index_t offset_k = cur_block_table * params.k_batch_stride;387        tKgK.data() = tKgK.data() + offset_k;388        flash::copy</*Is_even_MN=*/false, /*Is_even_K=*/true, /*Clear_OOB_MN=*/true>(gmem_tiled_copy_K, tKgK, tKsK, tKcK, tKpK,389                                                                                        seqlen_k - n_block * kBlockN);390        tKgK.data() = tKgK.data() + -offset_k;391        cute::cp_async_fence();392 393        if (n_block - 1 >= n_block_min) {394            cur_block_table = __ldg(&block_table[n_block - 1]);395        }396 397#pragma unroll 1398        for (; n_block >= n_block_min; --n_block) {399            flash::cp_async_wait<0>();400            __syncthreads();401 402            if (n_block - 1 >= n_block_min) {403                // Double buffer for sK404                const int sK_offset = n_block % 2 == 0 ? size(sK) : -size(sK);405                tKsK.data() = tKsK.data() + sK_offset;406 407                const index_t offset_k = cur_block_table * params.k_batch_stride;408                tKgK.data() = tKgK.data() + offset_k;409                flash::copy</*Is_even_MN=*/true, /*Is_even_K=*/true>(gmem_tiled_copy_K, tKgK, tKsK, tKcK, tKpK);410                tKgK.data() = tKgK.data() + -offset_k;411                cute::cp_async_fence();412            }413 414            cutlass::arch::NamedBarrier::sync(kNThreads, static_cast<int>(NamedBarriers::SReady));415 416            if (n_block - 2 >= n_block_min) {417                cur_block_table = __ldg(&block_table[n_block - 2]);418            }419 420            typename Kernel_traits::TiledMma tiled_mma;421            auto tSrS_layout = partition_fragment_C(tiled_mma, Shape<Int<kBlockM>, Int<kBlockN>>{}).layout();422            Tensor rP = make_tensor<Element>(tSrS_layout);423            Tensor scale_o = make_tensor<float>(Shape<_2>{});424            cute::copy(tScale_osScale_o, scale_o);425            cute::copy(tPsP, rP);426 427            flash::rescale_o(tOrO, scale_o);428 429            Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<Kernel_traits::TiledMma>(rP.layout()));430            flash::gemm</*zero_init=*/false, /*wg_wait=*/0>(tiled_mma_o, tOrP, tOrVt, tOrO);431 432            // Double buffer for sK433            const int sK_offset = n_block % 2 == 0 ? size(sK) : -size(sK);434            tOrVt.data() = tOrVt.data() + sK_offset / 8;435        }436 437        cutlass::arch::NamedBarrier::sync(kNThreads, static_cast<int>(NamedBarriers::SoftmaxReady));438        cute::copy(tRow_maxsRow_max, softmax.row_max);439        cute::copy(tRow_sumsRow_sum, softmax.row_sum);440    }441 442    if (NoSplit)443        store<Kernel_traits, false>(params, bidb, bidh, m_block, n_split_idx, shared_storage, tOrO, softmax);444    else445        store<Kernel_traits, true>(params, bidb, bidh, m_block, n_split_idx, shared_storage, tOrO, softmax);446}447 448template<typename Kernel_traits, bool Is_causal, typename SharedStorage>449__global__ void __launch_bounds__(Kernel_traits::kNThreads, 1, 1)450flash_fwd_splitkv_mla_kernel(__grid_constant__ const Flash_fwd_mla_params params) {451    constexpr int kBlockN = Kernel_traits::kBlockN;452    const int m_block = blockIdx.x;453    const int bidh = blockIdx.y;454    const int partition_idx = blockIdx.z;455 456    extern __shared__ char shared_memory[];457    auto &shared_storage = *reinterpret_cast<SharedStorage *>(shared_memory);458 459    int *tile_scheduler_metadata_ptr = params.tile_scheduler_metadata_ptr + partition_idx * TileSchedulerMetaDataSize;460    int4 tile_scheduler_metadata = __ldg(reinterpret_cast<int4 *>(tile_scheduler_metadata_ptr));461    int begin_idx = tile_scheduler_metadata.x;462    int begin_seqlen = tile_scheduler_metadata.y;463    int end_idx = tile_scheduler_metadata.z;464    int end_seqlen = tile_scheduler_metadata.w;465    if (begin_idx >= params.b) return;466    int begin_n_split_idx = __ldg(tile_scheduler_metadata_ptr + 4);467 468#pragma unroll 1469    for (int batch_id = begin_idx; batch_id <= end_idx; ++batch_id) {470        const int n_split_idx = batch_id == begin_idx ? begin_n_split_idx : 0;471        const int seqlen_k = __ldg(params.cu_seqlens_k + batch_id);472        const int n_block_min = batch_id == begin_idx ? begin_seqlen / kBlockN : 0;473        const int n_block_max = batch_id == end_idx ? cute::ceil_div(end_seqlen, kBlockN) : cute::ceil_div(seqlen_k, kBlockN);474        const bool NoSplit = n_block_min == 0 && n_block_max == cute::ceil_div(seqlen_k, kBlockN);475        if (batch_id > begin_idx) {476            __syncthreads();  // Barrier between two tiles.477        }478        flash::compute_attn_1rowblock_splitkv_mla<Kernel_traits, Is_causal>(params, batch_id, bidh, m_block, n_split_idx, seqlen_k, n_block_min, n_block_max, NoSplit, shared_storage);479    }480}481 482////////////////////////////////////////////////////////////////////////////////////////////////////483 484template<typename Element, typename ElementAccum, typename index_t, int kHeadDimV, int kMaxSplits>485__global__ void __launch_bounds__(256, 1, 1)486flash_fwd_splitkv_mla_combine_kernel(__grid_constant__ const Flash_fwd_mla_params params) {487    constexpr int kNThreads = 128;488 489    const int tidx = threadIdx.x;490    const int bidx = blockIdx.x;491    const int hs = params.h * params.seqlen_q;492    const int batch_idx = bidx / hs;493    const int hs_idx = bidx % hs;494 495    const int split_offset = __ldg(params.num_splits_ptr + batch_idx);496    const int actual_num_splits = __ldg(params.num_splits_ptr + batch_idx + 1) - split_offset;497    FLASH_DEVICE_ASSERT(actual_num_splits <= kMaxSplits);498    if (actual_num_splits == 1) return;499 500    __shared__ ElementAccum sLseScale[kMaxSplits];501 502    const index_t row_offset_lseaccum = split_offset * hs + hs_idx;503    const index_t row_offset_lse = bidx;504    Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),505                                   Shape<Int<kMaxSplits>>{}, make_stride(hs));506    Tensor gLSE = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lse_ptr) + row_offset_lse),507                              Shape<_1>{}, Stride<_1>{});508 509    int warp_idx = cutlass::canonical_warp_idx_sync();510    if (warp_idx == 0) {511        constexpr int kNLsePerThread = cute::ceil_div(kMaxSplits, 32);512 513        float local_lse[kNLsePerThread];514        for (int i = 0; i < kNLsePerThread; ++i) {515            const int split = i * 32 + tidx;516            local_lse[i] = split < actual_num_splits ? gLSEaccum(split) : -INFINITY;517        }518 519        float max_lse = -INFINITY;520        for (int i = 0; i < kNLsePerThread; ++i) max_lse = max(max_lse, local_lse[i]);521        for (int offset = 16; offset >= 1; offset /= 2) max_lse = max(max_lse, __shfl_xor_sync(uint32_t(-1), max_lse, offset));522        max_lse = max_lse == -INFINITY ? 0.0f : max_lse;  // In case all local LSEs are -inf523 524        float sum_lse = 0;525        for (int i = 0; i < kNLsePerThread; ++i) sum_lse = sum_lse + expf(local_lse[i] - max_lse);526        for (int offset = 16; offset >= 1; offset /= 2) sum_lse = sum_lse + __shfl_xor_sync(uint32_t(-1), sum_lse, offset);527 528        float global_lse = (sum_lse == 0.f || sum_lse != sum_lse) ? INFINITY : logf(sum_lse) + max_lse;529        if (tidx == 0) gLSE(0) = global_lse;530 531        for (int i = 0; i < kNLsePerThread; ++i) {532            const int split = i * 32 + tidx;533            if (split < actual_num_splits) sLseScale[split] = expf(local_lse[i] - global_lse);534        }535    }536    __syncthreads();537 538    static_assert(kHeadDimV % kNThreads == 0);539    constexpr int Elements = kHeadDimV / kNThreads;540    const index_t row_offset_oaccum = (split_offset * hs + hs_idx) * kHeadDimV;541    Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.oaccum_ptr) + row_offset_oaccum),542                                 Shape<Int<kHeadDimV>>{}, Stride<_1>{});543    using GmemTiledCopyOaccum = decltype(make_tiled_copy(544            Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementAccum>{},545            Layout<Shape<Int<kNThreads>>>{},546            Layout<Shape<Int<Elements>>>{}));547    GmemTiledCopyOaccum gmem_tiled_copy_Oaccum;548    auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx);549    Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_S(gOaccum);550    Tensor tOrOaccum = make_tensor<ElementAccum>(shape(tOgOaccum));551    Tensor tOrO = make_tensor<ElementAccum>(shape(tOgOaccum));552    clear(tOrO);553 554    for (int split = 0; split < actual_num_splits; ++split) {555        cute::copy(tOgOaccum, tOrOaccum);556        ElementAccum lse_scale = sLseScale[split];557        for (int i = 0; i < size(tOrO); ++i) {558            tOrO(i) += lse_scale * tOrOaccum(i);559        }560        tOgOaccum.data() = tOgOaccum.data() + hs * kHeadDimV;561    }562 563    Tensor rO = flash::convert_type<Element>(tOrO);564    const int head_idx = (bidx - batch_idx * hs) / params.seqlen_q;565    const int row = bidx - batch_idx * hs - head_idx * params.seqlen_q;566    auto o_ptr = reinterpret_cast<Element *>(params.o_ptr) + batch_idx * params.o_batch_stride + head_idx * params.o_head_stride + row * params.o_row_stride;567    Tensor gO = make_tensor(make_gmem_ptr(o_ptr + tidx * Elements), Shape<Int<decltype(size<0>(rO))::value>>{}, Stride<_1>{});568    cute::copy(rO, gO);569}570 571} // namespace flash572 573////////////////////////////////////////////////////////////////////////////////////////////////////574 575template<typename Kernel_traits, typename SharedStorage>576void run_flash_splitkv_fwd_mla(Flash_fwd_mla_params &params, cudaStream_t stream) {577    FLASH_ASSERT(params.page_block_size == Kernel_traits::kBlockN);578    const int num_m_block = cute::ceil_div(params.seqlen_q, Kernel_traits::kBlockM);579    BOOL_SWITCH(params.is_causal, Is_causal, [&] {580        auto kernel = &flash::flash_fwd_splitkv_mla_kernel<Kernel_traits, Is_causal, SharedStorage>;581        constexpr size_t smem_size = sizeof(SharedStorage);582        CHECK_CUDA(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));583        kernel<<<dim3(num_m_block, params.h, params.num_sm_parts), Kernel_traits::kNThreads, smem_size, stream>>>(params);584    });585    CHECK_CUDA_KERNEL_LAUNCH();586 587    dim3 grid_combine(params.b * params.h * params.seqlen_q);588    MLA_NUM_SPLITS_SWITCH(params.num_sm_parts, kMaxSplits, [&] {589        auto combine_kernel = &flash::flash_fwd_splitkv_mla_combine_kernel<590                typename Kernel_traits::Element, typename Kernel_traits::ElementAccum, typename Kernel_traits::index_t, Kernel_traits::kHeadDimV, kMaxSplits>;591        combine_kernel<<<grid_combine, 128, 0, stream>>>(params);592    });593    CHECK_CUDA_KERNEL_LAUNCH();594}595 596template<typename T, int Headdim>597void run_mha_fwd_splitkv_mla(Flash_fwd_mla_params &params, cudaStream_t stream) {598    static_assert(Headdim == 576);599    FLASH_ASSERT(params.d_v == 512);600    FLASH_ASSERT(params.k_ptr == params.v_ptr);  // Shared_KV601    using Kernel_traits = Flash_fwd_kernel_traits_mla<576, 64, 64, 8, T, 512>;602    run_flash_splitkv_fwd_mla<Kernel_traits, flash::SharedStorageMLA<Kernel_traits>>(params, stream);603}604