kernels-community/flash-mla
51k
1#pragma once2 3////////////////////////////////////////////////////////////////////////////////////////////////////4 5struct Flash_fwd_mla_params {6 using index_t = int64_t;7 8 int b, seqlen_q, d, d_v;9 int h, h_h_k_ratio, ngroups;10 bool is_causal;11 float scale_softmax, scale_softmax_log2;12 int *__restrict__ cu_seqlens_k;13 14 void *__restrict__ q_ptr;15 void *__restrict__ k_ptr;16 void *__restrict__ v_ptr;17 void *__restrict__ o_ptr;18 void *__restrict__ softmax_lse_ptr;19 20 index_t q_batch_stride;21 index_t k_batch_stride;22 index_t v_batch_stride;23 index_t o_batch_stride;24 index_t q_row_stride;25 index_t k_row_stride;26 index_t v_row_stride;27 index_t o_row_stride;28 index_t q_head_stride;29 index_t k_head_stride;30 index_t v_head_stride;31 index_t o_head_stride;32 33 int *__restrict__ block_table;34 index_t block_table_batch_stride;35 int page_block_size;36 37 int *__restrict__ tile_scheduler_metadata_ptr;38 int num_sm_parts;39 int *__restrict__ num_splits_ptr;40 41 void *__restrict__ softmax_lseaccum_ptr;42 void *__restrict__ oaccum_ptr;43};44 45static constexpr int TileSchedulerMetaDataSize = 8;46// [begin_idx, begin_seqlen, end_idx, end_seqlen, begin_n_split_idx, _, _, _]47 48////////////////////////////////////////////////////////////////////////////////////////////////////49 50template<typename T, int Headdim>51void run_mha_fwd_splitkv_mla(Flash_fwd_mla_params ¶ms, cudaStream_t stream);52 53struct Mla_metadata_params {54 int *__restrict__ seqlens_k_ptr;55 int *__restrict__ tile_scheduler_metadata_ptr;56 int *__restrict__ num_splits_ptr;57 int batch_size;58 int block_size_n;59 int fixed_overhead_num_blocks;60 int num_sm_parts;61};62 63void get_mla_metadata_func(Mla_metadata_params ¶ms, cudaStream_t stream);64 