diff --git a/.codespellrc b/.codespellrc index 16a3b983a..4f60f6084 100644 --- a/.codespellrc +++ b/.codespellrc @@ -1,3 +1,3 @@ [codespell] -ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather, MIS, medias, allready, inout, nd, fo, visibles, nothink, renderD, ond, tbe, CopyIn, notin, subtile, subtiles, IST +ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather, MIS, medias, allready, inout, nd, fo, visibles, nothink, renderD, ond, tbe, CopyIn, notin, subtile, subtiles, dout, IST skip = *.json, *.jsonl, *.patch, *.txt, *.lock diff --git a/sgl-kernel/CMakeLists.txt b/sgl-kernel/CMakeLists.txt index e45aa68b9..294a61b88 100644 --- a/sgl-kernel/CMakeLists.txt +++ b/sgl-kernel/CMakeLists.txt @@ -275,6 +275,8 @@ set(SOURCES "csrc/gemm/gptq/gptq_kernel.cu" "csrc/grammar/apply_token_bitmask_inplace_cuda.cu" + "csrc/infllm_v2/max_pooling.cu" + "csrc/kvcacheio/transfer.cu" "csrc/mamba/causal_conv1d.cu" "csrc/memory/weak_ref_tensor.cpp" @@ -475,6 +477,68 @@ if (SGL_KERNEL_ENABLE_FA3) target_compile_definitions(flash_ops PRIVATE ${FLASH_OPS_COMPILE_DEFS}) endif() +# ===================== InfLLM-V2 FlashAttention backend ===================== # +# Standalone pybind extension `infllm_ops`, vendored from +# 3rdparty/infllmv2_cuda_impl. Kept as its own module so its `flash::` symbols +# stay isolated from sgl-kernel's own flash attention. Mirrors the original +# setup.py: only hdim 64/128 bf16 forward instantiations are compiled (the +# vendored static_switch.h forces bf16 and dispatches headdim to {64, 128} +# only). Backward kernels are intentionally omitted because SGLang only uses +# these ops for inference. +set(INFLLM_FLASH_CUDA_FLAGS + "-DNDEBUG" + "-O3" + "-std=c++17" + "-Xcompiler" + "-fPIC" + "-U__CUDA_NO_HALF_OPERATORS__" + "-U__CUDA_NO_HALF_CONVERSIONS__" + "-U__CUDA_NO_HALF2_OPERATORS__" + "-U__CUDA_NO_BFLOAT16_CONVERSIONS__" + "--expt-relaxed-constexpr" + "--expt-extended-lambda" + "--use_fast_math" + "-DFLASHATTENTION_DISABLE_DROPOUT" + "-DFLASHATTENTION_DISABLE_ALIBI" + "-DFLASHATTENTION_DISABLE_SOFTCAP" + "-DFLASHATTENTION_DISABLE_UNEVEN_K" + "-DFLASHATTENTION_DISABLE_LOCAL" + "--threads=${SGL_KERNEL_COMPILE_THREADS}" +) + +# Arch gencodes: match the original setup.py auto-detection +# (80 always; 90 for CUDA>=11.8; 120 for CUDA>=12.8). +if (ENABLE_BELOW_SM90) + list(APPEND INFLLM_FLASH_CUDA_FLAGS "-gencode=arch=compute_80,code=sm_80") +endif() +list(APPEND INFLLM_FLASH_CUDA_FLAGS "-gencode=arch=compute_90,code=sm_90") +if ("${CUDA_VERSION}" VERSION_GREATER_EQUAL "12.8" OR SGL_KERNEL_ENABLE_SM100A) + list(APPEND INFLLM_FLASH_CUDA_FLAGS "-gencode=arch=compute_120a,code=sm_120a") +endif() + +set(INFLLM_FLASH_SOURCES + "csrc/infllm_v2/flash_extension.cc" + "csrc/infllm_v2/flash_attn/flash_api.cpp" + "csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_sm80.cu" + "csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_sm80.cu" + "csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_causal_sm80.cu" + "csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_causal_sm80.cu" +) + +Python_add_library(infllm_ops MODULE WITH_SOABI ${INFLLM_FLASH_SOURCES}) +target_compile_options(infllm_ops PRIVATE $<$:${INFLLM_FLASH_CUDA_FLAGS}>) +target_include_directories(infllm_ops PRIVATE + ${repo-cutlass_SOURCE_DIR}/include + ${repo-cutlass_SOURCE_DIR}/tools/util/include + ${CMAKE_CURRENT_LIST_DIR}/csrc/infllm_v2/flash_attn + ${CMAKE_CURRENT_LIST_DIR}/csrc/infllm_v2/flash_attn/src +) +# The pybind module binds functions taking at::Generator, which pulls in +# THPGeneratorClass from libtorch_python (not part of TORCH_LIBRARIES). +find_library(TORCH_PYTHON_LIBRARY torch_python PATHS "${TORCH_INSTALL_PREFIX}/lib" REQUIRED) +target_link_libraries(infllm_ops PRIVATE ${TORCH_LIBRARIES} ${TORCH_PYTHON_LIBRARY} c10 cuda) +install(TARGETS infllm_ops LIBRARY DESTINATION "sgl_kernel") + # Build spatial_ops as a separate, optional extension for green contexts set(SPATIAL_SOURCES "csrc/spatial/greenctx_stream.cu" diff --git a/sgl-kernel/csrc/common_extension.cc b/sgl-kernel/csrc/common_extension.cc index f6261dbc7..c3a02943e 100644 --- a/sgl-kernel/csrc/common_extension.cc +++ b/sgl-kernel/csrc/common_extension.cc @@ -49,6 +49,15 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { m.impl("cutlass_mla_decode", torch::kCUDA, &cutlass_mla_decode); m.def("cutlass_mla_get_workspace_size", &cutlass_mla_get_workspace_size); + /* + * From csrc/infllm_v2 + */ + m.def( + "infllm_v2_max_pooling_1d_varlen(Tensor input, Tensor! output, Tensor cu_seqlens_q, Tensor cu_seqlens_k, " + "Tensor cache_lens, int max_seqlen_q, int max_seqlen_k, int kernel_size, int stride, int padding, " + "int block_size, int local_blocks, int init_blocks, int total_q) -> ()"); + m.impl("infllm_v2_max_pooling_1d_varlen", torch::kCUDA, &infllm_v2_max_pooling_1d_varlen); + /* * From csrc/elementwise */ diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/flash_api.cpp b/sgl-kernel/csrc/infllm_v2/flash_attn/flash_api.cpp new file mode 100644 index 000000000..50a2b64db --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/flash_api.cpp @@ -0,0 +1,485 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +// Include these 2 headers instead of torch/extension.h since we don't need all of the torch headers. +#include // For at::Generator and at::PhiloxCudaState +#include +#include +#include +#include +#include + +#include "flash.h" +#include "hardware_info.h" +#include "philox_unpack.cuh" // For at::cuda::philox::unpack +#include "static_switch.h" + +#define CHECK_DEVICE(x) TORCH_CHECK(x.is_cuda(), #x " must be on CUDA") +#define CHECK_SHAPE(x, ...) \ + TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), #x " must have shape (" #__VA_ARGS__ ")") +#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous") + +void set_params_fprop( + Flash_fwd_params& params, + // sizes + const size_t b, + const size_t seqlen_q, + const size_t seqlen_k, + const size_t seqlen_q_rounded, + const size_t seqlen_k_rounded, + const size_t h, + const size_t h_k, + const size_t d, + const size_t d_rounded, + // device pointers + const at::Tensor q, + const at::Tensor k, + const at::Tensor v, + at::Tensor out, + void* cu_seqlens_q_d, + void* cu_seqlens_k_d, + void* seqused_k, + void* p_d, + void* softmax_lse_d, + float p_dropout, + float softmax_scale, + int window_size_left, + int window_size_right, + const float softcap, + bool seqlenq_ngroups_swapped = false, + const bool unpadded_lse = false) { + // Reset the parameters + params = {}; + + params.is_bf16 = q.dtype() == torch::kBFloat16; + + // Set the pointers and strides. + params.q_ptr = q.data_ptr(); + params.k_ptr = k.data_ptr(); + params.v_ptr = v.data_ptr(); + // All stride are in elements, not bytes. + params.q_row_stride = q.stride(-3); + params.k_row_stride = k.stride(-3); + params.v_row_stride = v.stride(-3); + params.q_head_stride = q.stride(-2); + params.k_head_stride = k.stride(-2); + params.v_head_stride = v.stride(-2); + params.o_ptr = out.data_ptr(); + params.o_row_stride = params.o_ptr ? out.stride(-3) : 0; + params.o_head_stride = params.o_ptr ? out.stride(-2) : 0; + + if (cu_seqlens_q_d == nullptr) { + params.q_batch_stride = q.stride(0); + params.k_batch_stride = k.stride(0); + params.v_batch_stride = v.stride(0); + params.o_batch_stride = params.o_ptr ? out.stride(0) : 0; + if (seqlenq_ngroups_swapped) { + params.q_batch_stride *= seqlen_q; + params.o_batch_stride *= seqlen_q; + } + } + + params.cu_seqlens_q = static_cast(cu_seqlens_q_d); + params.cu_seqlens_k = static_cast(cu_seqlens_k_d); + params.seqused_k = static_cast(seqused_k); + + // P = softmax(QK^T) + params.p_ptr = p_d; + + // Softmax sum + params.softmax_lse_ptr = softmax_lse_d; + + // Set the dimensions. + params.b = b; + params.h = h; + params.h_k = h_k; + params.h_h_k_ratio = h / h_k; + params.seqlen_q = seqlen_q; + params.seqlen_k = seqlen_k; + params.seqlen_q_rounded = seqlen_q_rounded; + params.seqlen_k_rounded = seqlen_k_rounded; + params.d = d; + params.d_rounded = d_rounded; + +// Set the different scale values. +#ifdef FLASHATTENTION_DISABLE_SOFTCAP + TORCH_CHECK(softcap <= 0.0, "This flash attention build does not support softcap."); +#endif + if (softcap > 0.0) { + params.softcap = softmax_scale / softcap; + params.scale_softmax = softcap; + params.scale_softmax_log2 = softcap * M_LOG2E; + } else { + // Remove potential NaN + params.softcap = 0.0; + params.scale_softmax = softmax_scale; + params.scale_softmax_log2 = softmax_scale * M_LOG2E; + } + + // Set this to probability of keeping an element to simplify things. + params.p_dropout = 1.f - p_dropout; + // Convert p from float to int so we don't have to convert the random uint to float to compare. + // [Minor] We want to round down since when we do the comparison we use <= instead of < + // params.p_dropout_in_uint = uint32_t(std::floor(params.p_dropout * 4294967295.0)); + // params.p_dropout_in_uint16_t = uint16_t(std::floor(params.p_dropout * 65535.0)); + params.p_dropout_in_uint8_t = uint8_t(std::floor(params.p_dropout * 255.0)); + params.rp_dropout = 1.f / params.p_dropout; + params.scale_softmax_rp_dropout = params.rp_dropout * params.scale_softmax; + TORCH_CHECK(p_dropout < 1.f); +#ifdef FLASHATTENTION_DISABLE_DROPOUT + TORCH_CHECK(p_dropout == 0.0f, "This flash attention build does not support dropout."); +#endif + + // Causal is the special case where window_size_right == 0 and window_size_left < 0. + // Local is the more general case where window_size_right >= 0 or window_size_left >= 0. + params.is_causal = window_size_left < 0 && window_size_right == 0; + + if (window_size_left < 0 && window_size_right >= 0) { + window_size_left = seqlen_k; + } + if (window_size_left >= 0 && window_size_right < 0) { + window_size_right = seqlen_k; + } + params.window_size_left = window_size_left; + params.window_size_right = window_size_right; + +#ifdef FLASHATTENTION_DISABLE_LOCAL + TORCH_CHECK( + params.is_causal || (window_size_left < 0 && window_size_right < 0), + "This flash attention build does not support local attention."); +#endif + + params.is_seqlens_k_cumulative = true; + +#ifdef FLASHATTENTION_DISABLE_UNEVEN_K + TORCH_CHECK(d == d_rounded, "This flash attention build does not support headdim not being a multiple of 32."); +#endif + + params.unpadded_lse = unpadded_lse; + params.seqlenq_ngroups_swapped = seqlenq_ngroups_swapped; +} + +void run_mha_fwd_split_stage1(Flash_fwd_params& params, cudaStream_t stream) { + FP16_SWITCH(!params.is_bf16, [&] { + HEADDIM_SWITCH(params.d, [&] { + BOOL_SWITCH(params.is_causal, Is_causal, [&] { + run_mha_fwd_splitkv_dispatch(params, stream); + }); + }); + }); +} + +void set_params_alibi( + Flash_fwd_params& params, c10::optional& alibi_slopes_, int batch_size, int num_heads) { +#ifdef FLASHATTENTION_DISABLE_ALIBI + TORCH_CHECK(!alibi_slopes_.has_value(), "This flash attention build does not support alibi."); + params.alibi_slopes_ptr = nullptr; +#else + if (alibi_slopes_.has_value()) { + auto alibi_slopes = alibi_slopes_.value(); + TORCH_CHECK(alibi_slopes.dtype() == torch::kFloat32, "ALiBi slopes must have dtype fp32"); + CHECK_DEVICE(alibi_slopes); + TORCH_CHECK(alibi_slopes.stride(-1) == 1, "ALiBi slopes tensor must have contiguous last dimension"); + TORCH_CHECK( + alibi_slopes.sizes() == torch::IntArrayRef({num_heads}) || + alibi_slopes.sizes() == torch::IntArrayRef({batch_size, num_heads})); + params.alibi_slopes_ptr = alibi_slopes.data_ptr(); + params.alibi_slopes_batch_stride = alibi_slopes.dim() == 2 ? alibi_slopes.stride(0) : 0; + } else { + params.alibi_slopes_ptr = nullptr; + } +#endif +} + +std::vector mha_varlen_fwd_stage1( + at::Tensor& q, // total_q x num_heads x head_size, total_q := \sum_{i=0}^{b} s_i + const at::Tensor& k, // total_k x num_heads_k x head_size, total_k := \sum_{i=0}^{b} s_i or num_blocks x + // page_block_size x num_heads_k x head_size if there's a block_table. + const at::Tensor& v, // total_k x num_heads_k x head_size, total_k := \sum_{i=0}^{b} s_i or num_blocks x + // page_block_size x num_heads_k x head_size if there's a block_table. + c10::optional& out_, // total_q x num_heads x head_size, total_k := \sum_{i=0}^{b} s_i + const at::Tensor& cu_seqlens_q, // b+1 + const at::Tensor& cu_seqlens_k, // b+1 + const at::Tensor& cu_seqlens_v, // b+1 + c10::optional& + seqused_k, // b. If given, only this many elements of each batch element's keys are used. + c10::optional& leftpad_k_, // batch_size + c10::optional& block_table_, // batch_size x max_num_blocks_per_seq + c10::optional& alibi_slopes_, // num_heads or b x num_heads + int max_seqlen_q, + const int max_seqlen_k, + const float p_dropout, + const float softmax_scale, + const bool zero_tensors, + bool is_causal, + int window_size_left, + int window_size_right, + const float softcap, + const bool return_softmax, + c10::optional gen_) { + // Otherwise the kernel will be launched from cuda:0 device + at::cuda::CUDAGuard device_guard{q.device()}; + + auto [cc_major, cc_minor] = get_compute_capability(get_current_device()); + // bool is_sm75 = cc_major == 7 && cc_minor == 5; + bool is_sm8x = cc_major == 8 && cc_minor >= 0; + bool is_sm90 = cc_major == 9 && cc_minor == 0; + // TORCH_CHECK(is_sm90 || is_sm8x, "FlashAttention only supports Ampere GPUs or newer."); + // We will support Turing in the near future + // TORCH_CHECK(is_sm90 || is_sm8x || is_sm75, "FlashAttention only supports Turing GPUs or newer."); + + auto q_dtype = q.dtype(); + TORCH_CHECK( + q_dtype == torch::kFloat16 || q_dtype == torch::kBFloat16, "FlashAttention only support fp16 and bf16 data type"); + if (q_dtype == torch::kBFloat16) { + // TORCH_CHECK(is_sm90 || is_sm8x, "bfloat16 is only supported on Ampere GPUs or newer"); + } + TORCH_CHECK(k.dtype() == q_dtype, "query and key must have the same dtype"); + TORCH_CHECK(v.dtype() == q_dtype, "query and value must have the same dtype"); + TORCH_CHECK(cu_seqlens_q.dtype() == torch::kInt32, "cu_seqlens_q must have dtype int32"); + TORCH_CHECK(cu_seqlens_k.dtype() == torch::kInt32, "cu_seqlens_k must have dtype int32"); + TORCH_CHECK(cu_seqlens_v.dtype() == torch::kInt32, "cu_seqlens_v must have dtype int32"); + + CHECK_DEVICE(q); + CHECK_DEVICE(k); + CHECK_DEVICE(v); + CHECK_DEVICE(cu_seqlens_q); + CHECK_DEVICE(cu_seqlens_k); + CHECK_DEVICE(cu_seqlens_v); + + at::Tensor block_table; + const bool paged_KV = block_table_.has_value(); + if (paged_KV) { + block_table = block_table_.value(); + CHECK_DEVICE(block_table); + TORCH_CHECK(block_table.dtype() == torch::kInt32, "block_table must have dtype torch.int32"); + TORCH_CHECK(block_table.stride(-1) == 1, "block_table must have contiguous last dimension"); + } + + TORCH_CHECK(q.stride(-1) == 1, "Input tensor must have contiguous last dimension"); + TORCH_CHECK(k.stride(-1) == 1, "Input tensor must have contiguous last dimension"); + TORCH_CHECK(v.stride(-1) == 1, "Input tensor must have contiguous last dimension"); + CHECK_CONTIGUOUS(cu_seqlens_q); + CHECK_CONTIGUOUS(cu_seqlens_k); + CHECK_CONTIGUOUS(cu_seqlens_v); + + const auto sizes = q.sizes(); + + const int batch_size = cu_seqlens_q.numel() - 1; + int num_heads = sizes[1]; + const int head_size = sizes[2]; + const int num_heads_k = paged_KV ? k.size(2) : k.size(1); + + if (softcap > 0.f) { + TORCH_CHECK(p_dropout == 0.f, "Softcapping does not support dropout for now"); + } + + const int max_num_blocks_per_seq = !paged_KV ? 0 : block_table.size(1); + const int num_blocks = !paged_KV ? 0 : k.size(0); + const int page_block_size = !paged_KV ? 1 : k.size(1); + TORCH_CHECK(!paged_KV || page_block_size % 256 == 0, "Paged KV cache block size must be divisible by 256"); + + if (max_seqlen_q == 1 && !alibi_slopes_.has_value()) { + is_causal = false; + } // causal=true is the same as causal=false in this case + if (is_causal) { + window_size_right = 0; + } + + void* cu_seqlens_q_d = cu_seqlens_q.data_ptr(); + + // Faster to transpose q from (b, 1, (nheads_kv ngroups), d) to (b, ngroups, nheads_kv, d) in this case + // H/t Daniel Haziza + const int seqlenq_ngroups_swapped = max_seqlen_q == 1 && num_heads > num_heads_k && window_size_left < 0 && + window_size_right < 0 && p_dropout == 0.f && head_size % 8 == 0 && + !alibi_slopes_.has_value(); + const int ngroups = num_heads / num_heads_k; + if (seqlenq_ngroups_swapped) { + q = q.reshape({batch_size, num_heads_k, ngroups, head_size}) + .transpose(1, 2) + .reshape({batch_size * ngroups, num_heads_k, head_size}); + max_seqlen_q = ngroups; + num_heads = num_heads_k; + cu_seqlens_q_d = nullptr; + } + + const int total_q = q.sizes()[0]; + + TORCH_CHECK(batch_size > 0, "batch size must be positive"); + TORCH_CHECK(head_size <= 256, "FlashAttention forward only supports head dimension at most 256"); + TORCH_CHECK(head_size % 8 == 0, "query, key, value, and out_ must have a head_size that is a multiple of 8"); + TORCH_CHECK(num_heads % num_heads_k == 0, "Number of heads in key/value must divide number of heads in query"); + + if (window_size_left >= max_seqlen_k) { + window_size_left = -1; + } + if (window_size_right >= max_seqlen_k) { + window_size_right = -1; + } + + CHECK_SHAPE(q, total_q, num_heads, head_size); + if (!paged_KV) { + const int total_k = k.size(0); + CHECK_SHAPE(k, total_k, num_heads_k, head_size); + // CHECK_SHAPE(v, total_k, num_heads_k, head_size); + } else { + CHECK_SHAPE(k, num_blocks, page_block_size, num_heads_k, head_size); + // CHECK_SHAPE(v, num_blocks, page_block_size, num_heads_k, head_size); + CHECK_SHAPE(block_table, batch_size, max_num_blocks_per_seq); + } + + CHECK_SHAPE(cu_seqlens_q, batch_size + 1); + CHECK_SHAPE(cu_seqlens_k, batch_size + 1); + CHECK_SHAPE(cu_seqlens_v, batch_size + 1); + if (seqused_k.has_value()) { + auto seqused_k_ = seqused_k.value(); + TORCH_CHECK(seqused_k_.dtype() == torch::kInt32, "seqused_k must have dtype int32"); + TORCH_CHECK(seqused_k_.is_cuda(), "seqused_k must be on CUDA device"); + TORCH_CHECK(seqused_k_.is_contiguous(), "seqused_k must be contiguous"); + CHECK_SHAPE(seqused_k_, batch_size); + } + + auto opts = q.options(); + at::Tensor out; + out = torch::empty({0}, opts); + // if (out_.has_value()) { + // out = out_.value(); + // TORCH_CHECK(out.dtype() == q_dtype, "Output must have the same dtype as inputs"); + // CHECK_DEVICE(out); + // TORCH_CHECK(out.stride(-1) == 1, "Output tensor must have contiguous last dimension"); + // CHECK_SHAPE(out, sizes[0], sizes[1], head_size); + // if (seqlenq_ngroups_swapped) { + // out = out.reshape({batch_size, num_heads_k, ngroups, head_size}).transpose(1, 2).reshape({batch_size * + // ngroups, num_heads_k, head_size}); + // } + // } else { + // out = torch::empty_like(q); + // } + + auto round_multiple = [](int x, int m) { return (x + m - 1) / m * m; }; + const int head_size_rounded = head_size <= 192 ? round_multiple(head_size, 32) : 256; + const int seqlen_q_rounded = round_multiple(max_seqlen_q, 128); + const int seqlen_k_rounded = round_multiple(max_seqlen_k, 128); + + // auto softmax_lse = torch::empty({num_heads, total_q}, opts.dtype(at::kFloat)); + at::Tensor p; + // Only return softmax if there's dropout to reduce compilation time + if (return_softmax) { + // Return tensor with shape (num_heads_k, total_q, max_seqlen_k) + p = torch::full({num_heads_k, total_q / 16, seqlen_k_rounded}, 0, opts); + } else { + p = torch::empty({0}, opts); + } + + if (zero_tensors) { + // out.zero_(); + // softmax_lse.fill_(-std::numeric_limits::infinity()); + if (return_softmax) { + p.zero_(); + } + } + + Flash_fwd_params params; + set_params_fprop( + params, + batch_size, + max_seqlen_q, + max_seqlen_k, + seqlen_q_rounded, + seqlen_k_rounded, + num_heads, + num_heads_k, + head_size, + head_size_rounded, + q, + k, + v, + out, + cu_seqlens_q_d, + cu_seqlens_k.data_ptr(), + seqused_k.has_value() ? seqused_k.value().data_ptr() : nullptr, + return_softmax ? p.data_ptr() : nullptr, + nullptr, // softmax_lse.data_ptr(), + p_dropout, + softmax_scale, + window_size_left, + window_size_right, + softcap, + seqlenq_ngroups_swapped, + /*unpadded_lse*/ true); + + params.cu_seqlens_v = static_cast(cu_seqlens_v.data_ptr()); + params.is_seqlens_v_cumulative = true; // Treat cu_seqlens_v as cumulative sequence lengths + // { + // // Copy cu_seqlens_v to CPU for printing + // at::Tensor cu_seqlens_v_cpu = cu_seqlens_v.to(torch::kCPU); + // const int* cu_seqlens_v_data = cu_seqlens_v_cpu.data_ptr(); + // printf("params.cu_seqlens_v: "); + // for (int i = 0; i < batch_size + 1; ++i) { + // printf("%d ", cu_seqlens_v_data[i]); + // } + // printf("\n"); + // } + params.total_q = total_q; + + params.m_block_dim = 16; + params.n_block_dim = 1; + + if (paged_KV) { + params.block_table = block_table.data_ptr(); + params.block_table_batch_stride = block_table.stride(0); + params.k_batch_stride = k.stride(0); + // params.v_batch_stride = v.stride(0); + } + params.page_block_size = page_block_size; + // Keep references to these tensors to extend their lifetime + + if (leftpad_k_.has_value()) { + auto leftpad_k = leftpad_k_.value(); + TORCH_CHECK(!paged_KV, "We don't support Paged KV and leftpad_k running at the same time yet"); + TORCH_CHECK(leftpad_k.dtype() == torch::kInt32, "leftpad_k must have dtype int32"); + CHECK_DEVICE(leftpad_k); + CHECK_CONTIGUOUS(leftpad_k); + CHECK_SHAPE(leftpad_k, batch_size); + params.leftpad_k = static_cast(leftpad_k.data_ptr()); + } + + // number of times random will be generated per thread, to offset philox counter in thc random + // state + // We use a custom RNG that increases the offset by batch_size * nheads * 32. + int64_t counter_offset = params.b * params.h * 32; + auto options = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); + auto rng_state = torch::empty({2}, options.dtype(torch::kInt64)); + // Forward kernel will populate memory with the seed and offset. + params.rng_state = reinterpret_cast(rng_state.data_ptr()); + + if (p_dropout > 0.0) { + auto gen = at::get_generator_or_default(gen_, at::cuda::detail::getDefaultCUDAGenerator()); + // See Note [Acquire lock when using random generators] + std::lock_guard lock(gen->mutex_); + params.philox_args = gen->philox_cuda_state(counter_offset); + } + + set_params_alibi(params, alibi_slopes_, batch_size, num_heads); + + if (max_seqlen_k > 0) { + auto stream = at::cuda::getCurrentCUDAStream().stream(); + params.num_splits = 1; + run_mha_fwd_split_stage1(params, stream); + } else { + // If seqlen_k == 0, then we have an empty tensor. We need to set the output to 0. + // out.zero_(); + // softmax_lse.fill_(std::numeric_limits::infinity()); + } + + if (seqlenq_ngroups_swapped) { + int64_t size_before[] = {batch_size, max_seqlen_q, num_heads_k, head_size}; + int64_t size_after[] = {batch_size, num_heads_k * max_seqlen_q, head_size}; + // out = out.reshape(size_before).transpose(1, 2).reshape(size_after); + q = q.reshape(size_before).transpose(1, 2).reshape(size_after); + // softmax_lse = softmax_lse.reshape({num_heads * max_seqlen_q, batch_size}); + } + + return {p}; +} diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/block_info.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/block_info.h new file mode 100644 index 000000000..64fef60fa --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/block_info.h @@ -0,0 +1,90 @@ +/****************************************************************************** + * Copyright (c) 2023, Tri Dao. + ******************************************************************************/ + +#pragma once + +namespace flash { + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct BlockInfo { + template + __device__ BlockInfo(const Params& params, const int bidb) + : sum_s_q(!Varlen || params.cu_seqlens_q == nullptr ? -1 : params.cu_seqlens_q[bidb]), + sum_s_k( + !Varlen || params.cu_seqlens_k == nullptr || !params.is_seqlens_k_cumulative ? -1 + : params.cu_seqlens_k[bidb]), + sum_s_v( + !Varlen || params.cu_seqlens_v == nullptr || !params.is_seqlens_v_cumulative ? -1 + : params.cu_seqlens_v[bidb]), + actual_seqlen_q( + !Varlen || params.cu_seqlens_q == nullptr ? params.seqlen_q : params.cu_seqlens_q[bidb + 1] - sum_s_q) + // If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb]. + // Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K. + , + leftpad_k(params.leftpad_k == nullptr ? 0 : params.leftpad_k[bidb]), + seqlen_k_cache( + (!Varlen || params.cu_seqlens_k == nullptr + ? params.seqlen_k + : (params.is_seqlens_k_cumulative ? params.cu_seqlens_k[bidb + 1] - sum_s_k + : params.cu_seqlens_k[bidb])) - + leftpad_k), + actual_seqlen_k( + params.seqused_k ? params.seqused_k[bidb] - leftpad_k + : seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew)) + // If is_seqlens_v_cumulative, then seqlen_v is cu_seqlens_v[bidb + 1] - cu_seqlens_v[bidb]. + // Otherwise it's cu_seqlens_v[bidb], i.e., we use cu_seqlens_v to store the sequence lengths of V. + , + leftpad_v(params.leftpad_v == nullptr ? 0 : params.leftpad_v[bidb]), + seqlen_v_cache( + (!Varlen || params.cu_seqlens_v == nullptr + ? params.seqlen_v + : (params.is_seqlens_v_cumulative ? params.cu_seqlens_v[bidb + 1] - sum_s_v + : params.cu_seqlens_v[bidb])) - + leftpad_v), + actual_seqlen_c( + params.seqused_v ? params.seqused_v[bidb] - leftpad_v + : seqlen_v_cache + (params.vnew_ptr == nullptr ? 0 : params.seqlen_vnew)) {} + + template + __forceinline__ __device__ index_t + q_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const { + return sum_s_q == -1 ? bidb * batch_stride : uint32_t(sum_s_q) * row_stride; + } + + template + __forceinline__ __device__ index_t + k_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const { + return sum_s_k == -1 ? bidb * batch_stride + leftpad_k * row_stride : uint32_t(sum_s_k + leftpad_k) * row_stride; + } + + template + __forceinline__ __device__ index_t + v_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const { + return sum_s_v == -1 ? bidb * batch_stride + leftpad_v * row_stride : uint32_t(sum_s_v + leftpad_v) * row_stride; + } + + template + inline __device__ index_t blockmask_q_offset(const index_t m_block_dim, const int bidb) const { + return sum_s_q == -1 ? bidb * (actual_seqlen_q / m_block_dim) : uint32_t(sum_s_q) / m_block_dim; + } + + const int sum_s_q; + const int sum_s_k; + const int sum_s_v; + const int actual_seqlen_q; + // We have to have seqlen_k_cache declared before actual_seqlen_k, otherwise actual_seqlen_k is set to 0. + const int leftpad_k; + const int seqlen_k_cache; + const int actual_seqlen_k; + // We have to have seqlen_v_cache declared before actual_seqlen_c, otherwise actual_seqlen_c is set to 0. + const int leftpad_v; + const int seqlen_v_cache; + const int actual_seqlen_c; +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/dropout.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/dropout.h new file mode 100644 index 000000000..dd8ea18ea --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/dropout.h @@ -0,0 +1,103 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include "philox.cuh" +#include "utils.h" + +namespace flash { + +struct Dropout { + const unsigned long long seed, offset; + const uint8_t p_dropout_in_uint8_t; + + __forceinline__ __device__ Dropout( + const unsigned long long seed, + const unsigned long long offset, + const uint8_t p_dropout_in_uint8_t, + const int bid, + const int hid, + const int tid, + const int nheads) + : seed(seed), offset(offset + (bid * nheads + hid) * 32 + tid % 32), p_dropout_in_uint8_t(p_dropout_in_uint8_t) {} + + template + __forceinline__ __device__ void + apply_dropout(Tensor& tensor_, int block_row_start, int block_col_start, int block_row_stride) { + // convert shape from (4, MMA_M, MMA_N) to (8, MMA_M, MMA_N / 2) + Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_dropout(tensor_.layout())); + using T = typename Engine::value_type; + auto encode_dropout = [](bool keep, T val) { return keep ? val : (encode_dropout_in_sign_bit ? -val : T(0)); }; + static_assert(decltype(size<2>(tensor))::value % 2 == 0); + const uint16_t p_dropout_8bit_in_uint16_t = uint16_t(p_dropout_in_uint8_t); + const uint32_t p_dropout_8bit_in_uint32_t = + (uint32_t(p_dropout_8bit_in_uint16_t) << 16) | uint32_t(p_dropout_8bit_in_uint16_t); +// if (cute::thread0()) { printf("threshold2 = 0x%x\n", p_dropout_8bit_in_uint32_t); } +#pragma unroll + for (int m = 0; m < size<1>(tensor); ++m, block_row_start += block_row_stride) { + uint2 rowcol = make_uint2(block_row_start, block_col_start); +#pragma unroll + for (int n = 0; n < size<2>(tensor) / 2; ++n, ++rowcol.y) { + // if (cute::thread(32, 0)) { printf("m = %d, n = %d, row = %d, col = %d\n", m, n, int(rowcol.x), + // int(rowcol.y));} + uint4 random_uint4 = flash::philox(seed, reinterpret_cast(rowcol), offset); + // if (cute::thread0()) { printf("philox = %u, %d, %d, %d\n", random_uint4.x, random_uint4.y, random_uint4.z, + // random_uint4.w);} + uint8_t (&rnd_8)[16] = reinterpret_cast(random_uint4); + // Special implementation for 16-bit types: we duplicate the threshold to the + // low and high 16 bits of a 32-bit value, then use the f16x2 comparison instruction + // to get a mask. The low 16 bits of the mask will be either 0xffff or 0x0000, + // and the high 16 bits will be either 0xffff or 0x0000, depending on whether + // the random value is less than the threshold. + // We then do a bit-wise AND between the mask and the original value (in 32-bit). + // We're exploiting the fact that floating point comparison is equivalent to integer + // comparison, since we're comparing unsigned integers whose top 8-bits are zero. + if (!encode_dropout_in_sign_bit && + (std::is_same::value || std::is_same::value)) { + uint16_t rnd_16[16]; +#pragma unroll + for (int i = 0; i < 16; i++) { + rnd_16[i] = uint16_t(rnd_8[i]); + } + uint32_t (&rnd_32)[8] = reinterpret_cast(rnd_16); +#pragma unroll + for (int j = 0; j < 2; j++) { + Tensor tensor_uint32 = recast(tensor(_, m, n * 2 + j)); +// if (cute::thread0()) { printf("random = 0x%x, 0x%x, 0x%x, 0x%x\n", rnd_32[j * 4 + 0], rnd_32[j * 4 + 1], rnd_32[j * 4 +// + 2], rnd_32[j * 4 + 3]); } if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", +// tensor_uint32(0), tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); } +#pragma unroll + for (int i = 0; i < 4; i++) { + uint32_t mask; + asm volatile("set.le.u32.f16x2 %0, %1, %2;\n" + : "=r"(mask) + : "r"(rnd_32[j * 4 + i]), "r"(p_dropout_8bit_in_uint32_t)); + tensor_uint32(i) &= mask; + } + // if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", tensor_uint32(0), + // tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); } + } + } else { +#pragma unroll + for (int j = 0; j < 2; j++) { +#pragma unroll + for (int i = 0; i < 8; i++) { + tensor(i, m, n * 2 + j) = + encode_dropout(rnd_8[j * 8 + i] <= p_dropout_in_uint8_t, tensor(i, m, n * 2 + j)); + } + Tensor tensor_uint32 = recast(tensor(_, m, n * 2 + j)); + // if (cute::thread0()) { printf("tensor_uint32 = 0x%x, 0x%x, 0x%x, 0x%x\n", tensor_uint32(0), + // tensor_uint32(1), tensor_uint32(2), tensor_uint32(3)); } + } + } + // // if ((threadIdx.x == 0) && (blockIdx.x == 0) && (blockIdx.y == 0)) { + // // printf("n = %d, ph Philox: %u, %u, %u, %u\n", n, rnd_8.x, rnd_8.y, rnd_8.z, rnd_8.w); + // // } + } + } + } +}; + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash.h new file mode 100644 index 000000000..2eeec97b1 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash.h @@ -0,0 +1,152 @@ +/****************************************************************************** + * Copyright (c) 2023, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include // For at::Generator and at::PhiloxCudaState +#include + +#include + +constexpr int TOTAL_DIM = 0; +constexpr int H_DIM = 1; +constexpr int D_DIM = 2; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +struct Qkv_params { + using index_t = int64_t; + // The QKV matrices. + void* __restrict__ q_ptr; + void* __restrict__ k_ptr; + void* __restrict__ v_ptr; + + // The stride between rows of the Q, K and V matrices. + index_t q_batch_stride; + index_t k_batch_stride; + index_t v_batch_stride; + index_t q_row_stride; + index_t k_row_stride; + index_t v_row_stride; + index_t q_head_stride; + index_t k_head_stride; + index_t v_head_stride; + + // The number of heads. + int h, h_k; + // In the case of multi-query and grouped-query attention (MQA/GQA), nheads_k could be + // different from nheads (query). + int h_h_k_ratio; // precompute h / h_k, +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +struct Flash_fwd_params : public Qkv_params { + // The O matrix (output). + void* __restrict__ o_ptr; + void* __restrict__ oaccum_ptr; + + // The stride between rows of O. + index_t o_batch_stride; + index_t o_row_stride; + index_t o_head_stride; + + // The pointer to the P matrix. + void* __restrict__ p_ptr; + + // The pointer to the softmax sum. + void* __restrict__ softmax_lse_ptr; + void* __restrict__ softmax_lseaccum_ptr; + + // The dimensions. + int b, seqlen_q, seqlen_k, seqlen_v, seqlen_knew, seqlen_vnew, d, seqlen_q_rounded, seqlen_k_rounded, d_rounded, + rotary_dim, total_q; + + // The scaling factors for the kernel. + float scale_softmax; + float scale_softmax_log2; + + // array of length b+1 holding starting offset of each sequence. + int* __restrict__ cu_seqlens_q; + int* __restrict__ cu_seqlens_k; + int* __restrict__ cu_seqlens_v; + int* __restrict__ leftpad_k; + int* __restrict__ leftpad_v; + + // If provided, the actual length of each k sequence. + int* __restrict__ seqused_k; + int* __restrict__ seqused_v; + uint64_t* __restrict__ blockmask; + int m_block_dim, n_block_dim, num_k_heads; + int num_blocks_m, num_blocks_n; + + // The K_new and V_new matrices. + void* __restrict__ knew_ptr; + void* __restrict__ vnew_ptr; + + // The stride between rows of the Q, K and V matrices. + index_t knew_batch_stride; + index_t vnew_batch_stride; + index_t knew_row_stride; + index_t vnew_row_stride; + index_t knew_head_stride; + index_t vnew_head_stride; + + // The cos and sin matrices for rotary embedding. + void* __restrict__ rotary_cos_ptr; + void* __restrict__ rotary_sin_ptr; + + // The indices to index into the KV cache. + int* __restrict__ cache_batch_idx; + + // Paged KV cache + int* __restrict__ block_table; + index_t block_table_batch_stride; + int page_block_size; + + // The dropout probability (probability of keeping an activation). + float p_dropout; + // uint32_t p_dropout_in_uint; + // uint16_t p_dropout_in_uint16_t; + uint8_t p_dropout_in_uint8_t; + + // Scale factor of 1 / (1 - p_dropout). + float rp_dropout; + float scale_softmax_rp_dropout; + + // Local window size + int window_size_left, window_size_right; + float softcap; + + // Random state. + at::PhiloxCudaState philox_args; + + // Pointer to the RNG seed (idx 0) and offset (idx 1). + uint64_t* rng_state; + + bool is_bf16; + bool is_causal; + + // If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb]. + // Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K. + bool is_seqlens_k_cumulative; + + // If is_seqlens_v_cumulative, then seqlen_v is cu_seqlens_v[bidb + 1] - cu_seqlens_v[bidb]. + // Otherwise it's cu_seqlens_v[bidb], i.e., we use cu_seqlens_v to store the sequence lengths of V. + bool is_seqlens_v_cumulative; + + bool is_rotary_interleaved; + + int num_splits; // For split-KV version + + void* __restrict__ alibi_slopes_ptr; + index_t alibi_slopes_batch_stride; + + bool unpadded_lse; // For varlen paths: LSE is in [nheads, total_seqlen_q] format instead of [b, nheads, seqlen_q]. + bool seqlenq_ngroups_swapped; // q has been transposed from (b, 1, (nheads_kv ngroups), d) to (b, ngroups, nheads_kv, + // d). +}; + +template +void run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream); diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_blockmask.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_blockmask.h new file mode 100644 index 000000000..abff41fe5 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_blockmask.h @@ -0,0 +1,108 @@ +#pragma once + +namespace flash { + +class fwdIterator { + public: + template + __device__ fwdIterator( + const Params& params, + const BlockInfo& binfo, + const int kBlockM, + const int kBlockN, + const int batch_idx, + const int head_idx, + const int loop_step_idx, + int n_block_min, + int n_block_max) { // row first + if (params.blockmask == nullptr) { + blockmask_ptr = nullptr; + return; + } + this->cache_seqlen_k = binfo.actual_seqlen_k - binfo.actual_seqlen_q / params.m_block_dim; + this->max_block_idx = cute::ceil_div(binfo.actual_seqlen_k, params.n_block_dim); + this->m_block_dim = params.m_block_dim; + this->n_block_dim = params.n_block_dim; + this->n_block_min = n_block_min; + this->n_block_max = n_block_max; + this->batch_idx = batch_idx; // Store batch_idx for debugging + this->head_idx = head_idx; + + // Calculate the offset for the uint64 blockmask + const int num_blocks_m = params.num_blocks_m; + const int num_blocks_n = params.num_blocks_n; + const int uint64_per_row = (num_blocks_n + 64 - 1) / 64; + const int row_offset = params.cu_seqlens_q != nullptr ? binfo.blockmask_q_offset(m_block_dim, batch_idx) + : batch_idx * params.num_k_heads * params.num_blocks_m; + + blockmask_ptr = params.blockmask + head_idx * params.num_blocks_m * uint64_per_row + row_offset * uint64_per_row + + loop_step_idx * uint64_per_row; + + // printf("blockmask_ptr = %d\n", blockmask_ptr); + + const int q_block_idx = loop_step_idx + cache_seqlen_k; + } + + __device__ int max_no_larger(int target) const { + if (blockmask_ptr == nullptr) { + // printf("blockmask_ptr is nullptr\n"); + return target; + } + // printf("blockmask_ptr is NOT!!!! nullptr\n"); + if (max_block_idx == 0) { + return -1; + }; + + // 目标值不能超过最大块索引 + target = min(target, max_block_idx - 1); + + // 计算相对于当前q_bit_position的实际位置 + int target_bit_pos = target; + + // 确定此块在哪个uint64中 + int uint64_offset = target_bit_pos / 64; + + // 确定此块在uint64中的哪一位 + int bit_pos = target_bit_pos % 64; + + // 创建一个掩码,保留target及更低位的所有位 + uint64_t mask = bit_pos != 63 ? (1ULL << (bit_pos + 1)) - 1 : 0xFFFFFFFFFFFFFFFFULL; + + // 检查当前uint64中target及以下的位 + uint64_t value = blockmask_ptr[uint64_offset] & mask; + + // 如果当前uint64中有设置的位 + int result = -1; + if (value != 0) { + // 找到最高位的1(即不大于target的最大设置位) + int highest_bit = 63 - __clzll(value); // __clzll计算前导0的数量 + result = highest_bit + (uint64_offset * 64); + } else { + // 如果当前uint64中没有找到,检查更低的uint64块 + for (int i = uint64_offset - 1; i >= 0; i--) { + value = blockmask_ptr[i]; + if (value != 0) { + // 找到最高位的1 + int highest_bit = 63 - __clzll(value); + // 计算相对于q_bit_position的偏移 + result = highest_bit + (i * 64); + break; + } + } + } + + // 没有找到设置位 + return result; + } + + uint64_t* blockmask_ptr; + int row_offset; // 行偏移量 + int uint64_per_row; // 每行使用的uint64数量 + int cache_seqlen_k; + int max_block_idx; + int m_block_dim, n_block_dim; + int n_block_min, n_block_max; + int batch_idx, head_idx; +}; + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_kernel.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_kernel.h new file mode 100644 index 000000000..4b1cd47ee --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_kernel.h @@ -0,0 +1,2355 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include +#include +#include + +#include + +#include "block_info.h" +#include "dropout.h" +#include "flash_blockmask.h" +#include "kernel_traits.h" +#include "mask.h" +#include "philox_unpack.cuh" // For at::cuda::philox::unpack +#include "rotary.h" +#include "softmax.h" +#include "utils.h" + +namespace flash { + +using namespace cute; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +__forceinline__ __device__ float minusinf_to_zero(float& input) { + if (input == -INFINITY) { + return 0; + } else { + return input; + } +} + +template +__forceinline__ __device__ void thread_element_wise_reduce_(Tensor& tensor, Operator& op) { + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); mi++) { +#pragma unroll + for (int ni = 0; ni < size<1>(tensor); ni++) { + tensor(mi, ni) = + op(minusinf_to_zero(tensor(mi, ni)), + __shfl_xor_sync(uint32_t(-1), minusinf_to_zero(tensor(mi, ni)), thr_offset)); + } + } +} + +template +__forceinline__ __device__ void hdim16_reduce( + Tensor& acc_S, + Tensor& g_Sh, + const int col_idx_offset_, + const int row_idx_offset_, + const int warp_row_stride) { + // Reshape tensor_ from (MMA=4, MMA_M, MMA_N) (or (_2,_2),_2,_16) for D=32) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + auto tensor = make_tensor(acc_S.data(), flash::convert_layout_acc_rowcol(acc_S.layout())); + const int warp_id = threadIdx.x / 32; + const int lane_id = threadIdx.x % 32; + // const int col_idx_offset = col_idx_offset_ + (lane_id % 4) * 2; + // const int row_idx_offset = row_idx_offset_ + warp_id * 16 + lane_id / 4; + const int col_idx_offset = 0 + (lane_id % 4) * 2; + const int row_idx_offset = 0 + warp_id * 16 + lane_id / 4; + + // step 1: 线程内部求和 (v0 + v2) + using TensorT = decltype(make_tensor(Shape(tensor)>, Int(tensor)>>{})); + TensorT v02; + clear(v02); +#pragma unroll + for (int mi = 0; mi < size<0, 1>(tensor); ++mi) { +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx_v02 = j * size<1, 1>(tensor) + nj; + const float v0 = minusinf_to_zero(tensor(make_coord(0, mi), make_coord(j, nj))); + const float v2 = minusinf_to_zero(tensor(make_coord(1, mi), make_coord(j, nj))); + v02(mi, col_idx_v02) = v0 + v2; + } + } + } + + // step 2: warp 内部蝶形求和 + SumOp sum_op; + thread_element_wise_reduce_<16>(v02, sum_op); + thread_element_wise_reduce_<8>(v02, sum_op); + thread_element_wise_reduce_<4>(v02, sum_op); + + // step 3: copy 到 global mem + cutlass::NumericConverter converter; + if (lane_id < 4) { +#pragma unroll + for (int mi = 0; mi < size<0, 1>(tensor); ++mi) { + const int row_idx_base = row_idx_offset + mi * warp_row_stride; +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; + const int col_idx_v02 = j * size<1, 1>(tensor) + nj; + g_Sh(row_idx_base / 16, col_idx) = converter(v02(mi, col_idx_v02)); // ignore /16 since it's too slow + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__forceinline__ __device__ auto get_lse_tile( + const Params& params, + const int bidb, + const int bidh, + const int m_block, + const BlockInfo& binfo) { + // When params.unpadded_lse is false, LSE is written as (b, h, seqlen_q) - this is non-variable seqlen path. + // Otherwise, when params.seqlenq_ngroups_swapped is true, it is written as (h, seqlen_q, b) to account for seqlen_q + // <-> h swapping trick. Otherwise, it's written as (h, b, seqlen_q). + const bool varlen_q = params.unpadded_lse && !params.seqlenq_ngroups_swapped; + auto lse_offset = varlen_q ? binfo.q_offset(params.seqlen_q, 1, bidb) : 0; + auto gmem_ptr_lse = make_gmem_ptr(reinterpret_cast(params.softmax_lse_ptr) + lse_offset); + + auto lse_shape = varlen_q ? make_shape(1, params.h, params.total_q) : make_shape(params.b, params.h, params.seqlen_q); + auto lse_stride = params.seqlenq_ngroups_swapped + ? make_stride(1, params.seqlen_q * params.b, params.b) + : (params.unpadded_lse ? make_stride(params.h * params.total_q, params.total_q, 1) + : make_stride(params.h * params.seqlen_q, params.seqlen_q, 1)); + + auto lse_layout = make_layout(lse_shape, lse_stride); + Tensor mLSE = make_tensor(gmem_ptr_lse, lse_layout); + auto mLSE_slice = varlen_q ? mLSE(0, bidh, _) : mLSE(bidb, bidh, _); + return local_tile(mLSE_slice, Shape>{}, make_coord(m_block)); +} + +template < + typename Kernel_traits, + bool Is_dropout, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Return_softmax, + typename Params> +inline __device__ void compute_attn_1rowblock(const Params& params, const int bidb, const int bidh, const int m_block) { + using Element = typename Kernel_traits::Element; + using ElementAccum = typename Kernel_traits::ElementAccum; + using index_t = typename Kernel_traits::index_t; + + // Shared memory. + extern __shared__ char smem_[]; + + // The thread index. + const int tidx = threadIdx.x; + + constexpr int kBlockM = Kernel_traits::kBlockM; + constexpr int kBlockN = Kernel_traits::kBlockN; + constexpr int kHeadDim = Kernel_traits::kHeadDim; + constexpr int kNWarps = Kernel_traits::kNWarps; + + auto seed_offset = at::cuda::philox::unpack(params.philox_args); + flash::Dropout dropout( + std::get<0>(seed_offset), std::get<1>(seed_offset), params.p_dropout_in_uint8_t, bidb, bidh, tidx, params.h); + + // Save seed and offset before any early exit so the caller can return RNG state. + if (Is_dropout && blockIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0 && tidx == 0) { + params.rng_state[0] = std::get<0>(seed_offset); + params.rng_state[1] = std::get<1>(seed_offset); + } + + const BlockInfo binfo(params, bidb); + if (m_block * kBlockM >= binfo.actual_seqlen_q) return; + + const int n_block_min = + !Is_local ? 0 + : std::max( + 0, + (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / + kBlockN); + int n_block_max = cute::ceil_div(binfo.actual_seqlen_k, kBlockN); + if (Is_causal || Is_local) { + n_block_max = std::min( + n_block_max, + cute::ceil_div( + (m_block + 1) * kBlockM / params.m_block_dim + binfo.actual_seqlen_k - + (binfo.actual_seqlen_q / params.m_block_dim) + params.window_size_right, + kBlockN)); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { + // printf("m_block = %d, n_block_max = %d\n", m_block, n_block_max); + // } + } + // We exit early and write 0 to gO and gLSE. This also covers the case where actual_seqlen_k == 0. + // Otherwise we might read OOB elements from gK and gV. + if ((Is_causal || Is_local || !Is_even_MN) && n_block_max <= n_block_min) { + Tensor mO = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.o_ptr) + + binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)), + make_shape(binfo.actual_seqlen_q, params.h, params.d), + make_stride(params.o_row_stride, params.o_head_stride, _1{})); + Tensor gO = local_tile( + mO(_, bidh, _), Shape, Int>{}, make_coord(m_block, 0)); // (kBlockM, kHeadDim) + + Tensor gLSE = get_lse_tile(params, bidb, bidh, m_block, binfo); + + typename Kernel_traits::GmemTiledCopyO gmem_tiled_copy_O; + auto gmem_thr_copy_O = gmem_tiled_copy_O.get_thread_slice(tidx); + Tensor tOgO = gmem_thr_copy_O.partition_D(gO); + Tensor tOrO = make_tensor(shape(tOgO)); + clear(tOrO); + // Construct identity layout for sO + Tensor cO = make_identity_tensor(make_shape(size<0>(gO), size<1>(gO))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcO = gmem_thr_copy_O.partition_D(cO); + Tensor tOpO = make_tensor(make_shape(size<2>(tOgO))); + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tOpO); ++k) { + tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d; + } + } + // Clear_OOB_K must be false since we don't want to write zeros to gmem + flash::copy( + gmem_tiled_copy_O, tOrO, tOgO, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM); +#pragma unroll + for (int m = 0; m < size<1>(tOgO); ++m) { + const int row = get<0>(tOcO(0, m, 0)); + if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { + gLSE(row) = INFINITY; + } + } + return; + } + // if (tidx == 0) { printf("m_block = %d, n_block_min = %d, n_block_max = %d\n", m_block, n_block_min, n_block_max); } + + // We iterate over the blocks in reverse order. This is because the last block is the only one + // that needs masking when we read K and V from global memory. Moreover, iterating in reverse + // might save us 1 register (we just need n_block instead of both n_block and n_block_max). + + const index_t row_offset_p = + ((bidb * params.h + bidh) * params.seqlen_q_rounded + m_block * kBlockM) * params.seqlen_k_rounded + + (n_block_max - 1) * kBlockN; + + Tensor mQ = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.q_ptr) + binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)), + make_shape(binfo.actual_seqlen_q, params.h, params.d), + make_stride(params.q_row_stride, params.q_head_stride, _1{})); + Tensor gQ = + local_tile(mQ(_, bidh, _), Shape, Int>{}, make_coord(m_block, 0)); // (kBlockM, kHeadDim) + Tensor mK = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.k_ptr) + binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb)), + make_shape(binfo.actual_seqlen_k, params.h_k, params.d), + make_stride(params.k_row_stride, params.k_head_stride, _1{})); + Tensor gK = local_tile( + mK(_, bidh / params.h_h_k_ratio, _), + Shape, Int>{}, + make_coord(_, 0)); // (kBlockN, kHeadDim, nblocksN) + Tensor mV = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.v_ptr) + binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb)), + make_shape(binfo.actual_seqlen_k, params.h_k, params.d), + make_stride(params.v_row_stride, params.v_head_stride, _1{})); + Tensor gV = local_tile( + mV(_, bidh / params.h_h_k_ratio, _), + Shape, Int>{}, + make_coord(_, 0)); // (kBlockN, kHeadDim, nblocksN) + Tensor gP = make_tensor( + make_gmem_ptr(reinterpret_cast(params.p_ptr) + row_offset_p), + Shape, Int>{}, + make_stride(params.seqlen_k_rounded, _1{})); + + Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutQ{}); + // Careful we're using the same smem for sQ and sK | sV if Share_Q_K_smem; + Tensor sK = + make_tensor(sQ.data() + (Kernel_traits::Share_Q_K_smem ? 0 : size(sQ)), typename Kernel_traits::SmemLayoutKV{}); + Tensor sV = make_tensor(sK.data() + size(sK), typename Kernel_traits::SmemLayoutKV{}); + Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{}); + Tensor sVtNoSwizzle = make_tensor(sV.data().get(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{}); + + typename Kernel_traits::GmemTiledCopyQKV gmem_tiled_copy_QKV; + auto gmem_thr_copy_QKV = gmem_tiled_copy_QKV.get_thread_slice(tidx); + + Tensor tQgQ = gmem_thr_copy_QKV.partition_S(gQ); + Tensor tQsQ = gmem_thr_copy_QKV.partition_D(sQ); + Tensor tKgK = gmem_thr_copy_QKV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K, nblocksN) + Tensor tKsK = gmem_thr_copy_QKV.partition_D(sK); + Tensor tVgV = gmem_thr_copy_QKV.partition_S(gV); // (VCPY, VCPY_N, VCPY_K, nblocksN) + Tensor tVsV = gmem_thr_copy_QKV.partition_D(sV); + + typename Kernel_traits::TiledMma tiled_mma; + auto thr_mma = tiled_mma.get_thread_slice(tidx); + Tensor tSrQ = thr_mma.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K) + Tensor tSrK = thr_mma.partition_fragment_B(sK); // (MMA,MMA_N,MMA_K) + Tensor tOrVt = thr_mma.partition_fragment_B(sVtNoSwizzle); // (MMA, MMA_K,MMA_N) + + Tensor tSgS = thr_mma.partition_C(gP); + + Tensor acc_o = partition_fragment_C(tiled_mma, Shape, Int>{}); // MMA, MMA_M, MMA_K + + // + // Copy Atom retiling + // + + auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx); + // if (cute::thread0()) {smem_thr_copy_Q.print_all();} + Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ); + // if (cute::thread0()) {print(tSsQ.layout()); printf("\n");} + + auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx); + Tensor tSsK = smem_thr_copy_K.partition_S(sK); + + auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma); + auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx); + Tensor tOsVt = smem_thr_copy_V.partition_S(sVt); + + // + // PREDICATES + // + + // // Allocate predicate tensors for m and n + // Tensor tQpQ = make_tensor(make_shape(size<1>(tQsQ), size<2>(tQsQ)), Stride<_1,_0>{}); + // Tensor tKVpKV = make_tensor(make_shape(size<1>(tKsK), size<2>(tKsK)), Stride<_1,_0>{}); + + // Construct identity layout for sQ and sK + Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k) + // Tensor tScQ = thr_mma.partition_A(cQ); // (MMA,MMA_M,MMA_K) + // if (cute::thread0()) { + // print(tScQ.layout()); printf("\n"); + // for (int i = 0; i < size(tScQ); ++i) { + // printf("%d ", get<0>(tScQ(i))); + // } + // printf("\n"); + // for (int i = 0; i < size(tScQ); ++i) { + // printf("%d ", get<1>(tScQ(i))); + // } + // printf("\n"); + // } + + // Repeat the partitioning with identity layouts + Tensor tQcQ = gmem_thr_copy_QKV.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) + Tensor tKVcKV = gmem_thr_copy_QKV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k) + + // Allocate predicate tensors for k + Tensor tQpQ = make_tensor(make_shape(size<2>(tQsQ))); + Tensor tKVpKV = make_tensor(make_shape(size<2>(tKsK))); + + // Set predicates for k bounds + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tQpQ); ++k) { + tQpQ(k) = get<1>(tQcQ(0, 0, k)) < params.d; + } +#pragma unroll + for (int k = 0; k < size(tKVpKV); ++k) { + tKVpKV(k) = get<1>(tKVcKV(0, 0, k)) < params.d; + } + } + + // Prologue + + // We don't need to clear the sQ smem tiles since we'll only write out the valid outputs + flash::copy( + gmem_tiled_copy_QKV, tQgQ, tQsQ, tQcQ, tQpQ, binfo.actual_seqlen_q - m_block * kBlockM); + if (Kernel_traits::Is_Q_in_regs) { + cute::cp_async_fence(); + } + + // // if (cute::thread(1, 0)) { print(tQsQ); } + // // Tensor sQNoSwizzle = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename + // Kernel_traits::SmemLayoutQNoSwizzle{}); + // // if (cute::thread0()) { print(sQNoSwizzle); } + + if (Kernel_traits::Share_Q_K_smem) { + flash::cp_async_wait<0>(); + __syncthreads(); + Tensor tSrQ_copy_view = smem_thr_copy_Q.retile_D(tSrQ); + CUTE_STATIC_ASSERT_V(size<1>(tSsQ) == size<1>(tSrQ_copy_view)); // M + cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ_copy_view); + __syncthreads(); + } + + int n_block = n_block_max - 1; + // We don't need to clear the sK smem tiles since we'll mask out the scores anyway. + flash::copy( + gmem_tiled_copy_QKV, tKgK(_, _, _, n_block), tKsK, tKVcKV, tKVpKV, binfo.actual_seqlen_k - n_block * kBlockN); + cute::cp_async_fence(); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z < 2) { print(tKgK); } + // __syncthreads(); + + if (Kernel_traits::Is_Q_in_regs && !Kernel_traits::Share_Q_K_smem) { + flash::cp_async_wait<1>(); + __syncthreads(); + Tensor tSrQ_copy_view = smem_thr_copy_Q.retile_D(tSrQ); + CUTE_STATIC_ASSERT_V(size<1>(tSsQ) == size<1>(tSrQ_copy_view)); // M + cute::copy(smem_tiled_copy_Q, tSsQ, tSrQ_copy_view); + } + + clear(acc_o); + + flash::Softmax<2 * size<1>(acc_o)> softmax; + + const float alibi_slope = + !Has_alibi || params.alibi_slopes_ptr == nullptr + ? 0.0f + : reinterpret_cast(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / + params.scale_softmax; + flash::Mask mask( + binfo.actual_seqlen_k, + binfo.actual_seqlen_q, + params.window_size_left, + params.window_size_right, + alibi_slope, + params.m_block_dim); + + // Initialize blockmask iterator if blockmask is enabled + fwdIterator blockmask(params, binfo, kBlockM, kBlockN, bidb, bidh, m_block, n_block_min, n_block_max); + int next_block_idx = blockmask.max_no_larger(n_block_max - 1); + int leap = 0; + + // For performance reason, we separate out two kinds of iterations: + // those that need masking on S, and those that don't. + // We need masking on S for the very last block when K and V has length not multiple of kBlockN. + // We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks. + // We will have at least 1 "masking" iteration. + + // If not even_N, then seqlen_k might end in the middle of a block. In that case we need to + // mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1. + constexpr int n_masking_steps = + (!Is_causal && !Is_local) + ? 1 + : ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1); +#pragma unroll + for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) { + const bool skip = (n_block != next_block_idx); + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + // Advance gV + if (masking_step > 0) { + flash::copy(gmem_tiled_copy_QKV, tVgV(_, _, _, n_block), tVsV, tKVcKV, tKVpKV); + } else { + // Clear the smem tiles to account for predicated off loads + flash::copy( + gmem_tiled_copy_QKV, tVgV(_, _, _, n_block), tVsV, tKVcKV, tKVpKV, binfo.actual_seqlen_k - n_block * kBlockN); + } + cute::cp_async_fence(); + + if (!skip) { + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + // if (cute::thread0()) { print(acc_s); } + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + mask.template apply_mask( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + + next_block_idx = blockmask.max_no_larger(n_block - 1); + } else { + mask.all_mask(acc_s); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + + leap = (masking_step + 1 == n_masking_steps) ? n_block - next_block_idx : 1; + + if (n_block > n_block_min && next_block_idx != -1) { + flash::copy( + gmem_tiled_copy_QKV, tKgK(_, _, _, n_block - leap), tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + // TODO: when we have key_padding_mask we'll need to Check_inf + masking_step == 0 ? softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2) + : softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2); + + if (!skip) { + // Convert acc_s from fp32 to fp16/bf16 + Tensor rP = flash::convert_type(acc_s); + int block_row_idx = m_block * (kBlockM / 16) + tidx / 32; + int block_col_idx = n_block * (kBlockN / 32); + if (Return_softmax) { + Tensor rP_drop = make_fragment_like(rP); + cute::copy(rP, rP_drop); + dropout.template apply_dropout( + rP_drop, block_row_idx, block_col_idx, kNWarps); + cute::copy(rP_drop, tSgS); + tSgS.data() = tSgS.data() + (-kBlockN); + } + if (Is_dropout) { + dropout.apply_dropout(rP, block_row_idx, block_col_idx, kNWarps); + } + + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs(rP.layout())); + // if (cute::thread0()) { print(tOrP); } + flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma, smem_tiled_copy_V, smem_thr_copy_V); + // if (cute::thread0()) { print(scores); } + } + + // This check is at the end of the loop since we always have at least 1 iteration + if (n_masking_steps > 1 && n_block <= n_block_min) { + --n_block; + break; + } + } + + leap = n_block - next_block_idx + 1; + + // These are the iterations where we don't need masking on S + for (n_block = next_block_idx; n_block != -1 && n_block >= n_block_min; n_block = next_block_idx) { + next_block_idx = blockmask.max_no_larger(n_block - 1); + + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + flash::copy(gmem_tiled_copy_QKV, tVgV(_, _, _, n_block), tVsV, tKVcKV, tKVpKV); + cute::cp_async_fence(); + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + + leap = n_block - next_block_idx; + if (next_block_idx != -1) { + flash::copy( + gmem_tiled_copy_QKV, tKgK(_, _, _, n_block - leap), tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + mask.template apply_mask( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + + softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2); + + Tensor rP = flash::convert_type(acc_s); + int block_row_idx = m_block * (kBlockM / 16) + tidx / 32; + int block_col_idx = n_block * (kBlockN / 32); + if (Return_softmax) { + Tensor rP_drop = make_fragment_like(rP); + cute::copy(rP, rP_drop); + dropout.template apply_dropout( + rP_drop, block_row_idx, block_col_idx, kNWarps); + cute::copy(rP_drop, tSgS); + tSgS.data() = tSgS.data() + (-kBlockN); + } + if (Is_dropout) { + dropout.apply_dropout(rP, block_row_idx, block_col_idx, kNWarps); + } + + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs(rP.layout())); + flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma, smem_tiled_copy_V, smem_thr_copy_V); + } + + // Epilogue + + Tensor lse = softmax.template normalize_softmax_lse(acc_o, params.scale_softmax, params.rp_dropout); + + // Convert acc_o from fp32 to fp16/bf16 + Tensor rO = flash::convert_type(acc_o); + Tensor sO = make_tensor(sQ.data(), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) + // Partition sO to match the accumulator partitioning + auto smem_tiled_copy_O = make_tiled_copy_C(typename Kernel_traits::SmemCopyAtomO{}, tiled_mma); + auto smem_thr_copy_O = smem_tiled_copy_O.get_thread_slice(tidx); + Tensor taccOrO = smem_thr_copy_O.retile_S(rO); // ((Atom,AtomNum), MMA_M, MMA_N) + Tensor taccOsO = smem_thr_copy_O.partition_D(sO); // ((Atom,AtomNum),PIPE_M,PIPE_N) + + // sO has the same size as sQ, so we don't need to sync here. + if (Kernel_traits::Share_Q_K_smem) { + __syncthreads(); + } + + cute::copy(smem_tiled_copy_O, taccOrO, taccOsO); + + Tensor mO = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.o_ptr) + binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb)), + make_shape(binfo.actual_seqlen_q, params.h, params.d), + make_stride(params.o_row_stride, params.o_head_stride, _1{})); + Tensor gO = + local_tile(mO(_, bidh, _), Shape, Int>{}, make_coord(m_block, 0)); // (kBlockM, kHeadDim) + Tensor gLSE = get_lse_tile(params, bidb, bidh, m_block, binfo); + + typename Kernel_traits::GmemTiledCopyO gmem_tiled_copy_O; + auto gmem_thr_copy_O = gmem_tiled_copy_O.get_thread_slice(tidx); + Tensor tOsO = gmem_thr_copy_O.partition_S(sO); // ((Atom,AtomNum),ATOM_M,ATOM_N) + Tensor tOgO = gmem_thr_copy_O.partition_D(gO); + + __syncthreads(); + + Tensor tOrO = make_tensor(shape(tOgO)); + cute::copy(gmem_tiled_copy_O, tOsO, tOrO); + + Tensor caccO = make_identity_tensor(Shape, Int>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k) + Tensor taccOcO = thr_mma.partition_C(caccO); // (MMA,MMA_M,MMA_K) + static_assert(decltype(size<0>(taccOcO))::value == 4); + // Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices. + Tensor taccOcO_row = logical_divide(taccOcO, Shape<_2>{})(make_coord(0, _), _, 0); + CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M + if (get<1>(taccOcO_row(0)) == 0) { +#pragma unroll + for (int mi = 0; mi < size(lse); ++mi) { + const int row = get<0>(taccOcO_row(mi)); + if (row < binfo.actual_seqlen_q - m_block * kBlockM) { + gLSE(row) = lse(mi); + } + } + } + + // Construct identity layout for sO + Tensor cO = make_identity_tensor(make_shape(size<0>(sO), size<1>(sO))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcO = gmem_thr_copy_O.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) + Tensor tOpO = make_tensor(make_shape(size<2>(tOgO))); + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tOpO); ++k) { + tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d; + } + } + // Clear_OOB_K must be false since we don't want to write zeros to gmem + flash::copy( + gmem_tiled_copy_O, tOrO, tOgO, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + typename Kernel_traits, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV, + typename Params> +inline __device__ void compute_attn_1rowblock_splitkv( + const Params& params, + const int bidb, + const int bidh, + const int m_block, + const int n_split_idx, + const int num_n_splits) { + using Element = typename Kernel_traits::Element; + using ElementAccum = typename Kernel_traits::ElementAccum; + using index_t = typename Kernel_traits::index_t; + + // Shared memory. + extern __shared__ char smem_[]; + + // The thread index. + const int tidx = threadIdx.x; + + constexpr int kBlockM = Kernel_traits::kBlockM; + constexpr int kBlockN = Kernel_traits::kBlockN; + constexpr int kHeadDim = Kernel_traits::kHeadDim; + constexpr int kNWarps = Kernel_traits::kNWarps; + + using GmemTiledCopyO = + std::conditional_t; + using ElementO = std::conditional_t; + + const BlockInfo binfo(params, bidb); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, + // seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, + // binfo.actual_seqlen_k); } if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = + // %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : + // params.seqlen_knew)); } + if (m_block * kBlockM >= binfo.actual_seqlen_q) return; + + const int n_blocks_per_split = ((params.seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits; + const int n_block_min = + !Is_local ? n_split_idx * n_blocks_per_split + : std::max( + n_split_idx * n_blocks_per_split, + (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / + kBlockN); + int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split); + if (Is_causal || Is_local) { + n_block_max = std::min( + n_block_max, + cute::ceil_div( + (m_block + 1) * kBlockM / params.m_block_dim + binfo.actual_seqlen_k - + (binfo.actual_seqlen_q / params.m_block_dim) + params.window_size_right, + kBlockN)); + } + if (n_block_min >= n_block_max) { // This also covers the case where n_block_max <= 0 + // We exit early and write 0 to gOaccum and -inf to gLSEaccum. + // Otherwise we might read OOB elements from gK and gV, + // or get wrong results when we combine gOaccum from different blocks. + const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb) + + m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride; + const index_t row_offset_oaccum = + (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM) * params.d_rounded; + const index_t row_offset_lseaccum = + ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM; + Tensor gOaccum = make_tensor( + make_gmem_ptr( + reinterpret_cast(Split ? params.oaccum_ptr : params.o_ptr) + + (Split ? row_offset_oaccum : row_offset_o)), + Shape, Int>{}, + make_stride(Split ? kHeadDim : params.o_row_stride, _1{})); + Tensor gLSEaccum = make_tensor( + make_gmem_ptr( + reinterpret_cast(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + + row_offset_lseaccum), + Shape>{}, + Stride<_1>{}); + + GmemTiledCopyO gmem_tiled_copy_Oaccum; + auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx); + Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + Tensor tOrOaccum = make_tensor(shape(tOgOaccum)); + clear(tOrOaccum); + // Construct identity layout for sO + Tensor cO = make_identity_tensor(make_shape(size<0>(gOaccum), size<1>(gOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); + Tensor tOpO = make_tensor(make_shape(size<2>(tOgOaccum))); + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tOpO); ++k) { + tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d; + } + } + // Clear_OOB_K must be false since we don't want to write zeros to gmem + flash::copy( + gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM); +#pragma unroll + for (int m = 0; m < size<1>(tOgOaccum); ++m) { + const int row = get<0>(tOcO(0, m, 0)); + if (row < binfo.actual_seqlen_q - m_block * kBlockM && get<1>(tOcO(0, m, 0)) == 0) { + gLSEaccum(row) = Split ? -INFINITY : INFINITY; + } + } + return; + } + + // We iterate over the blocks in reverse order. This is because the last block is the only one + // that needs masking when we read K and V from global memory. Moreover, iterating in reverse + // might save us 1 register (we just need n_block instead of both n_block and n_block_max). + + // We move K and V to the last block. + const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb]; + const int* block_table = + params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride; + const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size; + const int block_table_offset = + block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size; + const index_t row_offset_k = + block_table == nullptr + ? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache) + + (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride + : block_table[block_table_idx] * params.k_batch_stride + block_table_offset * params.k_row_stride + + (bidh / params.h_h_k_ratio) * params.k_head_stride; + const index_t row_offset_v = + block_table == nullptr + ? binfo.k_offset(params.v_batch_stride, params.v_row_stride, bidb_cache) + + (n_block_max - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride + : block_table[block_table_idx] * params.v_batch_stride + block_table_offset * params.v_row_stride + + (bidh / params.h_h_k_ratio) * params.v_head_stride; + + Tensor mQ = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.q_ptr) + binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)), + make_shape(binfo.actual_seqlen_q, params.h, params.d), + make_stride(params.q_row_stride, params.q_head_stride, _1{})); + Tensor gQ = + local_tile(mQ(_, bidh, _), Shape, Int>{}, make_coord(m_block, 0)); // (kBlockM, kHeadDim) + Tensor gK = make_tensor( + make_gmem_ptr(reinterpret_cast(params.k_ptr) + row_offset_k), + Shape, Int>{}, + make_stride(params.k_row_stride, _1{})); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = + // %p\n", params.k_ptr, row_offset_k, gK.data()); } + Tensor gV = make_tensor( + make_gmem_ptr(reinterpret_cast(params.v_ptr) + row_offset_v), + Shape, Int>{}, + make_stride(params.v_row_stride, _1{})); + + Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutQ{}); + Tensor sK = make_tensor(sQ.data() + size(sQ), typename Kernel_traits::SmemLayoutKV{}); + Tensor sV = make_tensor(sK.data() + size(sK), typename Kernel_traits::SmemLayoutKV{}); + Tensor sVt = make_tensor(sV.data(), typename Kernel_traits::SmemLayoutVtransposed{}); + Tensor sVtNoSwizzle = make_tensor(sV.data().get(), typename Kernel_traits::SmemLayoutVtransposedNoSwizzle{}); + + typename Kernel_traits::GmemTiledCopyQKV gmem_tiled_copy_QKV; + auto gmem_thr_copy_QKV = gmem_tiled_copy_QKV.get_thread_slice(tidx); + + Tensor tQgQ = gmem_thr_copy_QKV.partition_S(gQ); + Tensor tQsQ = gmem_thr_copy_QKV.partition_D(sQ); + Tensor tKgK = gmem_thr_copy_QKV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K) + Tensor tKsK = gmem_thr_copy_QKV.partition_D(sK); + Tensor tVgV = gmem_thr_copy_QKV.partition_S(gV); // (VCPY, VCPY_N, VCPY_K) + Tensor tVsV = gmem_thr_copy_QKV.partition_D(sV); + + typename Kernel_traits::TiledMma tiled_mma; + auto thr_mma = tiled_mma.get_thread_slice(tidx); + Tensor tSrQ = thr_mma.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K) + Tensor tSrK = thr_mma.partition_fragment_B(sK); // (MMA,MMA_N,MMA_K) + Tensor tOrVt = thr_mma.partition_fragment_B(sVtNoSwizzle); // (MMA, MMA_K,MMA_N) + + Tensor acc_o = partition_fragment_C(tiled_mma, Shape, Int>{}); // MMA, MMA_M, MMA_K + + // + // Copy Atom retiling + // + + auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx); + Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ); + + auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx); + Tensor tSsK = smem_thr_copy_K.partition_S(sK); + + auto smem_tiled_copy_V = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtomTransposed{}, tiled_mma); + auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(tidx); + Tensor tOsVt = smem_thr_copy_V.partition_S(sVt); + + // PREDICATES + // + + // // Allocate predicate tensors for m and n + // Tensor tQpQ = make_tensor(make_shape(size<1>(tQsQ), size<2>(tQsQ)), Stride<_1,_0>{}); + // Tensor tKVpKV = make_tensor(make_shape(size<1>(tKsK), size<2>(tKsK)), Stride<_1,_0>{}); + + // Construct identity layout for sQ and sK + Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k) + + // Repeat the partitioning with identity layouts + Tensor tQcQ = gmem_thr_copy_QKV.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) + Tensor tKVcKV = gmem_thr_copy_QKV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k) + + // Allocate predicate tensors for k + Tensor tQpQ = make_tensor(make_shape(size<2>(tQsQ))); + Tensor tKVpKV = make_tensor(make_shape(size<2>(tKsK))); + + // Set predicates for k bounds + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tQpQ); ++k) { + tQpQ(k) = get<1>(tQcQ(0, 0, k)) < params.d; + } +#pragma unroll + for (int k = 0; k < size(tKVpKV); ++k) { + tKVpKV(k) = get<1>(tKVcKV(0, 0, k)) < params.d; + } + } + + // Prologue + + // Copy from Knew to K, optionally apply rotary embedding. + typename Kernel_traits::GmemTiledCopyRotcossin gmem_tiled_copy_rotary; + auto gmem_thr_copy_rotary = gmem_tiled_copy_rotary.get_thread_slice(tidx); + typename Kernel_traits::GmemTiledCopyRotcossinCont gmem_tiled_copy_rotary_cont; + auto gmem_thr_copy_rotary_cont = gmem_tiled_copy_rotary_cont.get_thread_slice(tidx); + if constexpr (Append_KV) { + // Even if we have MQA / GQA, all threadblocks responsible for the same KV head are writing to + // gmem. Technically it's a race condition, but they all write the same content anyway, and it's safe. + // We want to do this so that all threadblocks can proceed right after they finish writing the KV cache. + const index_t row_offset_cossin = + ((n_block_max - 1) * kBlockN + (params.leftpad_k == nullptr ? 0 : params.leftpad_k[bidb])) * + (params.rotary_dim / 2); + Tensor gCos = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(params.rotary_dim / 2, _1{})); + Tensor gSin = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(params.rotary_dim / 2, _1{})); + Tensor gCosCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(params.rotary_dim / 2, _1{})); + Tensor gSinCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(params.rotary_dim / 2, _1{})); + Tensor tRgCos = gmem_thr_copy_rotary.partition_S(gCos); + Tensor tRgSin = gmem_thr_copy_rotary.partition_S(gSin); + Tensor tRgCosCont = gmem_thr_copy_rotary_cont.partition_S(gCosCont); + Tensor tRgSinCont = gmem_thr_copy_rotary_cont.partition_S(gSinCont); + // if (cute::thread(0, 0)) { printf("rotary_cos_ptr = %p, gCos.data() = %p, tRgCos.data() = %p, rotary_dim = %d\n", + // params.rotary_cos_ptr, gCos.data(), tRgCos.data(), params.rotary_dim); } if (cute::thread(8, 0)) { + // print_tensor(gCos); } if (cute::thread(0, 0)) { print_tensor(tRgCos); } + + // const index_t row_offset_knew = binfo.k_offset(params.knew_batch_stride, params.knew_row_stride, bidb) + const index_t row_offset_knew = bidb * params.knew_batch_stride + + ((n_block_max - 1) * kBlockN) * params.knew_row_stride + + (bidh / params.h_h_k_ratio) * params.knew_head_stride; + // const index_t row_offset_vnew = binfo.k_offset(params.vnew_batch_stride, params.vnew_row_stride, bidb) + const index_t row_offset_vnew = bidb * params.vnew_batch_stride + + ((n_block_max - 1) * kBlockN) * params.vnew_row_stride + + (bidh / params.h_h_k_ratio) * params.vnew_head_stride; + // Subtract seqlen_k_cache * row stride so that conceptually gK and gKnew "line up". When we access them, + // e.g. if gK has 128 rows and gKnew has 64 rows, we access gK[:128] and gKNew[128:128 + 64]. + // This maps to accessing the first 64 rows of knew_ptr. + Tensor gKnew = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.knew_ptr) + row_offset_knew - + binfo.seqlen_k_cache * params.knew_row_stride), + Shape, Int>{}, + make_stride(params.knew_row_stride, _1{})); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("knew_ptr = %p, row_offset_knew = %d, + // gKnew_ptr = %p\n", params.knew_ptr, row_offset_knew, gKnew.data()); } + Tensor gVnew = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.vnew_ptr) + row_offset_vnew - + binfo.seqlen_k_cache * params.vnew_row_stride), + Shape, Int>{}, + make_stride(params.vnew_row_stride, _1{})); + Tensor tKgKnew = gmem_thr_copy_QKV.partition_S(gKnew); // (KCPY, KCPY_N, KCPY_K) + Tensor tVgVnew = gmem_thr_copy_QKV.partition_S(gVnew); // (VCPY, VCPY_N, VCPY_K) + + const int n_block_copy_min = std::max(n_block_min, binfo.seqlen_k_cache / kBlockN); + auto tKgK_data = tKgK.data(); + auto tVgV_data = tVgV.data(); + for (int n_block = n_block_max - 1; n_block >= n_block_copy_min; n_block--) { + flash::copy_w_min_idx( + tVgVnew, + tVgV, + tKVcKV, + tKVpKV, + binfo.actual_seqlen_k - n_block * kBlockN, + binfo.seqlen_k_cache - n_block * kBlockN); + tVgVnew.data() = tVgVnew.data() + (-int(kBlockN * params.vnew_row_stride)); + if (params.rotary_dim == 0) { + flash::copy_w_min_idx( + tKgKnew, + tKgK, + tKVcKV, + tKVpKV, + binfo.actual_seqlen_k - n_block * kBlockN, + binfo.seqlen_k_cache - n_block * kBlockN); + } else { + if (params.is_rotary_interleaved) { + // Don't clear OOB_K because we're writing to global memory + flash::copy_rotary_interleaved( + tKgKnew, + tKgK, + tRgCos, + tRgSin, + tKVcKV, + binfo.actual_seqlen_k - n_block * kBlockN, + binfo.seqlen_k_cache - n_block * kBlockN, + params.d, + params.rotary_dim); + tRgCos.data() = tRgCos.data() + (-int(kBlockN * params.rotary_dim / 2)); + tRgSin.data() = tRgSin.data() + (-int(kBlockN * params.rotary_dim / 2)); + } else { + // Don't clear OOB_K because we're writing to global memory + flash::copy_rotary_contiguous( + tKgKnew, + tKgK, + tRgCosCont, + tRgSinCont, + tKVcKV, + binfo.actual_seqlen_k - n_block * kBlockN, + binfo.seqlen_k_cache - n_block * kBlockN, + params.d, + params.rotary_dim); + tRgCosCont.data() = tRgCosCont.data() + (-int(kBlockN * params.rotary_dim / 2)); + tRgSinCont.data() = tRgSinCont.data() + (-int(kBlockN * params.rotary_dim / 2)); + } + } + tKgKnew.data() = tKgKnew.data() + (-int(kBlockN * params.knew_row_stride)); + if (block_table == nullptr) { + tVgV.data() = tVgV.data() + (-int(kBlockN * params.v_row_stride)); + tKgK.data() = tKgK.data() + (-int(kBlockN * params.k_row_stride)); + } else { + if (n_block > n_block_copy_min) { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - 1) * kBlockN / params.page_block_size; + const int block_table_offset_next = (n_block - 1) * kBlockN - block_table_idx_next * params.page_block_size; + const int table_diff = block_table[block_table_idx_next] - block_table[block_table_idx_cur]; + const int offset_diff = block_table_offset_next - block_table_offset_cur; + tVgV.data() = tVgV.data() + table_diff * params.v_batch_stride + offset_diff * params.v_row_stride; + tKgK.data() = tKgK.data() + table_diff * params.k_batch_stride + offset_diff * params.k_row_stride; + } + } + } + // Need this before we can read in K again, so that we'll see the updated K values. + __syncthreads(); + tKgK.data() = tKgK_data; + tVgV.data() = tVgV_data; + } + + // Read Q from gmem to smem, optionally apply rotary embedding. + if (!Append_KV || params.rotary_dim == 0) { + // We don't need to clear the sQ smem tiles since we'll only write out the valid outputs + flash::copy( + gmem_tiled_copy_QKV, tQgQ, tQsQ, tQcQ, tQpQ, binfo.actual_seqlen_q - m_block * kBlockM); + } else { + const index_t row_offset_cossin = + (binfo.seqlen_k_cache + (params.leftpad_k == nullptr ? 0 : params.leftpad_k[bidb]) + + (Is_causal || Is_local ? m_block * kBlockM : 0)) * + (params.rotary_dim / 2); + // If not causal, all the queries get the same the cos/sin, taken at location seqlen_k_cache. + // We do this by setting the row stride of gCos / gSin to 0. + Tensor gCos = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gSin = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gCosCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gSinCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor tRgCos = gmem_thr_copy_rotary.partition_S(gCos); + Tensor tRgSin = gmem_thr_copy_rotary.partition_S(gSin); + Tensor tRgCosCont = gmem_thr_copy_rotary_cont.partition_S(gCosCont); + Tensor tRgSinCont = gmem_thr_copy_rotary_cont.partition_S(gSinCont); + if (params.is_rotary_interleaved) { + flash::copy_rotary_interleaved( + tQgQ, tQsQ, tRgCos, tRgSin, tQcQ, binfo.actual_seqlen_q - m_block * kBlockM, 0, params.d, params.rotary_dim); + } else { + flash::copy_rotary_contiguous( + tQgQ, + tQsQ, + tRgCosCont, + tRgSinCont, + tQcQ, + binfo.actual_seqlen_q - m_block * kBlockM, + 0, + params.d, + params.rotary_dim); + } + } + + int n_block = n_block_max - 1; + // We don't need to clear the sK smem tiles since we'll mask out the scores anyway. + flash::copy( + gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV, binfo.actual_seqlen_k - n_block * kBlockN); + cute::cp_async_fence(); + + // flash::cp_async_wait<0>(); + // __syncthreads(); + // if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tKsK); } + // __syncthreads(); + + clear(acc_o); + + flash::Softmax<2 * size<1>(acc_o)> softmax; + + const float alibi_slope = + !Has_alibi ? 0.0f + : reinterpret_cast(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / + params.scale_softmax; + flash::Mask mask( + binfo.actual_seqlen_k, + binfo.actual_seqlen_q, + params.window_size_left, + params.window_size_right, + alibi_slope, + params.m_block_dim); + + fwdIterator blockmask(params, binfo, kBlockM, kBlockN, bidb, bidh, m_block, n_block_min, n_block_max); + int next_block_idx = blockmask.max_no_larger(n_block_max - 1); + int leap = 0; + + // For performance reason, we separate out two kinds of iterations: + // those that need masking on S, and those that don't. + // We need masking on S for the very last block when K and V has length not multiple of kBlockN. + // We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks. + // We will have at least 1 "masking" iteration. + + // If not even_N, then seqlen_k might end in the middle of a block. In that case we need to + // mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1. + constexpr int n_masking_steps = + (!Is_causal && !Is_local) + ? 1 + : ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1); +#pragma unroll + for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) { + const bool skip = (n_block != next_block_idx); + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + // Advance gV + if (masking_step > 0) { + if (block_table == nullptr) { + tVgV.data() = tVgV.data() + (-int(kBlockN * params.v_row_stride)); + } else { + const int block_table_idx_cur = (n_block + 1) * kBlockN / params.page_block_size; + const int block_table_offset_cur = (n_block + 1) * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = n_block * kBlockN / params.page_block_size; + const int block_table_offset_next = n_block * kBlockN - block_table_idx_next * params.page_block_size; + tVgV.data() = tVgV.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.v_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.v_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tVgV, tVsV, tKVcKV, tKVpKV); + } else { + // Clear the smem tiles to account for predicated off loads + flash::copy( + gmem_tiled_copy_QKV, tVgV, tVsV, tKVcKV, tKVpKV, binfo.actual_seqlen_k - n_block * kBlockN); + } + cute::cp_async_fence(); + + if (!skip) { + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + // if (cute::thread0()) { print(acc_s); } + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + mask.template apply_mask( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + + next_block_idx = blockmask.max_no_larger(n_block - 1); + } else { + mask.all_mask(acc_s); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + // if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tVsV); } + // __syncthreads(); + + leap = (masking_step + 1 == n_masking_steps) ? n_block - next_block_idx : 1; + + if (n_block > n_block_min && next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + // We have key_padding_mask so we'll need to Check_inf + masking_step == 0 + ? softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2) + : softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2); + // if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); } + + if (!skip) { + // Convert acc_s from fp32 to fp16/bf16 + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs(rP.layout())); + + flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma, smem_tiled_copy_V, smem_thr_copy_V); + } + + // This check is at the end of the loop since we always have at least 1 iteration + if (n_masking_steps > 1 && n_block <= n_block_min) { + --n_block; + break; + } + } + + leap = n_block - next_block_idx + 1; + + // These are the iterations where we don't need masking on S + for (n_block = next_block_idx; n_block != -1 && n_block >= n_block_min; n_block = next_block_idx) { + next_block_idx = blockmask.max_no_larger(n_block - 1); + + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + // Advance gV + if (block_table == nullptr) { + tVgV.data() = tVgV.data() + (-int(kBlockN * leap * params.v_row_stride)); + } else { + const int block_table_idx_cur = (n_block + leap) * kBlockN / params.page_block_size; + const int block_table_offset_cur = (n_block + leap) * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = n_block * kBlockN / params.page_block_size; + const int block_table_offset_next = n_block * kBlockN - block_table_idx_next * params.page_block_size; + tVgV.data() = tVgV.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.v_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.v_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tVgV, tVsV, tKVcKV, tKVpKV); + cute::cp_async_fence(); + + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + + leap = n_block - next_block_idx; + if (next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + mask.template apply_mask( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + softmax.template softmax_rescale_o( + acc_s, acc_o, params.scale_softmax_log2); + + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + Tensor tOrP = make_tensor(rP.data(), flash::convert_layout_acc_Aregs(rP.layout())); + + flash::gemm_rs(acc_o, tOrP, tOrVt, tOsVt, tiled_mma, smem_tiled_copy_V, smem_thr_copy_V); + } + + // Epilogue + + Tensor lse = softmax.template normalize_softmax_lse(acc_o, params.scale_softmax); + // if (cute::thread0()) { print(lse); } + + Tensor sOaccum = make_tensor( + make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) + // Partition sO to match the accumulator partitioning + using SmemTiledCopyO = + std::conditional_t; + auto smem_tiled_copy_Oaccum = make_tiled_copy_C(SmemTiledCopyO{}, tiled_mma); + auto smem_thr_copy_Oaccum = smem_tiled_copy_Oaccum.get_thread_slice(tidx); + Tensor rO = flash::convert_type(acc_o); + Tensor taccOrOaccum = smem_thr_copy_Oaccum.retile_S(rO); // ((Atom,AtomNum), MMA_M, MMA_N) + Tensor taccOsOaccum = smem_thr_copy_Oaccum.partition_D(sOaccum); // ((Atom,AtomNum),PIPE_M,PIPE_N) + + // sOaccum is larger than sQ, so we need to syncthreads here + // TODO: allocate enough smem for sOaccum + if constexpr (Split) { + __syncthreads(); + } + + cute::copy(smem_tiled_copy_Oaccum, taccOrOaccum, taccOsOaccum); + + const index_t row_offset_o = binfo.q_offset(params.o_batch_stride, params.o_row_stride, bidb) + + m_block * kBlockM * params.o_row_stride + bidh * params.o_head_stride; + const index_t row_offset_oaccum = + (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM) * params.d_rounded; + const index_t row_offset_lseaccum = + (Split || !params.unpadded_lse ? ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + : bidh * params.total_q + binfo.q_offset(params.seqlen_q, 1, bidb)) + + m_block * kBlockM; + + Tensor gOaccum = make_tensor( + make_gmem_ptr( + reinterpret_cast(Split ? params.oaccum_ptr : params.o_ptr) + + (Split ? row_offset_oaccum : row_offset_o)), + Shape, Int>{}, + make_stride(Split ? kHeadDim : params.o_row_stride, _1{})); + Tensor gLSEaccum = make_tensor( + make_gmem_ptr( + reinterpret_cast(Split ? params.softmax_lseaccum_ptr : params.softmax_lse_ptr) + + row_offset_lseaccum), + Shape>{}, + Stride<_1>{}); + // if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); } + + GmemTiledCopyO gmem_tiled_copy_Oaccum; + auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx); + Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N) + Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + + __syncthreads(); + + Tensor tOrOaccum = make_tensor(shape(tOgOaccum)); + cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum); + + Tensor caccO = make_identity_tensor(Shape, Int>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k) + Tensor taccOcO = thr_mma.partition_C(caccO); // (MMA,MMA_M,MMA_K) + static_assert(decltype(size<0>(taccOcO))::value == 4); + // Convert to ((2, 2), MMA_M, MMA_K) then take only the row indices. + Tensor taccOcO_row = logical_divide(taccOcO, Shape<_2>{})(make_coord(0, _), _, 0); + CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M + if (get<1>(taccOcO_row(0)) == 0) { +#pragma unroll + for (int mi = 0; mi < size(lse); ++mi) { + const int row = get<0>(taccOcO_row(mi)); + if (row < binfo.actual_seqlen_q - m_block * kBlockM) { + gLSEaccum(row) = lse(mi); + } + } + } + + // Construct identity layout for sO + Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) + Tensor tOpO = make_tensor(make_shape(size<2>(tOgOaccum))); + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tOpO); ++k) { + tOpO(k) = get<1>(tOcO(0, 0, k)) < params.d; + } + } + // Clear_OOB_K must be false since we don't want to write zeros to gmem + flash::copy( + gmem_tiled_copy_Oaccum, tOrOaccum, tOgOaccum, tOcO, tOpO, binfo.actual_seqlen_q - m_block * kBlockM); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + typename Kernel_traits, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV, + typename Params> +inline __device__ void compute_attn_1rowblock_splitkv_stage1( + const Params& params, + const int bidb, + const int bidh, + const int m_block, + const int n_split_idx, + const int num_n_splits) { + using Element = typename Kernel_traits::Element; + using ElementAccum = typename Kernel_traits::ElementAccum; + using index_t = typename Kernel_traits::index_t; + + // Shared memory. + extern __shared__ char smem_[]; + + // The thread index. + const int tidx = threadIdx.x; + + constexpr int kBlockM = Kernel_traits::kBlockM; + constexpr int kBlockN = Kernel_traits::kBlockN; + constexpr int kHeadDim = Kernel_traits::kHeadDim; + constexpr int kNWarps = Kernel_traits::kNWarps; + + const BlockInfo binfo(params, bidb); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("Is_even_MN = %d, is_cumulativ = %d, + // seqlen_k_cache = %d, actual_seqlen_k = %d\n", Is_even_MN, params.is_seqlens_k_cumulative, binfo.seqlen_k_cache, + // binfo.actual_seqlen_k); } if (threadIdx.x == 0 && blockIdx.y == 1 && blockIdx.z == 0) { printf("params.knew_ptr = + // %p, seqlen_k_cache + seqlen_knew = %d\n", params.knew_ptr, binfo.seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : + // params.seqlen_knew)); } + if (m_block * kBlockM >= binfo.actual_seqlen_q) return; + + const int n_blocks_per_split = ((params.seqlen_k + kBlockN - 1) / kBlockN + num_n_splits - 1) / num_n_splits; + const int n_block_min = + !Is_local ? n_split_idx * n_blocks_per_split + : std::max( + n_split_idx * n_blocks_per_split, + (m_block * kBlockM + binfo.actual_seqlen_k - binfo.actual_seqlen_q - params.window_size_left) / + kBlockN); + int n_block_max = std::min(cute::ceil_div(binfo.actual_seqlen_k, kBlockN), (n_split_idx + 1) * n_blocks_per_split); + int n_block_max_c = std::min(cute::ceil_div(binfo.actual_seqlen_c, kBlockN), (n_split_idx + 1) * n_blocks_per_split); + // if (tidx == 0) { + // printf("n_block_max = %d, n_block_max_c = %d, binfo.actual_seqlen_k = %d, binfo.actual_seqlen_c = %d, kBlockN = + // %d, n_split_idx = %d, n_blocks_per_split = %d\n", n_block_max, n_block_max_c, binfo.actual_seqlen_k, + // binfo.actual_seqlen_c, kBlockN, n_split_idx, n_blocks_per_split); + // } + + int phase_1_stride = 64; + if (binfo.actual_seqlen_k == binfo.actual_seqlen_c) { + // 这种情况下根本就没做compressed lse + phase_1_stride = 16; + } + + if (Is_local || Is_causal) { + // 注意到存在q_len比k_len短的情况,避免截断原本需要的k + // k对应以16为步长移动的, 计算q的长度对应的k, 然后计算偏移 + // 用 len - 15 / 16 这样计算是因为 block_size = 32 + // (len - 32) / 16 + 1 + const int _ori_actual_seqlen_q = binfo.actual_seqlen_q / params.m_block_dim; + const int _k_stride = 16; + const int _k_actual_seqlen_q = (_ori_actual_seqlen_q - _k_stride + 1) / _k_stride; + const int _max_seqlen_k = binfo.actual_seqlen_k; + const int _offset_k = _max_seqlen_k - _k_actual_seqlen_q; + + // coarse kernel对应以64为步长移动的, 计算q的长度对应的k, 然后计算偏移 + const int _c_stride = phase_1_stride; + const int _c_actual_seqlen_q = (_ori_actual_seqlen_q - _c_stride + 1) / _c_stride; + // 分析可知上下界算出来的_max_seqlen_c是一样的 + const int _max_seqlen_c = ((_max_seqlen_k * 16 + 15) - phase_1_stride + 1) / phase_1_stride; + const int _offset_c = _max_seqlen_c - _c_actual_seqlen_q; + // 如果q原本和k等长,可以_k_actual_seqlen_q == _max_seqlen_k, _c_actual_seqlen_q == _max_seqlen_k / 4 + + const int max_q = (m_block + 1) * kBlockM / params.m_block_dim - 1; + const int max_k = (max_q - 16 + 1) / 16 + _offset_k; + n_block_max = std::min(n_block_max, cute::ceil_div(max_k, kBlockN)); + const int max_c = (max_q - phase_1_stride + 1) / phase_1_stride + _offset_c; + n_block_max_c = std::min(n_block_max_c, cute::ceil_div(max_c, kBlockN)); + + // if (cute::thread0()) { + // printf("binfo.actual_seqlen_c = %d, binfo.actual_seqlen_k = %d, _ori_actual_seqlen_q = %d, _k_actual_seqlen_q = + // %d, _max_seqlen_k = %d, _offset_k = %d, _c_actual_seqlen_q = %d, _max_seqlen_c = %d, _offset_c = %d, max_q = %d, + // max_k = %d, max_c = %d\n, n_block_max = %d, n_block_max_c = %d\n", binfo.actual_seqlen_c, binfo.actual_seqlen_k, + // _ori_actual_seqlen_q, _k_actual_seqlen_q, _max_seqlen_k, _offset_k, _c_actual_seqlen_q, _max_seqlen_c, _offset_c, + // max_q, max_k, max_c, n_block_max, n_block_max_c); //} + // flash::cp_async_wait<0>(); __syncthreads(); + } + if (n_block_min >= n_block_max_c) { // This also covers the case where n_block_max <= 0 + // We exit early and write 0 to gOaccum and -inf to gLSEaccum. + // Otherwise we might read OOB elements from gK and gV, + // or get wrong results when we combine gOaccum from different blocks. + return; + } + + // We iterate over the blocks in reverse order. This is because the last block is the only one + // that needs masking when we read K and V from global memory. Moreover, iterating in reverse + // might save us 1 register (we just need n_block instead of both n_block and n_block_max). + + // We move K and V to the last block. + const int bidb_cache = params.cache_batch_idx == nullptr ? bidb : params.cache_batch_idx[bidb]; + const int* block_table = + params.block_table == nullptr ? nullptr : params.block_table + bidb * params.block_table_batch_stride; + const int block_table_idx = block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN / params.page_block_size; + const int block_table_offset = + block_table == nullptr ? 0 : (n_block_max - 1) * kBlockN - block_table_idx * params.page_block_size; + const index_t row_offset_k = + block_table == nullptr + ? binfo.k_offset(params.k_batch_stride, params.k_row_stride, bidb_cache) + + (n_block_max - 1) * kBlockN * params.k_row_stride + (bidh / params.h_h_k_ratio) * params.k_head_stride + : block_table[block_table_idx] * params.k_batch_stride + block_table_offset * params.k_row_stride + + (bidh / params.h_h_k_ratio) * params.k_head_stride; + const index_t row_offset_v = + block_table == nullptr + ? binfo.v_offset(params.v_batch_stride, params.v_row_stride, bidb_cache) + + (n_block_max_c - 1) * kBlockN * params.v_row_stride + (bidh / params.h_h_k_ratio) * params.v_head_stride + : block_table[block_table_idx] * params.v_batch_stride + block_table_offset * params.v_row_stride + + (bidh / params.h_h_k_ratio) * params.v_head_stride; + + // const index_t row_offset_p = ((bidb * params.h + bidh) * params.seqlen_q_rounded/16 // TODO 16 is m_block_dim + // + m_block * kBlockM/16) * params.seqlen_k_rounded + (n_block_max - 1) * kBlockN; + // 获取当前 batch 在 total_query 中的起始位置 + const index_t query_offset_in_total = binfo.sum_s_q == -1 ? bidb * params.seqlen_q // 固定长度序列 + : binfo.sum_s_q; // 变长序列,已经是累积位置 + + // 计算 P 矩阵的偏移 + const index_t row_offset_p = (bidh * params.total_q / 16 // head 偏移 + + (query_offset_in_total + m_block * kBlockM) / 16) * + params.seqlen_k_rounded // query 偏移 + + (n_block_max - 1) * kBlockN; // key 偏移 + + Tensor mQ = make_tensor( + make_gmem_ptr( + reinterpret_cast(params.q_ptr) + binfo.q_offset(params.q_batch_stride, params.q_row_stride, bidb)), + make_shape(binfo.actual_seqlen_q, params.h, params.d), + make_stride(params.q_row_stride, params.q_head_stride, _1{})); + Tensor gQ = + local_tile(mQ(_, bidh, _), Shape, Int>{}, make_coord(m_block, 0)); // (kBlockM, kHeadDim) + Tensor gK = make_tensor( + make_gmem_ptr(reinterpret_cast(params.k_ptr) + row_offset_k), + Shape, Int>{}, + make_stride(params.k_row_stride, _1{})); + Tensor gC = make_tensor( + make_gmem_ptr(reinterpret_cast(params.v_ptr) + row_offset_v), + Shape, Int>{}, + make_stride(params.v_row_stride, _1{})); + // if (threadIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0) { printf("k_ptr = %p, row_offset_k = %d, gK_ptr = + // %p\n", params.k_ptr, row_offset_k, gK.data()); } + Tensor gP = make_tensor( + make_gmem_ptr(reinterpret_cast(params.p_ptr) + row_offset_p), + Shape, Int>{}, // TODO 16 is m_block_dim + make_stride(params.seqlen_k_rounded, _1{})); + + Tensor sQ = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutQ{}); + Tensor sK = make_tensor(sQ.data() + size(sQ), typename Kernel_traits::SmemLayoutKV{}); + + typename Kernel_traits::GmemTiledCopyQKV gmem_tiled_copy_QKV; + auto gmem_thr_copy_QKV = gmem_tiled_copy_QKV.get_thread_slice(tidx); + + Tensor tQgQ = gmem_thr_copy_QKV.partition_S(gQ); + Tensor tQsQ = gmem_thr_copy_QKV.partition_D(sQ); + Tensor tKgK = gmem_thr_copy_QKV.partition_S(gC); // (KCPY, KCPY_N, KCPY_K) + Tensor tKsK = gmem_thr_copy_QKV.partition_D(sK); + + typename Kernel_traits::TiledMma tiled_mma; + auto thr_mma = tiled_mma.get_thread_slice(tidx); + Tensor tSrQ = thr_mma.partition_fragment_A(sQ); // (MMA,MMA_M,MMA_K) + Tensor tSrK = thr_mma.partition_fragment_B(sK); // (MMA,MMA_N,MMA_K) + + Tensor acc_o = partition_fragment_C(tiled_mma, Shape, Int>{}); // MMA, MMA_M, MMA_K + + // + // Copy Atom retiling + // + + auto smem_tiled_copy_Q = make_tiled_copy_A(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(tidx); + Tensor tSsQ = smem_thr_copy_Q.partition_S(sQ); + + auto smem_tiled_copy_K = make_tiled_copy_B(typename Kernel_traits::SmemCopyAtom{}, tiled_mma); + auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(tidx); + Tensor tSsK = smem_thr_copy_K.partition_S(sK); + + // PREDICATES + // + + // // Allocate predicate tensors for m and n + // Tensor tQpQ = make_tensor(make_shape(size<1>(tQsQ), size<2>(tQsQ)), Stride<_1,_0>{}); + // Tensor tKVpKV = make_tensor(make_shape(size<1>(tKsK), size<2>(tKsK)), Stride<_1,_0>{}); + + // Construct identity layout for sQ and sK + Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK))); // (BLK_N,BLK_K) -> (blk_n,blk_k) + + // Repeat the partitioning with identity layouts + Tensor tQcQ = gmem_thr_copy_QKV.partition_S(cQ); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) + Tensor tKVcKV = gmem_thr_copy_QKV.partition_S(cKV); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k) + + // Allocate predicate tensors for k + Tensor tQpQ = make_tensor(make_shape(size<2>(tQsQ))); + Tensor tKVpKV = make_tensor(make_shape(size<2>(tKsK))); + + // Set predicates for k bounds + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tQpQ); ++k) { + tQpQ(k) = get<1>(tQcQ(0, 0, k)) < params.d; + } +#pragma unroll + for (int k = 0; k < size(tKVpKV); ++k) { + tKVpKV(k) = get<1>(tKVcKV(0, 0, k)) < params.d; + } + } + + // Prologue + + // Copy from Knew to K, optionally apply rotary embedding. + typename Kernel_traits::GmemTiledCopyRotcossin gmem_tiled_copy_rotary; + auto gmem_thr_copy_rotary = gmem_tiled_copy_rotary.get_thread_slice(tidx); + typename Kernel_traits::GmemTiledCopyRotcossinCont gmem_tiled_copy_rotary_cont; + auto gmem_thr_copy_rotary_cont = gmem_tiled_copy_rotary_cont.get_thread_slice(tidx); + + // Read Q from gmem to smem, optionally apply rotary embedding. + if (!Append_KV || params.rotary_dim == 0) { + // We don't need to clear the sQ smem tiles since we'll only write out the valid outputs + flash::copy( + gmem_tiled_copy_QKV, tQgQ, tQsQ, tQcQ, tQpQ, binfo.actual_seqlen_q - m_block * kBlockM); + } else { + const index_t row_offset_cossin = + (binfo.seqlen_k_cache + (params.leftpad_k == nullptr ? 0 : params.leftpad_k[bidb]) + + (Is_causal || Is_local ? m_block * kBlockM : 0)) * + (params.rotary_dim / 2); + // If not causal, all the queries get the same the cos/sin, taken at location seqlen_k_cache. + // We do this by setting the row stride of gCos / gSin to 0. + Tensor gCos = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gSin = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gCosCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_cos_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor gSinCont = make_tensor( + make_gmem_ptr(reinterpret_cast(params.rotary_sin_ptr) + row_offset_cossin), + Shape, Int>{}, + make_stride(Is_causal || Is_local ? params.rotary_dim / 2 : 0, _1{})); + Tensor tRgCos = gmem_thr_copy_rotary.partition_S(gCos); + Tensor tRgSin = gmem_thr_copy_rotary.partition_S(gSin); + Tensor tRgCosCont = gmem_thr_copy_rotary_cont.partition_S(gCosCont); + Tensor tRgSinCont = gmem_thr_copy_rotary_cont.partition_S(gSinCont); + if (params.is_rotary_interleaved) { + flash::copy_rotary_interleaved( + tQgQ, tQsQ, tRgCos, tRgSin, tQcQ, binfo.actual_seqlen_q - m_block * kBlockM, 0, params.d, params.rotary_dim); + } else { + flash::copy_rotary_contiguous( + tQgQ, + tQsQ, + tRgCosCont, + tRgSinCont, + tQcQ, + binfo.actual_seqlen_q - m_block * kBlockM, + 0, + params.d, + params.rotary_dim); + } + } + + int n_block = n_block_max_c - 1; + // We don't need to clear the sK smem tiles since we'll mask out the scores anyway. + flash::copy( + gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV, binfo.actual_seqlen_c - n_block * kBlockN); + cute::cp_async_fence(); + + flash::Softmax<2 * size<1>(acc_o)> softmax; + + const float alibi_slope = + !Has_alibi ? 0.0f + : reinterpret_cast(params.alibi_slopes_ptr)[bidb * params.alibi_slopes_batch_stride + bidh] / + params.scale_softmax; + flash::Mask mask( + binfo.actual_seqlen_c, + binfo.actual_seqlen_q, + params.window_size_left, + params.window_size_right, + alibi_slope, + params.m_block_dim); + + fwdIterator blockmask(params, binfo, kBlockM, kBlockN, bidb, bidh, m_block, n_block_min, n_block_max_c); + int next_block_idx = blockmask.max_no_larger(n_block_max_c - 1); + int leap = 0; + + // For performance reason, we separate out two kinds of iterations: + // those that need masking on S, and those that don't. + // We need masking on S for the very last block when K and V has length not multiple of kBlockN. + // We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks. + // We will have at least 1 "masking" iteration. + + // If not even_N, then seqlen_k might end in the middle of a block. In that case we need to + // mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1. + constexpr int n_masking_steps = + (!Is_causal && !Is_local) + ? 1 + : ((Is_even_MN && Is_causal) ? cute::ceil_div(kBlockM, kBlockN) : cute::ceil_div(kBlockM, kBlockN) + 1); +#pragma unroll + for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) { + const bool skip = (n_block != next_block_idx); + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + if (!skip) { + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + // if (cute::thread0()) { print(acc_s); } + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + mask.template apply_mask_stage1( + acc_s, + n_block * kBlockN, + m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, + kNWarps * 16, + phase_1_stride); + + next_block_idx = blockmask.max_no_larger(n_block - 1); + } else { + mask.all_mask(acc_s); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + // if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tVsV); } + // __syncthreads(); + + leap = (masking_step + 1 == n_masking_steps) ? n_block - next_block_idx : 1; + + if (n_block > n_block_min && next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + // We have key_padding_mask so we'll need to Check_inf + masking_step == 0 + ? softmax + .template softmax_rescale_simple( + acc_s, params.scale_softmax_log2) + : softmax + .template softmax_rescale_simple( + acc_s, params.scale_softmax_log2); + // if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); } + + if (!skip) { + // Convert acc_s from fp32 to fp16/bf16 + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + } + + // This check is at the end of the loop since we always have at least 1 iteration + if (n_masking_steps > 1 && n_block <= n_block_min) { + --n_block; + break; + } + } + + leap = n_block - next_block_idx + 1; + + // These are the iterations where we don't need masking on S + for (n_block = next_block_idx; n_block != -1 && n_block >= n_block_min; n_block = next_block_idx) { + next_block_idx = blockmask.max_no_larger(n_block - 1); + + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + + leap = n_block - next_block_idx; + if (next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + mask.template apply_mask_stage1( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16, phase_1_stride); + softmax.template softmax_rescale_simple( + acc_s, params.scale_softmax_log2); + + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + } + + // Epilogue + + softmax.get_row_sum(); + + { // second time + tKgK = gmem_thr_copy_QKV.partition_S(gK); // (KCPY, KCPY_N, KCPY_K) + + n_block = n_block_max - 1; + // We don't need to clear the sK smem tiles since we'll mask out the scores anyway. + flash::copy( + gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV, binfo.actual_seqlen_k - n_block * kBlockN); + cute::cp_async_fence(); + + flash::Mask mask( + binfo.actual_seqlen_k, + binfo.actual_seqlen_q, + params.window_size_left, + params.window_size_right, + alibi_slope, + params.m_block_dim); + fwdIterator blockmask(params, binfo, kBlockM, kBlockN, bidb, bidh, m_block, n_block_min, n_block_max); + + next_block_idx = blockmask.max_no_larger(n_block_max - 1); + leap = 0; + +// For performance reason, we separate out two kinds of iterations: +// those that need masking on S, and those that don't. +// We need masking on S for the very last block when K and V has length not multiple of kBlockN. +// We also need masking on S if it's causal, for the last ceil_div(kBlockM, kBlockN) blocks. +// We will have at least 1 "masking" iteration. + +// If not even_N, then seqlen_k might end in the middle of a block. In that case we need to +// mask 2 blocks (e.g. when kBlockM == kBlockN), not just 1. +#pragma unroll + for (int masking_step = 0; masking_step < n_masking_steps; ++masking_step, --n_block) { + const bool skip = (n_block != next_block_idx); + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + if (!skip) { + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + // if (cute::thread0()) { print(acc_s); } + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + mask.template apply_mask_stage1( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + + next_block_idx = blockmask.max_no_larger(n_block - 1); + } else { + mask.all_mask(acc_s); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + // if (tidx == 0 && blockIdx.y == 0 && blockIdx.z == 0) { print(tVsV); } + // __syncthreads(); + + leap = (masking_step + 1 == n_masking_steps) ? n_block - next_block_idx : 1; + + if (n_block > n_block_min && next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = + (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + // We have key_padding_mask so we'll need to Check_inf + softmax.template softmax_rescale_gt(acc_s, params.scale_softmax_log2); + // if (cute::thread0()) { print(scores_max); print(scores_sum); print(scores); } + + if (!skip) { + // Convert acc_s from fp32 to fp16/bf16 + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + + if (params.p_ptr != nullptr) { + hdim16_reduce(acc_s, gP, n_block * kBlockN, m_block * kBlockM, kNWarps * 16); + gP.data() = gP.data() + (-kBlockN); + } + } + + // This check is at the end of the loop since we always have at least 1 iteration + if (n_masking_steps > 1 && n_block <= n_block_min) { + --n_block; + break; + } + } + + leap = n_block - next_block_idx + 1; + + // These are the iterations where we don't need masking on S + for (n_block = next_block_idx; n_block != -1 && n_block >= n_block_min; n_block = next_block_idx) { + next_block_idx = blockmask.max_no_larger(n_block - 1); + + Tensor acc_s = partition_fragment_C(tiled_mma, Shape, Int>{}); // (MMA=4, MMA_M, MMA_N) + clear(acc_s); + flash::cp_async_wait<0>(); + __syncthreads(); + + flash::gemm( + acc_s, + tSrQ, + tSrK, + tSsQ, + tSsK, + tiled_mma, + smem_tiled_copy_Q, + smem_tiled_copy_K, + smem_thr_copy_Q, + smem_thr_copy_K); + if constexpr (Is_softcap) { + flash::apply_softcap(acc_s, params.softcap); + } + + flash::cp_async_wait<0>(); + __syncthreads(); + + leap = n_block - next_block_idx; + if (next_block_idx != -1) { + // Advance gK + if (block_table == nullptr) { + tKgK.data() = tKgK.data() + (-int(kBlockN * leap * params.k_row_stride)); + } else { + const int block_table_idx_cur = n_block * kBlockN / params.page_block_size; + const int block_table_offset_cur = n_block * kBlockN - block_table_idx_cur * params.page_block_size; + const int block_table_idx_next = (n_block - leap) * kBlockN / params.page_block_size; + const int block_table_offset_next = + (n_block - leap) * kBlockN - block_table_idx_next * params.page_block_size; + tKgK.data() = tKgK.data() + + (block_table[block_table_idx_next] - block_table[block_table_idx_cur]) * params.k_batch_stride + + (block_table_offset_next - block_table_offset_cur) * params.k_row_stride; + } + flash::copy(gmem_tiled_copy_QKV, tKgK, tKsK, tKVcKV, tKVpKV); + // This cp_async_fence needs to be in the if block, otherwise the synchronization + // isn't right and we get race conditions. + cute::cp_async_fence(); + } + + mask.template apply_mask_stage1( + acc_s, n_block * kBlockN, m_block * kBlockM + (tidx / 32) * 16 + (tidx % 32) / 4, kNWarps * 16); + softmax.template softmax_rescale_gt(acc_s, params.scale_softmax_log2); + + Tensor rP = flash::convert_type(acc_s); + // Reshape rP from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) + // if using m16n8k16 or (4, MMA_M, MMA_N) if using m16n8k8. + + if (params.p_ptr != nullptr) { + hdim16_reduce(acc_s, gP, n_block * kBlockN, m_block * kBlockM, kNWarps * 16); + gP.data() = gP.data() + (-kBlockN); + } + } + } +} +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + typename Kernel_traits, + bool Is_dropout, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Return_softmax, + typename Params> +inline __device__ void compute_attn(const Params& params) { + const int m_block = blockIdx.x; + // The block index for the batch. + const int bidb = blockIdx.y; + // The block index for the head. + const int bidh = blockIdx.z; + + // In the Philox RNG, we use the offset to store the batch, head, and the lane id + // (within a warp). We use the subsequence to store the location of the 16 x 32 blocks within + // the attention matrix. This way, as long as we have the batch, head, and the location of + // the 16 x 32 block within the attention matrix, we can generate the exact same dropout pattern. + + flash::compute_attn_1rowblock< + Kernel_traits, + Is_dropout, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Return_softmax>(params, bidb, bidh, m_block); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + typename Kernel_traits, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV, + typename Params> +inline __device__ void compute_attn_splitkv(const Params& params) { + const int m_block = blockIdx.x; + // The block index for the batch. + const int bidb = Split ? blockIdx.z / params.h : blockIdx.y; + // The block index for the head. + const int bidh = Split ? blockIdx.z - bidb * params.h : blockIdx.z; + const int n_split_idx = Split ? blockIdx.y : 0; + const int num_n_splits = Split ? gridDim.y : 1; + flash::compute_attn_1rowblock_splitkv< + Kernel_traits, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Split, + Append_KV>(params, bidb, bidh, m_block, n_split_idx, num_n_splits); +} + +template < + typename Kernel_traits, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV, + typename Params> +inline __device__ void compute_attn_splitkv_stage1(const Params& params) { + const int m_block = blockIdx.x; + // The block index for the batch. + const int bidb = Split ? blockIdx.z / params.h : blockIdx.y; + // The block index for the head. + const int bidh = Split ? blockIdx.z - bidb * params.h : blockIdx.z; + const int n_split_idx = Split ? blockIdx.y : 0; + const int num_n_splits = Split ? gridDim.y : 1; + flash::compute_attn_1rowblock_splitkv_stage1< + Kernel_traits, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Split, + Append_KV>(params, bidb, bidh, m_block, n_split_idx, num_n_splits); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +inline __device__ void combine_attn_seqk_parallel(const Params& params) { + using Element = typename Kernel_traits::Element; + using ElementAccum = typename Kernel_traits::ElementAccum; + using index_t = typename Kernel_traits::index_t; + constexpr int kMaxSplits = 1 << Log_max_splits; + constexpr int kHeadDim = Kernel_traits::kHeadDim; + constexpr int kNThreads = 128; + + static_assert(kMaxSplits <= 128, "kMaxSplits must be <= 128"); + static_assert(kBlockM == 4 || kBlockM == 8 || kBlockM == 16 || kBlockM == 32, "kBlockM must be 4, 8, 16 or 32"); + + // Shared memory. + // kBlockM + 1 instead of kBlockM to reduce bank conflicts. + __shared__ ElementAccum sLSE[kMaxSplits][kBlockM + 1]; + + // The thread and block index. + const int tidx = threadIdx.x; + const int bidx = blockIdx.x; + + const index_t lse_size = params.b * params.h * params.seqlen_q; + + const index_t row_offset_lse = bidx * kBlockM; + Tensor gLSEaccum = make_tensor( + make_gmem_ptr(reinterpret_cast(params.softmax_lseaccum_ptr) + row_offset_lse), + Shape, Int>{}, + make_stride(lse_size, _1{})); + + // LSE format is different depending on params.unpadded_lse and params.seqlenq_ngroups_swapped, see comment in + // get_lse_tile. This tensor's layout maps row_offset_lse to {bidb, bidh, q_offset}. + Tensor gLSE = make_tensor( + make_gmem_ptr(reinterpret_cast(params.softmax_lse_ptr) + row_offset_lse), + Shape>{}, + Stride<_1>{}); + + // This layout maps row_offset_lse to {bidh, q_offset, bidb} or {bidh, bidb, q_offset}. + Layout flat_layout = make_layout(lse_size); + Layout orig_layout = make_layout(make_shape(params.seqlen_q, params.h, params.b)); + auto transposed_stride = params.seqlenq_ngroups_swapped ? make_stride(params.b, params.seqlen_q * params.b, 1) + : make_stride(1, params.seqlen_q * params.b, params.seqlen_q); + Layout remapped_layout = make_layout(make_shape(params.seqlen_q, params.h, params.b), transposed_stride); + Layout final_layout = cute::composition(remapped_layout, cute::composition(orig_layout, flat_layout)); + + Tensor gLSE_unpadded = + make_tensor(make_gmem_ptr(reinterpret_cast(params.softmax_lse_ptr)), final_layout); + + constexpr int kNLsePerThread = (kMaxSplits * kBlockM + kNThreads - 1) / kNThreads; + + // Read the LSE values from gmem and store them in shared memory, then transpose them. + constexpr int kRowsPerLoadLSE = kNThreads / kBlockM; +#pragma unroll + for (int l = 0; l < kNLsePerThread; ++l) { + const int row = l * kRowsPerLoadLSE + tidx / kBlockM; + const int col = tidx % kBlockM; + ElementAccum lse = (row < params.num_splits && col < lse_size - bidx * kBlockM) ? gLSEaccum(row, col) : -INFINITY; + if (row < kMaxSplits) { + sLSE[row][col] = lse; + } + // if (bidx == 0 && tidx < 32) { printf("tidx = %d, row = %d, col = %d, lse = %f\n", tidx, row, col, lse); } + } + // if (bidx == 1 && tidx < 32) { printf("tidx = %d, row_offset_lse = %d, lse = %f\n", tidx, row_offset_lse, + // lse_accum(0)); } + __syncthreads(); + Tensor lse_accum = make_tensor(Shape>{}); + constexpr int kRowsPerLoadTranspose = std::min(kRowsPerLoadLSE, kMaxSplits); + // To make sure that kMaxSplits is within 1 warp: we decide how many elements within kMaxSplits + // each thread should hold. If kMaxSplits = 16, then each thread holds 2 elements (128 threads, + // kBlockM rows, so each time we load we can load 128 / kBlockM rows). + // constexpr int kThreadsPerSplit = kMaxSplits / kRowsPerLoadTranspose; + // static_assert(kThreadsPerSplit <= 32); + static_assert(kRowsPerLoadTranspose <= 32); + static_assert(kNLsePerThread * kRowsPerLoadTranspose <= kMaxSplits); +#pragma unroll + for (int l = 0; l < kNLsePerThread; ++l) { + const int row = l * kRowsPerLoadTranspose + tidx % kRowsPerLoadTranspose; + const int col = tidx / kRowsPerLoadTranspose; + lse_accum(l) = (row < kMaxSplits && col < kBlockM) ? sLSE[row][col] : -INFINITY; + // if (bidx == 0 && tidx < 32) { printf("tidx = %d, row = %d, col = %d, lse = %f\n", tidx, row, col, lse_accum(l)); + // } + } + + // Compute the logsumexp of the LSE along the split dimension. + ElementAccum lse_max = lse_accum(0); +#pragma unroll + for (int l = 1; l < kNLsePerThread; ++l) { + lse_max = max(lse_max, lse_accum(l)); + } + MaxOp max_op; + lse_max = Allreduce::run(lse_max, max_op); + lse_max = lse_max == -INFINITY ? 0.0f : lse_max; // In case all local LSEs are -inf + float lse_sum = expf(lse_accum(0) - lse_max); +#pragma unroll + for (int l = 1; l < kNLsePerThread; ++l) { + lse_sum += expf(lse_accum(l) - lse_max); + } + SumOp sum_op; + lse_sum = Allreduce::run(lse_sum, sum_op); + // For the case where all local lse == -INFINITY, we want to set lse_logsum to INFINITY. Otherwise + // lse_logsum is log(0.0) = -INFINITY and we get NaN when we do lse_accum(l) - lse_logsum. + ElementAccum lse_logsum = (lse_sum == 0.f || lse_sum != lse_sum) ? INFINITY : logf(lse_sum) + lse_max; + // if (bidx == 0 && tidx < 32) { printf("tidx = %d, lse = %f, lse_max = %f, lse_logsum = %f\n", tidx, lse_accum(0), + // lse_max, lse_logsum); } + if (tidx % kRowsPerLoadTranspose == 0 && tidx / kRowsPerLoadTranspose < kBlockM) { + if (params.unpadded_lse) { + const index_t lse_offset = row_offset_lse + tidx / kRowsPerLoadTranspose; + if (lse_offset < lse_size) { + gLSE_unpadded(lse_offset) = lse_logsum; + } + } else { + gLSE(tidx / kRowsPerLoadTranspose) = lse_logsum; + } + } +// Store the scales exp(lse - lse_logsum) in shared memory. +#pragma unroll + for (int l = 0; l < kNLsePerThread; ++l) { + const int row = l * kRowsPerLoadTranspose + tidx % kRowsPerLoadTranspose; + const int col = tidx / kRowsPerLoadTranspose; + if (row < params.num_splits && col < kBlockM) { + sLSE[row][col] = expf(lse_accum(l) - lse_logsum); + } + } + __syncthreads(); + + const index_t row_offset_oaccum = bidx * kBlockM * params.d_rounded; + Tensor gOaccum = make_tensor( + make_gmem_ptr(reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum), + Shape, Int>{}, + Stride, _1>{}); + constexpr int kBlockN = kNThreads / kBlockM; + using GmemLayoutAtomOaccum = Layout, Int>, Stride, _1>>; + using GmemTiledCopyOaccum = decltype(make_tiled_copy( + Copy_Atom, ElementAccum>{}, + GmemLayoutAtomOaccum{}, + Layout>{})); // Val layout, 4 vals per store + GmemTiledCopyOaccum gmem_tiled_copy_Oaccum; + auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx); + Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_S(gOaccum); + Tensor tOrO = make_tensor(shape(tOgOaccum)); + Tensor tOrOaccum = make_tensor(shape(tOgOaccum)); + clear(tOrO); + + // Predicates + Tensor cOaccum = make_identity_tensor(Shape, Int>{}); + // Repeat the partitioning with identity layouts + Tensor tOcOaccum = gmem_thr_copy_Oaccum.partition_S(cOaccum); + Tensor tOpOaccum = make_tensor(make_shape(size<2>(tOgOaccum))); + if (!Is_even_K) { +#pragma unroll + for (int k = 0; k < size(tOpOaccum); ++k) { + tOpOaccum(k) = get<1>(tOcOaccum(0, 0, k)) < params.d; + } + } + // Load Oaccum in then scale and accumulate to O + for (int split = 0; split < params.num_splits; ++split) { + flash::copy( + gmem_tiled_copy_Oaccum, + tOgOaccum, + tOrOaccum, + tOcOaccum, + tOpOaccum, + params.b * params.h * params.seqlen_q - bidx * kBlockM); +#pragma unroll + for (int m = 0; m < size<1>(tOrOaccum); ++m) { + int row = get<0>(tOcOaccum(0, m, 0)); + ElementAccum lse_scale = sLSE[split][row]; +#pragma unroll + for (int k = 0; k < size<2>(tOrOaccum); ++k) { +#pragma unroll + for (int i = 0; i < size<0>(tOrOaccum); ++i) { + tOrO(i, m, k) += lse_scale * tOrOaccum(i, m, k); + } + } + // if (cute::thread0()) { printf("lse_scale = %f, %f\n", sLSE[split][0], sLSE[split][1]); print(tOrOaccum); } + } + tOgOaccum.data() = tOgOaccum.data() + params.b * params.h * params.seqlen_q * params.d_rounded; + } + // if (cute::thread0()) { print_tensor(tOrO); } + + Tensor rO = flash::convert_type(tOrO); +// Write to gO +#pragma unroll + for (int m = 0; m < size<1>(rO); ++m) { + const int idx = bidx * kBlockM + get<0>(tOcOaccum(0, m, 0)); + if (idx < params.b * params.h * params.seqlen_q) { + const int batch_idx = idx / (params.h * params.seqlen_q); + const int head_idx = (idx - batch_idx * (params.h * params.seqlen_q)) / params.seqlen_q; + // The index to the rows of Q + const int row = idx - batch_idx * (params.h * params.seqlen_q) - head_idx * params.seqlen_q; + auto o_ptr = reinterpret_cast(params.o_ptr) + batch_idx * params.o_batch_stride + + head_idx * params.o_head_stride + row * params.o_row_stride; +#pragma unroll + for (int k = 0; k < size<2>(rO); ++k) { + if (Is_even_K || tOpOaccum(k)) { + const int col = get<1>(tOcOaccum(0, m, k)); + Tensor gO = make_tensor(make_gmem_ptr(o_ptr + col), Shape(rO))::value>>{}, Stride<_1>{}); + // TODO: Should check if this is using vectorized store, but it seems pretty fast + copy(rO(_, m, k), gO); + // if (bidx == 0 && tidx == 0) { printf("tidx = %d, idx = %d, batch_idx = %d, head_idx = %d, row = %d, col = + // %d\n", tidx, idx, batch_idx, head_idx, row, col); print(rO(_, m, k)); print(gO); } + // reinterpret_cast(o_ptr)[col / 4] = recast(rO)(0, m, k); + } + } + } + } +} + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_launch_template.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_launch_template.h new file mode 100644 index 000000000..4c6becc67 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_launch_template.h @@ -0,0 +1,335 @@ +/****************************************************************************** + * Copyright (c) 2023, Tri Dao. + ******************************************************************************/ + +#pragma once +#include // For C10_CUDA_CHECK and C10_CUDA_KERNEL_LAUNCH_CHECK + +#include "flash.h" +#include "flash_fwd_kernel.h" +#include "hardware_info.h" +#include "static_switch.h" + +// Determine if the architecture supports FLASH and define a macro to handle parameter modifiers +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 +#define ARCH_SUPPORTS_FLASH +#define KERNEL_PARAM_MODIFIER __grid_constant__ +#else +#define KERNEL_PARAM_MODIFIER +#endif + +// Define a macro for unsupported architecture handling to centralize the error message +#define FLASH_UNSUPPORTED_ARCH \ + printf("FATAL: FlashAttention requires building with sm version sm80-sm90, but was built for < 8.0!"); + +// Use a macro to clean up kernel definitions +#define DEFINE_FLASH_FORWARD_KERNEL(kernelName, ...) \ + template \ + __global__ void kernelName(KERNEL_PARAM_MODIFIER const Flash_fwd_params params) + +DEFINE_FLASH_FORWARD_KERNEL( + flash_fwd_kernel, + bool Is_dropout, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Return_softmax) { +#if defined(ARCH_SUPPORTS_FLASH) + static_assert(!(Is_causal && Is_local)); // Enforce constraints + flash::compute_attn< + Kernel_traits, + Is_dropout, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Return_softmax>(params); +#else + FLASH_UNSUPPORTED_ARCH +#endif +} + +DEFINE_FLASH_FORWARD_KERNEL( + flash_fwd_splitkv_kernel, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV) { +#if defined(ARCH_SUPPORTS_FLASH) + flash::compute_attn_splitkv< + Kernel_traits, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Split, + Append_KV>(params); +#else + FLASH_UNSUPPORTED_ARCH +#endif +} + +DEFINE_FLASH_FORWARD_KERNEL( + flash_fwd_splitkv_stage1_kernel, + bool Is_causal, + bool Is_local, + bool Has_alibi, + bool Is_even_MN, + bool Is_even_K, + bool Is_softcap, + bool Split, + bool Append_KV) { +#if defined(ARCH_SUPPORTS_FLASH) + flash::compute_attn_splitkv_stage1< + Kernel_traits, + Is_causal, + Is_local, + Has_alibi, + Is_even_MN, + Is_even_K, + Is_softcap, + Split, + Append_KV>(params); +#else + FLASH_UNSUPPORTED_ARCH +#endif +} + +DEFINE_FLASH_FORWARD_KERNEL(flash_fwd_splitkv_combine_kernel, int kBlockM, int Log_max_splits, bool Is_even_K) { + static_assert(Log_max_splits >= 1); + flash::combine_attn_seqk_parallel(params); +} + +template +void run_flash_fwd(Flash_fwd_params& params, cudaStream_t stream) { + constexpr size_t smem_size = Kernel_traits::kSmemSize; + // printf("smem_size = %d\n", smem_size); + + // Work-around for gcc 7. It doesn't like nested BOOL_SWITCH. + // https://github.com/kokkos/kokkos-kernels/issues/349 + // https://github.com/HazyResearch/flash-attention/issues/21 + + const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM; + dim3 grid(num_m_block, params.b, params.h); + const bool is_even_MN = params.cu_seqlens_q == nullptr && params.cu_seqlens_k == nullptr && + params.seqlen_k % Kernel_traits::kBlockN == 0 && + params.seqlen_q % Kernel_traits::kBlockM == 0; + const bool is_even_K = params.d == Kernel_traits::kHeadDim; + // const bool return_softmax = params.p_ptr != nullptr; + BOOL_SWITCH(is_even_MN, IsEvenMNConst, [&] { + EVENK_SWITCH(is_even_K, IsEvenKConst, [&] { + // LOCAL_SWITCH((params.window_size_left >= 0 || params.window_size_right >= 0) && !Is_causal, Is_local, [&] { + constexpr static bool Is_local = false; + { // TODO remove debug info + // BOOL_SWITCH(return_softmax, ReturnSoftmaxConst, [&] { + constexpr static bool ReturnSoftmaxConst = false; + { // TODO remove debug info + // ALIBI_SWITCH(params.alibi_slopes_ptr != nullptr, Has_alibi, [&] { + constexpr static bool Has_alibi = false; + { // TODO remove debug info + // SOFTCAP_SWITCH(params.softcap > 0.0, Is_softcap, [&] { + constexpr static bool Is_softcap = false; + { + // Will only return softmax if dropout, to reduce compilation time. + // If not IsEvenKConst, we also set IsEvenMNConst to false to reduce number of templates. + // If return_softmax, set IsEvenMNConst to false to reduce number of templates + // If head dim > 128, set IsEvenMNConst to false to reduce number of templates + // If Is_local, set Is_causal to false + auto kernel = &flash_fwd_kernel < Kernel_traits, Is_dropout && !Is_softcap, Is_causal, + Is_local && !Is_causal, Has_alibi, + IsEvenMNConst && IsEvenKConst && !Is_local && !ReturnSoftmaxConst && Kernel_traits::kHeadDim <= 128, + IsEvenKConst, Is_softcap, ReturnSoftmaxConst && Is_dropout && !Is_softcap > ; + // auto kernel = &flash_fwd_kernel; + // printf("IsEvenMNConst = %d, IsEvenKConst = %d, Is_local = %d, Is_causal = %d, ReturnSoftmaxConst = %d, + // Is_dropout = %d\n", int(IsEvenMNConst), int(IsEvenKConst), int(Is_local), int(Is_causal), + // int(ReturnSoftmaxConst), int(Is_dropout)); auto kernel = &flash_fwd_kernel; + if (smem_size >= 48 * 1024) { + C10_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size)); + } + // int ctas_per_sm; + // cudaError status_ = cudaOccupancyMaxActiveBlocksPerMultiprocessor( + // &ctas_per_sm, kernel, Kernel_traits::kNThreads, smem_size); + // printf("smem_size = %d, CTAs per SM = %d\n", int(smem_size), ctas_per_sm); + kernel<<>>(params); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + } + } + } + }); + }); +} + +template +void run_flash_splitkv_fwd(Flash_fwd_params& params, cudaStream_t stream) { + static_assert(!Kernel_traits::Is_Q_in_regs, "SplitKV implementation does not support Is_Q_in_regs"); + static_assert(!Kernel_traits::Share_Q_K_smem, "SplitKV implementation does not support Share_Q_K_smem"); + constexpr size_t smem_size = Kernel_traits::kSmemSize; + const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM; + dim3 grid( + num_m_block, + params.num_splits > 1 ? params.num_splits : params.b, + params.num_splits > 1 ? params.b * params.h : params.h); + const bool is_even_MN = params.cu_seqlens_q == nullptr && params.cu_seqlens_k == nullptr && + params.seqlen_k % Kernel_traits::kBlockN == 0 && + params.seqlen_q % Kernel_traits::kBlockM == 0; + const bool is_even_K = params.d == Kernel_traits::kHeadDim; + BOOL_SWITCH(is_even_MN, IsEvenMNConst, [&] { + EVENK_SWITCH(is_even_K, IsEvenKConst, [&] { + // LOCAL_SWITCH((params.window_size_left >= 0 || params.window_size_right >= 0) && !Is_causal, Is_local, [&] { + constexpr static bool Is_local = false; + { // TODO remove debug info + BOOL_SWITCH(params.num_splits > 1, Split, [&] { + BOOL_SWITCH(params.knew_ptr != nullptr, Append_KV, [&] { + // ALIBI_SWITCH(params.alibi_slopes_ptr != nullptr, Has_alibi, [&] { + constexpr static bool Has_alibi = false; + { // TODO remove debug info + // SOFTCAP_SWITCH(params.softcap > 0.0, Is_softcap, [&] { + constexpr static bool Is_softcap = false; + { // TODO remove debug info + // If Append_KV, then we must have seqlen_offsets, which means cu_seqlens_k != nullptr. + // If not IsEvenKConst, we also set IsEvenMNConst to false to reduce number of templates. + // If Is_local, set Is_causal to false + auto kernel = &flash_fwd_splitkv_kernel < Kernel_traits, Is_causal, Is_local && !Is_causal, Has_alibi, + IsEvenMNConst && !Append_KV && IsEvenKConst && !Is_local && Kernel_traits::kHeadDim <= 128, + IsEvenKConst, Is_softcap, Split, Append_KV > ; + // auto kernel = &flash_fwd_splitkv_kernel; + // auto kernel = &flash_fwd_splitkv_kernel; + if (smem_size >= 48 * 1024) { + C10_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size)); + } + kernel<<>>(params); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + } + }); + }); + } + }); + }); + if (params.num_splits > 1) { + // We want kBlockM to be as small as possible for more parallelism. + // With 128 threads we can load 512 elements at a time, so if headdim is divisible by 128, kBlockM = 4. + // If headdim is divisible by 64, then we set kBlockM = 8, etc. + constexpr static int kBlockM = + Kernel_traits::kHeadDim % 128 == 0 ? 4 : (Kernel_traits::kHeadDim % 64 == 0 ? 8 : 16); + dim3 grid_combine((params.b * params.h * params.seqlen_q + kBlockM - 1) / kBlockM); + EVENK_SWITCH(is_even_K, IsEvenKConst, [&] { + if (params.num_splits <= 2) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 4) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 8) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 16) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 32) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 64) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } else if (params.num_splits <= 128) { + flash_fwd_splitkv_combine_kernel + <<>>(params); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + }); + } +} + +template +void run_flash_splitkv_fwd_stage1(Flash_fwd_params& params, cudaStream_t stream) { + static_assert(!Kernel_traits::Is_Q_in_regs, "SplitKV implementation does not support Is_Q_in_regs"); + static_assert(!Kernel_traits::Share_Q_K_smem, "SplitKV implementation does not support Share_Q_K_smem"); + constexpr size_t smem_size = Kernel_traits::kSmemSize; + const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM; + dim3 grid( + num_m_block, + params.num_splits > 1 ? params.num_splits : params.b, + params.num_splits > 1 ? params.b * params.h : params.h); + const bool is_even_MN = params.cu_seqlens_q == nullptr && params.cu_seqlens_k == nullptr && + params.seqlen_k % Kernel_traits::kBlockN == 0 && + params.seqlen_q % Kernel_traits::kBlockM == 0; + const bool is_even_K = params.d == Kernel_traits::kHeadDim; + BOOL_SWITCH(is_even_MN, IsEvenMNConst, [&] { + EVENK_SWITCH(is_even_K, IsEvenKConst, [&] { + // LOCAL_SWITCH((params.window_size_left >= 0 || params.window_size_right >= 0) && !Is_causal, Is_local, [&] { + constexpr static bool Is_local = false; + { // TODO remove debug info + // BOOL_SWITCH(params.num_splits > 1, Split, [&] { + constexpr static bool Split = false; + { // TODO remove debug info + // BOOL_SWITCH(params.knew_ptr != nullptr, Append_KV, [&] { + constexpr static bool Append_KV = false; + { // TODO remove debug info + // ALIBI_SWITCH(params.alibi_slopes_ptr != nullptr, Has_alibi, [&] { + constexpr static bool Has_alibi = false; + { // TODO remove debug info + // SOFTCAP_SWITCH(params.softcap > 0.0, Is_softcap, [&] { + constexpr static bool Is_softcap = false; + { // TODO remove debug info + // If Append_KV, then we must have seqlen_offsets, which means cu_seqlens_k != nullptr. + // If not IsEvenKConst, we also set IsEvenMNConst to false to reduce number of templates. + // If Is_local, set Is_causal to false + auto kernel = &flash_fwd_splitkv_stage1_kernel < Kernel_traits, Is_causal, Is_local && !Is_causal, + Has_alibi, + IsEvenMNConst && !Append_KV && IsEvenKConst && !Is_local && Kernel_traits::kHeadDim <= 128, + IsEvenKConst, Is_softcap, Split, Append_KV > ; + if (smem_size >= 48 * 1024) { + C10_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size)); + } + kernel<<>>(params); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + } + } + } + } + }); + }); +} + +template +void run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream) { + if (params.blockmask == nullptr) { + constexpr static int kBlockM = 64; // Fixed for all head dimensions + // TD [2023-08-28]: nvcc segfaults for headdim 96 with block size 64 x 256, + // and for headdim 192 with block size 64 x 128. + // Also for headdim 160 with block size 64 x 128 after the rotary addition. + constexpr static int kBlockN = Headdim <= 64 ? 256 : (Headdim <= 128 ? 128 : 64); + if (params.m_block_dim == 1) { + run_flash_splitkv_fwd, Is_causal>( + params, stream); + } else { + run_flash_splitkv_fwd_stage1, Is_causal>( + params, stream); + } + } else if (params.cu_seqlens_q != nullptr) { + constexpr static int kBlockM = 16; + constexpr static int kBlockN = 64; + run_flash_splitkv_fwd, Is_causal>( + params, stream); + } else { + constexpr static int kBlockM = 64; + constexpr static int kBlockN = 64; + run_flash_splitkv_fwd, Is_causal>( + params, stream); + } +} diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_causal_sm80.cu b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_causal_sm80.cu new file mode 100644 index 000000000..2172074d3 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_causal_sm80.cu @@ -0,0 +1,8 @@ +// Copyright (c) 2024, Tri Dao. +// Splitting the different head dimensions to different files to speed up compilation. +// This file is auto-generated. See "generate_kernels.py" + +#include "flash_fwd_launch_template.h" + +template void +run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream); diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_sm80.cu b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_sm80.cu new file mode 100644 index 000000000..88fb11388 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim128_bf16_sm80.cu @@ -0,0 +1,8 @@ +// Copyright (c) 2024, Tri Dao. +// Splitting the different head dimensions to different files to speed up compilation. +// This file is auto-generated. See "generate_kernels.py" + +#include "flash_fwd_launch_template.h" + +template void +run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream); diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_causal_sm80.cu b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_causal_sm80.cu new file mode 100644 index 000000000..9f430154a --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_causal_sm80.cu @@ -0,0 +1,8 @@ +// Copyright (c) 2024, Tri Dao. +// Splitting the different head dimensions to different files to speed up compilation. +// This file is auto-generated. See "generate_kernels.py" + +#include "flash_fwd_launch_template.h" + +template void +run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream); diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_sm80.cu b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_sm80.cu new file mode 100644 index 000000000..d14a2b680 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/flash_fwd_split_hdim64_bf16_sm80.cu @@ -0,0 +1,8 @@ +// Copyright (c) 2024, Tri Dao. +// Splitting the different head dimensions to different files to speed up compilation. +// This file is auto-generated. See "generate_kernels.py" + +#include "flash_fwd_launch_template.h" + +template void +run_mha_fwd_splitkv_dispatch(Flash_fwd_params& params, cudaStream_t stream); diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/hardware_info.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/hardware_info.h new file mode 100644 index 000000000..3ee73103c --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/hardware_info.h @@ -0,0 +1,39 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include + +#if !defined(__CUDACC_RTC__) +#include "cuda_runtime.h" +#endif + +#define CHECK_CUDA(call) \ + do { \ + cudaError_t status_ = call; \ + if (status_ != cudaSuccess) { \ + fprintf(stderr, "CUDA error (%s:%d): %s\n", __FILE__, __LINE__, cudaGetErrorString(status_)); \ + exit(1); \ + } \ + } while (0) + +inline int get_current_device() { + int device; + CHECK_CUDA(cudaGetDevice(&device)); + return device; +} + +inline std::tuple get_compute_capability(int device) { + int capability_major, capability_minor; + CHECK_CUDA(cudaDeviceGetAttribute(&capability_major, cudaDevAttrComputeCapabilityMajor, device)); + CHECK_CUDA(cudaDeviceGetAttribute(&capability_minor, cudaDevAttrComputeCapabilityMinor, device)); + return {capability_major, capability_minor}; +} + +inline int get_num_sm(int device) { + int multiprocessor_count; + CHECK_CUDA(cudaDeviceGetAttribute(&multiprocessor_count, cudaDevAttrMultiProcessorCount, device)); + return multiprocessor_count; +} diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/kernel_traits.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/kernel_traits.h new file mode 100644 index 000000000..5925b3761 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/kernel_traits.h @@ -0,0 +1,157 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include + +#include "cute/tensor.hpp" +#include "cutlass/cutlass.h" +#include "cutlass/layout/layout.h" + +using namespace cute; + +template +struct Flash_kernel_traits { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 + using Element = elem_type; + static constexpr bool Has_cp_async = true; +#else + using Element = cutlass::half_t; + static constexpr bool Has_cp_async = false; +#endif + + using ElementAccum = float; + using index_t = int64_t; + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 + using MMA_Atom_Arch = std::conditional_t< + std::is_same_v, + MMA_Atom, + MMA_Atom>; +#else + using MMA_Atom_Arch = MMA_Atom; +#endif + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 750 + using SmemCopyAtom = Copy_Atom; + using SmemCopyAtomTransposed = Copy_Atom; +#else + using SmemCopyAtom = Copy_Atom; + using SmemCopyAtomTransposed = Copy_Atom; +#endif +}; + +// If Share_Q_K_smem is true, that forces Is_Q_in_regs to be true +template < + int kHeadDim_, + int kBlockM_, + int kBlockN_, + int kNWarps_, + bool Is_Q_in_regs_ = false, + bool Share_Q_K_smem_ = false, + typename elem_type = cutlass::half_t, + typename Base = Flash_kernel_traits> +struct Flash_fwd_kernel_traits : public Base { + using Element = typename Base::Element; + using ElementAccum = typename Base::ElementAccum; + using index_t = typename Base::index_t; + static constexpr bool Has_cp_async = Base::Has_cp_async; + using SmemCopyAtom = typename Base::SmemCopyAtom; + using SmemCopyAtomTransposed = typename Base::SmemCopyAtomTransposed; + + static constexpr bool Share_Q_K_smem = Share_Q_K_smem_; + static constexpr bool Is_Q_in_regs = Is_Q_in_regs_ || Share_Q_K_smem; + + // The number of threads. + static constexpr int kNWarps = kNWarps_; + static constexpr int kNThreads = kNWarps * 32; + + static constexpr int kBlockM = kBlockM_; + static constexpr int kBlockN = kBlockN_; + static constexpr int kHeadDim = kHeadDim_; + static_assert(kHeadDim % 32 == 0); + static constexpr int kBlockKSmem = kHeadDim % 64 == 0 ? 64 : 32; + static constexpr int kBlockKGmem = kHeadDim % 128 == 0 ? 128 : (kHeadDim % 64 == 0 ? 64 : 32); + static constexpr int kSwizzle = kBlockKSmem == 32 ? 2 : 3; + + using TiledMma = TiledMMA< + typename Base::MMA_Atom_Arch, + Layout, _1, _1>>, // 4x1x1 or 8x1x1 thread group + Tile, _16, _16>>; + + using SmemLayoutAtomQ = decltype(composition( + Swizzle{}, + // This has to be kBlockKSmem, using kHeadDim gives wrong results for d=128 + Layout>, Stride, _1>>{})); + using SmemLayoutQ = decltype(tile_to_shape(SmemLayoutAtomQ{}, Shape, Int>{})); + + using SmemLayoutKV = decltype(tile_to_shape(SmemLayoutAtomQ{}, Shape, Int>{})); + + // https://github.com/ColfaxResearch/cutlass-kernels/blob/a222587e6d59b93ba704853d3946fb686d8b8892/src/fmha/fmha_forward.cu#L434 + using SmemLayoutVtransposed = + decltype(composition(SmemLayoutKV{}, make_layout(Shape, Int>{}, GenRowMajor{}))); + using SmemLayoutVtransposedNoSwizzle = decltype(get_nonswizzle_portion(SmemLayoutVtransposed{})); + + using SmemLayoutAtomO = decltype(composition( + Swizzle{}, Layout, Int>, Stride, _1>>{})); + using SmemLayoutO = decltype(tile_to_shape(SmemLayoutAtomO{}, Shape, Int>{})); + using SmemCopyAtomO = Copy_Atom, Element>; + using SmemCopyAtomOaccum = Copy_Atom, ElementAccum>; + + static constexpr int kSmemQSize = size(SmemLayoutQ{}) * sizeof(Element); + static constexpr int kSmemKVSize = size(SmemLayoutKV{}) * 2 * sizeof(Element); + static constexpr int kSmemSize = Share_Q_K_smem ? std::max(kSmemQSize, kSmemKVSize) : kSmemQSize + kSmemKVSize; + + static constexpr int kGmemElemsPerLoad = sizeof(cute::uint128_t) / sizeof(Element); + static_assert(kHeadDim % kGmemElemsPerLoad == 0, "kHeadDim must be a multiple of kGmemElemsPerLoad"); + // Using kBlockKSmem here is 6-10% faster than kBlockKGmem for d=128 because of bank conflicts. + // For example, for d=128, smem is split into 2 "pages", each page takes care of columns + // 0-63 and 64-127. If we have 16 threads per row for gmem read, when we write to smem, + // thread 0 - 7 will write to the first page and thread 8 - 15 will write to the second page, + // to the same banks. + static constexpr int kGmemThreadsPerRow = kBlockKSmem / kGmemElemsPerLoad; + static_assert(kNThreads % kGmemThreadsPerRow == 0, "kNThreads must be a multiple of kGmemThreadsPerRow"); + using GmemLayoutAtom = + Layout, Int>, Stride, _1>>; + + // We use CACHEGLOBAL instead of CACHEALWAYS for both Q and K/V, since we won't be reading + // from the same address by the same threadblock. This is slightly faster. + using Gmem_copy_struct = std::conditional_t< + Has_cp_async, + SM80_CP_ASYNC_CACHEGLOBAL, + AutoVectorizingCopyWithAssumedAlignment<128>>; + using GmemTiledCopyQKV = decltype(make_tiled_copy( + Copy_Atom{}, + GmemLayoutAtom{}, + Layout>{})); // Val layout, 8 vals per read + using GmemTiledCopyO = decltype(make_tiled_copy( + Copy_Atom, Element>{}, + GmemLayoutAtom{}, + Layout>{})); // Val layout, 8 vals per store + + using GmemLayoutAtomOaccum = std::conditional_t< + kBlockKSmem == 32, + Layout< + Shape<_16, _8>, // Thread layout, 8 threads per row + Stride<_8, _1>>, + Layout< + Shape<_8, _16>, // Thread layout, 16 threads per row + Stride<_16, _1>>>; + using GmemTiledCopyOaccum = decltype(make_tiled_copy( + Copy_Atom, ElementAccum>{}, + GmemLayoutAtomOaccum{}, + Layout>{})); // Val layout, 4 vals per store + using GmemLayoutAtomRotcossin = GmemLayoutAtom; + using GmemTiledCopyRotcossin = decltype(make_tiled_copy( + Copy_Atom, Element>{}, + GmemLayoutAtomRotcossin{}, + Layout>{})); // Val layout, 4 vals per load + using GmemTiledCopyRotcossinCont = decltype(make_tiled_copy( + Copy_Atom, Element>{}, + GmemLayoutAtomRotcossin{}, + Layout>{})); // Val layout, 8 vals per load +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/mask.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/mask.h new file mode 100644 index 000000000..ec059b9ce --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/mask.h @@ -0,0 +1,365 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include + +namespace flash { + +using namespace cute; + +template +__forceinline__ __device__ void +apply_mask(Tensor& tensor, const int max_seqlen_k, const int col_idx_offset_ = 0) { + // tensor has shape (nrow=(2, MMA_M), ncol=(2, MMA_N)) + static_assert(Layout::rank == 2, "Only support 2D Tensor"); + const int lane_id = threadIdx.x % 32; + const int col_idx_offset = col_idx_offset_ + (lane_id % 4) * 2; +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; + if (col_idx >= max_seqlen_k) { +// Without the "make_coord" we get wrong results +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + tensor(mi, make_coord(j, nj)) = -INFINITY; + } + } + } + } +} + +template +__forceinline__ __device__ void apply_mask_local( + Tensor& tensor, + const int col_idx_offset_, + const int max_seqlen_k, + const int row_idx_offset, + const int max_seqlen_q, + const int warp_row_stride, + const int window_size_left, + const int window_size_right, + const int m_block_dim = 1) { + // tensor has shape (nrow=(2, MMA_M), ncol=(2, MMA_N)) + static_assert(Layout::rank == 2, "Only support 2D Tensor"); + const int lane_id = threadIdx.x % 32; + const int col_idx_offset = col_idx_offset_ + (lane_id % 4) * 2; +#pragma unroll + for (int mi = 0; mi < size<0, 1>(tensor); ++mi) { + const int row_idx_base = row_idx_offset + mi * warp_row_stride; +#pragma unroll + for (int i = 0; i < size<0, 0>(tensor); ++i) { + const int row_idx = row_idx_base + i * 8; + + // Apply m_block_dim scaling to get logical indices + const int orig_row_idx = row_idx / m_block_dim; + const int orig_max_seqlen_q = max_seqlen_q / m_block_dim; + + const int col_idx_limit_left = std::max(0, orig_row_idx + max_seqlen_k - orig_max_seqlen_q - window_size_left); + const int col_idx_limit_right = + std::min(max_seqlen_k, orig_row_idx + 1 + max_seqlen_k - orig_max_seqlen_q + window_size_right); +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; + if (col_idx >= col_idx_limit_right || (HasWSLeft && col_idx < col_idx_limit_left)) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + } + // if (cute::thread0()) { + // printf("mi = %d, i = %d, row_idx = %d, max_seqlen_k = %d\n", mi, i, row_idx, max_seqlen_k); + // print(tensor(make_coord(i, mi), _)); + // // print(tensor(_, j + nj * size<1, 0>(tensor))); + // } + } + } +} + +template +__forceinline__ __device__ void apply_mask_causal( + Tensor& tensor, + const int col_idx_offset_, + const int max_seqlen_k, + const int row_idx_offset, + const int max_seqlen_q, + const int warp_row_stride, + const int m_block_dim = 1) { + // Causal masking is equivalent to local masking with window_size_left = infinity and window_size_right = 0 + apply_mask_local( + tensor, col_idx_offset_, max_seqlen_k, row_idx_offset, max_seqlen_q, warp_row_stride, -1, 0, m_block_dim); +} + +template +__forceinline__ __device__ void apply_mask_causal_w_idx( + Tensor& tensor, + Tensor const& idx_rowcol, + const int col_idx_offset_, + const int max_seqlen_k, + const int row_idx_offset) { + // tensor has shape (nrow=(2, MMA_M), ncol=(2, MMA_N)) + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); + static_assert(Layout1::rank == 2, "Only support 2D Tensor"); + CUTE_STATIC_ASSERT_V(size<0>(tensor) == size<0>(idx_rowcol)); + CUTE_STATIC_ASSERT_V(size<1>(tensor) == size<1>(idx_rowcol)); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + const int col_idx_limit = std::min(max_seqlen_k, 1 + row_idx_offset + get<0>(idx_rowcol(mi, 0))); +#pragma unroll + for (int ni = 0; ni < size<1, 1>(tensor); ++ni) { + if (col_idx_offset_ + get<1>(idx_rowcol(0, ni)) >= col_idx_limit) { + tensor(mi, ni) = -INFINITY; + } + } + // if (cute::thread0()) { + // printf("ni = %d, j = %d, col_idx = %d, max_seqlen_k = %d\n", ni, j, col_idx, max_seqlen_k); + // print(tensor(_, make_coord(j, ni))); + // // print(tensor(_, j + ni * size<1, 0>(tensor))); + // } + } +} + +template +struct Mask { + const int max_seqlen_k, max_seqlen_q; + const int window_size_left, window_size_right; + const float alibi_slope; + const int m_block_dim; + + __forceinline__ __device__ Mask( + const int max_seqlen_k, + const int max_seqlen_q, + const int window_size_left, + const int window_size_right, + const float alibi_slope = 0.f, + const int m_block_dim = 1) + : max_seqlen_k(max_seqlen_k), + max_seqlen_q(max_seqlen_q), + window_size_left(window_size_left), + window_size_right(window_size_right), + alibi_slope(!Has_alibi ? 0.0 : alibi_slope), + m_block_dim(m_block_dim) {}; + + // Causal_mask: whether this particular iteration needs causal masking + template + __forceinline__ __device__ void apply_mask( + Tensor& tensor_, const int col_idx_offset_, const int row_idx_offset, const int warp_row_stride) { + static_assert(!(Causal_mask && Is_local), "Cannot be both causal and local"); + static_assert(Layout::rank == 3, "Only support 3D Tensor"); + static_assert(decltype(size<0>(tensor_))::value == 4, "First dimension must be 4"); + static constexpr bool Need_masking = Has_alibi || Causal_mask || Is_local || !Is_even_MN; + // if (cute::thread0()) { printf("Has_alibi = %d, Causal_mask=%d, Is_local=%d, Is_even_MN = %d, Need_masking = + // %d\n", Has_alibi, Causal_mask, Is_local, Is_even_MN, Need_masking); } + if constexpr (Need_masking) { + // Reshape tensor_ from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_rowcol(tensor_.layout())); + // Do we need both row and column indices, or just column incides? + static constexpr bool Col_idx_only = !(Has_alibi && !Is_causal) && !Is_local && !Causal_mask; + const int lane_id = threadIdx.x % 32; + const int col_idx_offset = col_idx_offset_ + (lane_id % 4) * 2; + if constexpr (Col_idx_only) { +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + // No causal, no local + if constexpr (Has_alibi) { + tensor(mi, make_coord(j, nj)) += alibi_slope * col_idx; + } + if constexpr (!Is_even_MN) { + if (col_idx >= max_seqlen_k) { + tensor(mi, make_coord(j, nj)) = -INFINITY; + } + } + } + } + } + } else { +#pragma unroll + for (int mi = 0; mi < size<0, 1>(tensor); ++mi) { + const int row_idx_base = row_idx_offset + mi * warp_row_stride; +#pragma unroll + for (int i = 0; i < size<0, 0>(tensor); ++i) { + const int row_idx = row_idx_base + i * 8; + + const int orig_row_idx = row_idx / this->m_block_dim; + const int orig_max_seqlen_q = max_seqlen_q / this->m_block_dim; + + const int col_idx_limit_left = + std::max(0, orig_row_idx + max_seqlen_k - orig_max_seqlen_q - window_size_left); + const int col_idx_limit_right = + std::min(max_seqlen_k, orig_row_idx + 1 + max_seqlen_k - orig_max_seqlen_q + window_size_right); +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; + if constexpr (Has_alibi) { + if constexpr (Is_causal) { + tensor(make_coord(i, mi), make_coord(j, nj)) += alibi_slope * col_idx; + } else { + tensor(make_coord(i, mi), make_coord(j, nj)) -= + alibi_slope * abs(orig_row_idx + max_seqlen_k - orig_max_seqlen_q - col_idx); + } + } + if constexpr (Causal_mask) { + if (col_idx >= col_idx_limit_right) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + if constexpr (Is_local) { + if (col_idx >= col_idx_limit_right || col_idx < col_idx_limit_left) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + if constexpr (!Causal_mask && !Is_local && !Is_even_MN) { + // Causal and Local already handles MN masking + if (col_idx >= max_seqlen_k) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + } + } + } + } + } + } + }; + + // Causal_mask: whether this particular iteration needs causal masking + template + __forceinline__ __device__ void apply_mask_stage1( + Tensor& tensor_, + const int col_idx_offset_, + const int row_idx_offset, + const int warp_row_stride, + const int stride = 16) { + static_assert(!(Causal_mask && Is_local), "Cannot be both causal and local"); + static_assert(Layout::rank == 3, "Only support 3D Tensor"); + static_assert(decltype(size<0>(tensor_))::value == 4, "First dimension must be 4"); + static constexpr bool Need_masking = Has_alibi || Causal_mask || Is_local || !Is_even_MN; + // if (cute::thread0()) { printf("Has_alibi = %d, Causal_mask=%d, Is_local=%d, Is_even_MN = %d, Need_masking = + // %d\n", Has_alibi, Causal_mask, Is_local, Is_even_MN, Need_masking); } + if constexpr (Need_masking) { + // Reshape tensor_ from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_rowcol(tensor_.layout())); + // Do we need both row and column indices, or just column incides? + static constexpr bool Col_idx_only = !(Has_alibi && !Is_causal) && !Is_local && !Causal_mask; + const int lane_id = threadIdx.x % 32; + const int col_idx_offset = col_idx_offset_ + (lane_id % 4) * 2; + if constexpr (Col_idx_only) { +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + // No causal, no local + if constexpr (Has_alibi) { + tensor(mi, make_coord(j, nj)) += alibi_slope * col_idx; + } + if constexpr (!Is_even_MN) { + if (col_idx >= max_seqlen_k) { + tensor(mi, make_coord(j, nj)) = -INFINITY; + } + } + } + } + } + } else { +#pragma unroll + for (int mi = 0; mi < size<0, 1>(tensor); ++mi) { + const int row_idx_base = row_idx_offset + mi * warp_row_stride; +#pragma unroll + for (int i = 0; i < size<0, 0>(tensor); ++i) { + const int row_idx = row_idx_base + i * 8; + + const int orig_row_idx = row_idx / this->m_block_dim; + const int orig_max_seqlen_q = max_seqlen_q / this->m_block_dim; + + // 计算压缩后的max_seqlen_q + const int compress_stride = stride; + const int compressed_max_seqlen_q = (orig_max_seqlen_q - compress_stride + 1) / compress_stride; + const int compressed_row_idx = (orig_row_idx - stride + 1) / stride; + int _max_seqlen_k = max_seqlen_k; // compressed lse的时候max_seqlen_k是压缩后的长度 + const int offset_row_idx = + std::max(0, (orig_row_idx + 1) / stride - 1 + _max_seqlen_k - compressed_max_seqlen_q); + + const int col_idx_limit_left = std::max(0, (orig_row_idx - stride + 1) / stride - window_size_left); + const int col_idx_limit_right = std::min(_max_seqlen_k, (offset_row_idx + window_size_right)); + // const int col_idx_limit_right = std::min(max_seqlen_k, (orig_row_idx - stride + 1) / stride + + // window_size_right); + + // if (cute::thread0()) { + // if (stride == 64) { + // printf("orig_row_idx = %d, orig_max_seqlen_q = %d, compressed_max_seqlen_q = %d, compressed_row_idx = %d, + // _max_seqlen_k = %d, offset_row_idx = %d, col_idx_limit_left = %d, col_idx_limit_right = %d\n", + // orig_row_idx, orig_max_seqlen_q, compressed_max_seqlen_q, compressed_row_idx, _max_seqlen_k, + // offset_row_idx, col_idx_limit_left, col_idx_limit_right); + // } + // } + // flash::cp_async_wait<0>(); __syncthreads(); + +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { + const int col_idx_base = col_idx_offset + nj * 8; +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { + const int col_idx = col_idx_base + j; + if constexpr (Causal_mask) { + if (col_idx >= col_idx_limit_right) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + if constexpr (Is_local) { + if (col_idx >= col_idx_limit_right || col_idx < col_idx_limit_left) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + if constexpr (!Causal_mask && !Is_local && !Is_even_MN) { + // Causal and Local already handles MN masking + if (col_idx >= max_seqlen_k) { + tensor(make_coord(i, mi), make_coord(j, nj)) = -INFINITY; + } + } + } + } + } + } + } + } + }; + + template + __forceinline__ __device__ void all_mask(Tensor& tensor_) { + static_assert(Layout::rank == 3, "Only support 3D Tensor"); + static_assert(decltype(size<0>(tensor_))::value == 4, "First dimension must be 4"); + Tensor tensor = make_tensor(tensor_.data(), flash::convert_layout_acc_rowcol(tensor_.layout())); +#pragma unroll + for (int nj = 0; nj < size<1, 1>(tensor); ++nj) { +#pragma unroll + for (int j = 0; j < size<1, 0>(tensor); ++j) { +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + tensor(mi, make_coord(j, nj)) = -INFINITY; + } + } + } + }; +}; + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox.cuh b/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox.cuh new file mode 100644 index 000000000..d52b8d29e --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox.cuh @@ -0,0 +1,49 @@ +// Pytorch also has an implementation of Philox RNG: +// https://github.com/pytorch/pytorch/blob/8ca3c881db3e3510fcb7725389f6a0633c9b992c/torch/csrc/jit/tensorexpr/cuda_random.h +#pragma once +// Philox CUDA. + +namespace flash { + +struct ull2 { + unsigned long long x; + unsigned long long y; +}; + +__forceinline__ __device__ uint2 mulhilo32(const unsigned int a, const unsigned int b) { + uint2* res; + unsigned long long tmp; + asm("mul.wide.u32 %0, %1, %2;\n\t" : "=l"(tmp) : "r"(a), "r"(b)); + res = (uint2*)(&tmp); + return *res; +} + +__forceinline__ __device__ uint4 philox_single_round(const uint4 ctr, const uint2 key) { + constexpr unsigned long kPhiloxSA = 0xD2511F53; + constexpr unsigned long kPhiloxSB = 0xCD9E8D57; + uint2 res0 = mulhilo32(kPhiloxSA, ctr.x); + uint2 res1 = mulhilo32(kPhiloxSB, ctr.z); + uint4 ret = {res1.y ^ ctr.y ^ key.x, res1.x, res0.y ^ ctr.w ^ key.y, res0.x}; + return ret; +} + +__forceinline__ __device__ uint4 +philox(unsigned long long seed, unsigned long long subsequence, unsigned long long offset) { + constexpr unsigned long kPhilox10A = 0x9E3779B9; + constexpr unsigned long kPhilox10B = 0xBB67AE85; + uint2 key = reinterpret_cast(seed); + uint4 counter; + ull2* tmp = reinterpret_cast(&counter); + tmp->x = offset; + tmp->y = subsequence; +#pragma unroll + for (int i = 0; i < 6; i++) { + counter = philox_single_round(counter, key); + key.x += (kPhilox10A); + key.y += (kPhilox10B); + } + uint4 output = philox_single_round(counter, key); + return output; +} + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox_unpack.cuh b/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox_unpack.cuh new file mode 100644 index 000000000..3a54f45cb --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/philox_unpack.cuh @@ -0,0 +1,4 @@ +// This is purely so that it works with torch 2.1. For torch 2.2+ we can include ATen/cuda/PhiloxUtils.cuh + +#pragma once +#include diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/rotary.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/rotary.h new file mode 100644 index 000000000..234f06398 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/rotary.h @@ -0,0 +1,175 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include + +#include "utils.h" + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +namespace flash { + +using namespace cute; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool Is_even_K = true, + bool Clear_OOB_K = true, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Engine2, + typename Layout2, + typename Engine3, + typename Layout3> +__forceinline__ __device__ void copy_rotary_interleaved( + Tensor const& S, + Tensor& D, + Tensor const& Cos, + Tensor const& Sin, + Tensor const& identity_MN, + const int max_MN, + const int min_MN, + const int dim, + const int rotary_dim) { + CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{}); + CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{}); + CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K + CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin)); // MMA_K + static_assert(decltype(size<0>(S))::value == decltype(size<0>(Cos))::value * 2); + static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32 + Tensor rCos = make_fragment_like(Cos); + Tensor rSin = make_fragment_like(Sin); + Tensor rS = make_fragment_like(S); +#pragma unroll + for (int m = 0; m < size<1>(S); ++m) { + if (get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN) { +#pragma unroll + for (int k = 0; k < size<2>(S); ++k) { + if (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim) { + cute::copy(S(_, m, k), rS(_, m, k)); + if (get<1>(identity_MN(0, 0, k)) < rotary_dim) { + cute::copy(Cos(_, m, k), rCos(_, m, k)); + cute::copy(Sin(_, m, k), rSin(_, m, k)); + Tensor S_fp32 = convert_type(rS(_, m, k)); + Tensor cos_fp32 = convert_type(rCos(_, m, k)); + Tensor sin_fp32 = convert_type(rSin(_, m, k)); +#pragma unroll + for (int i = 0; i < size<0>(rS) / 2; ++i) { + float real = S_fp32(2 * i) * cos_fp32(i) - S_fp32(2 * i + 1) * sin_fp32(i); + float imag = S_fp32(2 * i) * sin_fp32(i) + S_fp32(2 * i + 1) * cos_fp32(i); + S_fp32(2 * i) = real; + S_fp32(2 * i + 1) = imag; + } + // Idk but I need to copy for the convert_type to work + Tensor S_fp32_copy = make_fragment_like(S_fp32); + cute::copy(S_fp32, S_fp32_copy); + using T = typename Engine0::value_type; + Tensor S_og_type = convert_type(S_fp32_copy); + cute::copy(S_og_type, rS(_, m, k)); + } + cute::copy(rS(_, m, k), D(_, m, k)); + } else if (Clear_OOB_K) { + cute::clear(D(_, m, k)); + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool Is_even_K = true, + bool Clear_OOB_K = true, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Engine2, + typename Layout2, + typename Engine3, + typename Layout3> +__forceinline__ __device__ void copy_rotary_contiguous( + Tensor const& S, + Tensor& D, + Tensor const& Cos, + Tensor const& Sin, + Tensor const& identity_MN, + const int max_MN, + const int min_MN, + const int dim, + const int rotary_dim) { + CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{}); + CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{}); + CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Cos)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Cos)); // MMA_K + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(Sin)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(Sin)); // MMA_K + CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(Cos)); // MMA + CUTE_STATIC_ASSERT_V(size<0>(Cos) == size<0>(Sin)); + static_assert(decltype(size<0>(Cos))::value % 2 == 0); // Since we do fast conversion from fp16/bf16 to fp32 + Tensor rCos = make_fragment_like(Cos); + Tensor rSin = make_fragment_like(Sin); + Tensor rS = make_fragment_like(S); + Tensor rS_other = make_fragment_like(rS(_, 0, 0)); +#pragma unroll + for (int m = 0; m < size<1>(S); ++m) { + if (get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN) { +#pragma unroll + for (int k = 0; k < size<2>(S); ++k) { + if (Is_even_K || get<1>(identity_MN(0, 0, k)) < dim) { + cute::copy(S(_, m, k), rS(_, m, k)); + if (get<1>(identity_MN(0, 0, k)) < rotary_dim) { + const bool is_left = get<1>(identity_MN(0, 0, k)) < rotary_dim / 2; + Tensor gS_other = + make_tensor(S(_, m, k).data() + (is_left ? rotary_dim / 2 : -rotary_dim / 2), S(_, m, k).layout()); + cute::copy(gS_other, rS_other); + // if (cute::thread0()) { print_tensor(rS(_, m, k)); print_tensor(rS_other); } + Tensor gCos = make_tensor(Cos(_, m, k).data() + (is_left ? 0 : -rotary_dim / 2), Cos(_, m, k).layout()); + Tensor gSin = make_tensor(Sin(_, m, k).data() + (is_left ? 0 : -rotary_dim / 2), Sin(_, m, k).layout()); + cute::copy(gCos, rCos(_, m, k)); + cute::copy(gSin, rSin(_, m, k)); + // if (cute::thread0()) { print_tensor(rCos(_, m, k)); print_tensor(rSin(_, m, k)); } + Tensor S_fp32 = convert_type(rS(_, m, k)); + Tensor S_other_fp32 = convert_type(rS_other); + Tensor cos_fp32 = convert_type(rCos(_, m, k)); + Tensor sin_fp32 = convert_type(rSin(_, m, k)); +#pragma unroll + for (int i = 0; i < size<0>(rS); ++i) { + S_fp32(i) = S_fp32(i) * cos_fp32(i) + S_other_fp32(i) * (is_left ? -sin_fp32(i) : sin_fp32(i)); + } + // Idk but I need to copy for the convert_type to work + Tensor S_fp32_copy = make_fragment_like(S_fp32); + cute::copy(S_fp32, S_fp32_copy); + using T = typename Engine0::value_type; + Tensor S_og_type = convert_type(S_fp32_copy); + cute::copy(S_og_type, rS(_, m, k)); + // if (cute::thread0()) { print_tensor(rS(_, m, k)); } + } + cute::copy(rS(_, m, k), D(_, m, k)); + } else if (Clear_OOB_K) { + cute::clear(D(_, m, k)); + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/softmax.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/softmax.h new file mode 100644 index 000000000..b777ed5df --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/softmax.h @@ -0,0 +1,279 @@ +/****************************************************************************** + * Copyright (c) 2024, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include + +#include +#include + +#include "philox.cuh" +#include "utils.h" + +namespace flash { + +using namespace cute; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool zero_init = true, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Operator> +__device__ __forceinline__ void +thread_reduce_(Tensor const& tensor, Tensor& summary, Operator& op) { + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); + static_assert(Layout1::rank == 1, "Only support 1D Tensor"); + CUTE_STATIC_ASSERT_V(size<0>(summary) == size<0>(tensor)); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); mi++) { + summary(mi) = zero_init ? tensor(mi, 0) : op(summary(mi), tensor(mi, 0)); +#pragma unroll + for (int ni = 1; ni < size<1>(tensor); ni++) { + summary(mi) = op(summary(mi), tensor(mi, ni)); + } + } +} + +template +__device__ __forceinline__ void +quad_allreduce_(Tensor& dst, Tensor& src, Operator& op) { + CUTE_STATIC_ASSERT_V(size(dst) == size(src)); +#pragma unroll + for (int i = 0; i < size(dst); i++) { + dst(i) = Allreduce<4>::run(src(i), op); + } +} + +template < + bool zero_init = true, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Operator> +__device__ __forceinline__ void +reduce_(Tensor const& tensor, Tensor& summary, Operator& op) { + thread_reduce_(tensor, summary, op); + quad_allreduce_(summary, summary, op); +} + +template +__device__ __forceinline__ void reduce_max(Tensor const& tensor, Tensor& max) { + MaxOp max_op; + reduce_(tensor, max, max_op); +} + +template +__device__ __forceinline__ void reduce_sum(Tensor const& tensor, Tensor& sum) { + SumOp sum_op; + thread_reduce_(tensor, sum, sum_op); +} + +// Apply the exp to all the elements. +template +__forceinline__ __device__ void +scale_apply_exp2(Tensor& tensor, Tensor const& max, const float scale) { + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); + static_assert(Layout1::rank == 1, "Only support 1D Tensor"); + CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor)); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + // If max is -inf, then all elements must have been -inf (possibly due to masking). + // We don't want (-inf - (-inf)) since that would give NaN. + // If we don't have float around M_LOG2E the multiplication is done in fp64. + const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * (Scale_max ? scale : float(M_LOG2E)); +#pragma unroll + for (int ni = 0; ni < size<1>(tensor); ++ni) { +// Instead of computing exp(x - max), we compute exp2(x * log_2(e) - +// max * log_2(e)) This allows the compiler to use the ffma +// instruction instead of fadd and fmul separately. +// The following macro will disable the use of fma. +// See: https://github.com/pytorch/pytorch/issues/121558 for more details +// This macro is set in PyTorch and not FlashAttention +#ifdef UNFUSE_FMA + tensor(mi, ni) = exp2f(__fmul_rn(tensor(mi, ni), scale) - max_scaled); +#else + tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled); +#endif + } + } +} + +// Apply the exp to all the elements. +template +__forceinline__ __device__ void get_softmax( + Tensor& tensor, + Tensor const& max, + Tensor const& sum, + const float scale) { + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); + static_assert(Layout1::rank == 1, "Only support 1D Tensor"); + CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor)); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + // If max is -inf, then all elements must have been -inf (possibly due to masking). + // We don't want (-inf - (-inf)) since that would give NaN. + // If we don't have float around M_LOG2E the multiplication is done in fp64. + const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * (Scale_max ? scale : float(M_LOG2E)); + const float sum_scaled = 1. / sum(mi); +#pragma unroll + for (int ni = 0; ni < size<1>(tensor); ++ni) { +// Instead of computing exp(x - max), we compute exp2(x * log_2(e) - +// max * log_2(e)) This allows the compiler to use the ffma +// instruction instead of fadd and fmul separately. +// The following macro will disable the use of fma. +// See: https://github.com/pytorch/pytorch/issues/121558 for more details +// This macro is set in PyTorch and not FlashAttention +#ifdef UNFUSE_FMA + tensor(mi, ni) = exp2f(__fmul_rn(tensor(mi, ni), scale) - max_scaled) * sum_scaled; +#else + tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled) * sum_scaled; +#endif + } + } +} + +// Apply the exp to all the elements. +template +__forceinline__ __device__ void max_scale_exp2_sum( + Tensor& tensor, Tensor& max, Tensor& sum, const float scale) { + static_assert(Layout0::rank == 2, "Only support 2D Tensor"); + static_assert(Layout1::rank == 1, "Only support 1D Tensor"); + CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor)); +#pragma unroll + for (int mi = 0; mi < size<0>(tensor); ++mi) { + MaxOp max_op; + max(mi) = zero_init ? tensor(mi, 0) : max_op(max(mi), tensor(mi, 0)); +#pragma unroll + for (int ni = 1; ni < size<1>(tensor); ni++) { + max(mi) = max_op(max(mi), tensor(mi, ni)); + } + max(mi) = Allreduce<4>::run(max(mi), max_op); + // If max is -inf, then all elements must have been -inf (possibly due to masking). + // We don't want (-inf - (-inf)) since that would give NaN. + const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * scale; + sum(mi) = 0; +#pragma unroll + for (int ni = 0; ni < size<1>(tensor); ++ni) { + // Instead of computing exp(x - max), we compute exp2(x * log_2(e) - + // max * log_2(e)) This allows the compiler to use the ffma + // instruction instead of fadd and fmul separately. + tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled); + sum(mi) += tensor(mi, ni); + } + SumOp sum_op; + sum(mi) = Allreduce<4>::run(sum(mi), sum_op); + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct Softmax { + using TensorT = decltype(make_tensor(Shape>{})); + TensorT row_max, row_sum; + + __forceinline__ __device__ Softmax() {}; + + template + __forceinline__ __device__ void softmax_rescale_o(Tensor0& acc_s, Tensor1& acc_o, float softmax_scale_log2) { + // Reshape acc_s from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout())); + static_assert(decltype(size<0>(scores))::value == kNRows); + if (Is_first) { + flash::template reduce_max(scores, row_max); + flash::scale_apply_exp2(scores, row_max, softmax_scale_log2); + flash::reduce_sum(scores, row_sum); + } else { + Tensor scores_max_prev = make_fragment_like(row_max); + cute::copy(row_max, scores_max_prev); + flash::template reduce_max(scores, row_max); + // Reshape acc_o from (MMA=4, MMA_M, MMA_K) to (nrow=(2, MMA_M), ncol=(2, MMA_K)) + Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout())); + static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows); +#pragma unroll + for (int mi = 0; mi < size(row_max); ++mi) { + float scores_max_cur = !Check_inf ? row_max(mi) : (row_max(mi) == -INFINITY ? 0.0f : row_max(mi)); + float scores_scale = exp2f((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2); + row_sum(mi) *= scores_scale; +#pragma unroll + for (int ni = 0; ni < size<1>(acc_o_rowcol); ++ni) { + acc_o_rowcol(mi, ni) *= scores_scale; + } + } + flash::scale_apply_exp2(scores, row_max, softmax_scale_log2); + // We don't do the reduce across threads here since we don't need to use the row_sum. + // We do that reduce at the end when we need to normalize the softmax. + flash::reduce_sum(scores, row_sum); + } + }; + + template + __forceinline__ __device__ void softmax_rescale_simple(Tensor0& acc_s, float softmax_scale_log2) { + // Reshape acc_s from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout())); + static_assert(decltype(size<0>(scores))::value == kNRows); + if (Is_first) { + flash::template reduce_max(scores, row_max); + flash::scale_apply_exp2(scores, row_max, softmax_scale_log2); + flash::reduce_sum(scores, row_sum); + } else { + Tensor scores_max_prev = make_fragment_like(row_max); + cute::copy(row_max, scores_max_prev); + flash::template reduce_max(scores, row_max); +#pragma unroll + for (int mi = 0; mi < size(row_max); ++mi) { + float scores_max_cur = !Check_inf ? row_max(mi) : (row_max(mi) == -INFINITY ? 0.0f : row_max(mi)); + float scores_scale = exp2f((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2); + row_sum(mi) *= scores_scale; + } + flash::scale_apply_exp2(scores, row_max, softmax_scale_log2); + // We don't do the reduce across threads here since we don't need to use the row_sum. + // We do that reduce at the end when we need to normalize the softmax. + flash::reduce_sum(scores, row_sum); + } + }; + + template + __forceinline__ __device__ void softmax_rescale_gt(Tensor0& acc_s, float softmax_scale_log2) { + // Reshape acc_s from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) + Tensor scores = make_tensor(acc_s.data(), flash::convert_layout_acc_rowcol(acc_s.layout())); + static_assert(decltype(size<0>(scores))::value == kNRows); + flash::get_softmax(scores, row_max, row_sum, softmax_scale_log2); + }; + + __forceinline__ __device__ void get_row_sum() { + SumOp sum_op; + quad_allreduce_(row_sum, row_sum, sum_op); + } + + template + __forceinline__ __device__ TensorT + normalize_softmax_lse(Tensor0& acc_o, float softmax_scale, float rp_dropout = 1.0) { + SumOp sum_op; + quad_allreduce_(row_sum, row_sum, sum_op); + TensorT lse = make_fragment_like(row_sum); + Tensor acc_o_rowcol = make_tensor(acc_o.data(), flash::convert_layout_acc_rowcol(acc_o.layout())); + static_assert(decltype(size<0>(acc_o_rowcol))::value == kNRows); +#pragma unroll + for (int mi = 0; mi < size<0>(acc_o_rowcol); ++mi) { + float sum = row_sum(mi); + float inv_sum = (sum == 0.f || sum != sum) ? 1.f : 1.f / sum; + lse(mi) = (sum == 0.f || sum != sum) ? (Split ? -INFINITY : INFINITY) : row_max(mi) * softmax_scale + __logf(sum); + float scale = !Is_dropout ? inv_sum : inv_sum * rp_dropout; +#pragma unroll + for (int ni = 0; ni < size<1>(acc_o_rowcol); ++ni) { + acc_o_rowcol(mi, ni) *= scale; + } + } + return lse; + }; +}; + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/static_switch.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/static_switch.h new file mode 100644 index 000000000..f26fe49b6 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/static_switch.h @@ -0,0 +1,94 @@ +// Inspired by +// https://github.com/NVIDIA/DALI/blob/main/include/dali/core/static_switch.h +// and https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/Dispatch.h + +#pragma once + +/// @param COND - a boolean expression to switch by +/// @param CONST_NAME - a name given for the constexpr bool variable. +/// @param ... - code to execute for true and false +/// +/// Usage: +/// ``` +/// BOOL_SWITCH(flag, BoolConst, [&] { +/// some_function(...); +/// }); +/// ``` + +#define BOOL_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + if (COND) { \ + constexpr static bool CONST_NAME = true; \ + return __VA_ARGS__(); \ + } else { \ + constexpr static bool CONST_NAME = false; \ + return __VA_ARGS__(); \ + } \ + }() + +#ifdef FLASHATTENTION_DISABLE_DROPOUT +#define DROPOUT_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + constexpr static bool CONST_NAME = false; \ + return __VA_ARGS__(); \ + }() +#else +#define DROPOUT_SWITCH BOOL_SWITCH +#endif + +#ifdef FLASHATTENTION_DISABLE_ALIBI +#define ALIBI_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + constexpr static bool CONST_NAME = false; \ + return __VA_ARGS__(); \ + }() +#else +#define ALIBI_SWITCH BOOL_SWITCH +#endif + +#ifdef FLASHATTENTION_DISABLE_UNEVEN_K +#define EVENK_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + constexpr static bool CONST_NAME = true; \ + return __VA_ARGS__(); \ + }() +#else +#define EVENK_SWITCH BOOL_SWITCH +#endif + +#ifdef FLASHATTENTION_DISABLE_SOFTCAP +#define SOFTCAP_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + constexpr static bool CONST_NAME = false; \ + return __VA_ARGS__(); \ + }() +#else +#define SOFTCAP_SWITCH BOOL_SWITCH +#endif + +#ifdef FLASHATTENTION_DISABLE_LOCAL +#define LOCAL_SWITCH(COND, CONST_NAME, ...) \ + [&] { \ + constexpr static bool CONST_NAME = false; \ + return __VA_ARGS__(); \ + }() +#else +#define LOCAL_SWITCH BOOL_SWITCH +#endif + +#define FP16_SWITCH(COND, ...) \ + [&] { \ + using elem_type = cutlass::bfloat16_t; \ + return __VA_ARGS__(); \ + }() + +#define HEADDIM_SWITCH(HEADDIM, ...) \ + [&] { \ + if (HEADDIM == 64) { \ + constexpr static int kHeadDim = 64; \ + return __VA_ARGS__(); \ + } else { \ + constexpr static int kHeadDim = 128; \ + return __VA_ARGS__(); \ + } \ + }() diff --git a/sgl-kernel/csrc/infllm_v2/flash_attn/src/utils.h b/sgl-kernel/csrc/infllm_v2/flash_attn/src/utils.h new file mode 100644 index 000000000..7599eef50 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_attn/src/utils.h @@ -0,0 +1,483 @@ +/****************************************************************************** + * Copyright (c) 2023, Tri Dao. + ******************************************************************************/ + +#pragma once + +#include +#include +#include +#include + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 +#include +#endif + +#include +#include +#include +#include + +#include + +using namespace cute; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +namespace flash { + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__forceinline__ __device__ uint32_t relu2(const uint32_t x); + +template <> +__forceinline__ __device__ uint32_t relu2(const uint32_t x) { + uint32_t res; + const uint32_t zero = 0u; +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 + asm volatile("max.f16x2 %0, %1, %2;\n" : "=r"(res) : "r"(x), "r"(zero)); +#else + asm volatile( + "{\n" + "\t .reg .f16x2 sela;\n" + "\t set.gtu.u32.f16x2 sela, %1, %2;\n" + "\t and.b32 %0, sela, %1;\n" + "}\n" + : "=r"(res) + : "r"(x), "r"(zero)); +#endif + return res; +} + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 +template <> +__forceinline__ __device__ uint32_t relu2(const uint32_t x) { + uint32_t res; + const uint32_t zero = 0u; + asm volatile("max.bf16x2 %0, %1, %2;\n" : "=r"(res) : "r"(x), "r"(zero)); + return res; +} +#endif + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 + +template +__forceinline__ __device__ uint32_t convert_relu2(const float2 x); + +template <> +__forceinline__ __device__ uint32_t convert_relu2(const float2 x) { + uint32_t res; + const uint32_t a = reinterpret_cast(x.x); + const uint32_t b = reinterpret_cast(x.y); + asm volatile("cvt.rn.relu.f16x2.f32 %0, %1, %2;\n" : "=r"(res) : "r"(b), "r"(a)); + return res; +} + +template <> +__forceinline__ __device__ uint32_t convert_relu2(const float2 x) { + uint32_t res; + const uint32_t a = reinterpret_cast(x.x); + const uint32_t b = reinterpret_cast(x.y); + asm volatile("cvt.rn.relu.bf16x2.f32 %0, %1, %2;\n" : "=r"(res) : "r"(b), "r"(a)); + return res; +} + +#endif + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct MaxOp { + __device__ __forceinline__ T operator()(T const& x, T const& y) { + return x > y ? x : y; + } +}; + +template <> +struct MaxOp { + // This is slightly faster + __device__ __forceinline__ float operator()(float const& x, float const& y) { + return max(x, y); + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct SumOp { + __device__ __forceinline__ T operator()(T const& x, T const& y) { + return x + y; + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +struct Allreduce { + static_assert(THREADS == 32 || THREADS == 16 || THREADS == 8 || THREADS == 4); + template + static __device__ __forceinline__ T run(T x, Operator& op) { + constexpr int OFFSET = THREADS / 2; + x = op(x, __shfl_xor_sync(uint32_t(-1), x, OFFSET)); + return Allreduce::run(x, op); + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template <> +struct Allreduce<2> { + template + static __device__ __forceinline__ T run(T x, Operator& op) { + x = op(x, __shfl_xor_sync(uint32_t(-1), x, 1)); + return x; + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool A_in_regs = false, + bool B_in_regs = false, + typename Tensor0, + typename Tensor1, + typename Tensor2, + typename Tensor3, + typename Tensor4, + typename TiledMma, + typename TiledCopyA, + typename TiledCopyB, + typename ThrCopyA, + typename ThrCopyB> +__forceinline__ __device__ void gemm( + Tensor0& acc, + Tensor1& tCrA, + Tensor2& tCrB, + Tensor3 const& tCsA, + Tensor4 const& tCsB, + TiledMma tiled_mma, + TiledCopyA smem_tiled_copy_A, + TiledCopyB smem_tiled_copy_B, + ThrCopyA smem_thr_copy_A, + ThrCopyB smem_thr_copy_B) { + CUTE_STATIC_ASSERT_V(size<1>(tCrA) == size<1>(acc)); // MMA_M + CUTE_STATIC_ASSERT_V(size<1>(tCrB) == size<2>(acc)); // MMA_N + CUTE_STATIC_ASSERT_V(size<2>(tCrA) == size<2>(tCrB)); // MMA_K + Tensor tCrA_copy_view = smem_thr_copy_A.retile_D(tCrA); + CUTE_STATIC_ASSERT_V(size<1>(tCsA) == size<1>(tCrA_copy_view)); // M + Tensor tCrB_copy_view = smem_thr_copy_B.retile_D(tCrB); + CUTE_STATIC_ASSERT_V(size<1>(tCsB) == size<1>(tCrB_copy_view)); // N + if (!A_in_regs) { + cute::copy(smem_tiled_copy_A, tCsA(_, _, _0{}), tCrA_copy_view(_, _, _0{})); + } + if (!B_in_regs) { + cute::copy(smem_tiled_copy_B, tCsB(_, _, _0{}), tCrB_copy_view(_, _, _0{})); + } +#pragma unroll + for (int i = 0; i < size<2>(tCrA); ++i) { + if (i < size<2>(tCrA) - 1) { + if (!A_in_regs) { + cute::copy(smem_tiled_copy_A, tCsA(_, _, i + 1), tCrA_copy_view(_, _, i + 1)); + } + if (!B_in_regs) { + cute::copy(smem_tiled_copy_B, tCsB(_, _, i + 1), tCrB_copy_view(_, _, i + 1)); + } + } + cute::gemm(tiled_mma, tCrA(_, _, i), tCrB(_, _, i), acc); + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + typename Tensor0, + typename Tensor1, + typename Tensor2, + typename Tensor3, + typename TiledMma, + typename TiledCopy, + typename ThrCopy> +__forceinline__ __device__ void gemm_rs( + Tensor0& acc, + Tensor1& tCrA, + Tensor2& tCrB, + Tensor3 const& tCsB, + TiledMma tiled_mma, + TiledCopy smem_tiled_copy_B, + ThrCopy smem_thr_copy_B) { + CUTE_STATIC_ASSERT_V(size<1>(tCrA) == size<1>(acc)); // MMA_M + CUTE_STATIC_ASSERT_V(size<1>(tCrB) == size<2>(acc)); // MMA_N + CUTE_STATIC_ASSERT_V(size<2>(tCrA) == size<2>(tCrB)); // MMA_K + Tensor tCrB_copy_view = smem_thr_copy_B.retile_D(tCrB); + CUTE_STATIC_ASSERT_V(size<1>(tCsB) == size<1>(tCrB_copy_view)); // N + cute::copy(smem_tiled_copy_B, tCsB(_, _, _0{}), tCrB_copy_view(_, _, _0{})); +#pragma unroll + for (int i = 0; i < size<2>(tCrA); ++i) { + if (i < size<2>(tCrA) - 1) { + cute::copy(smem_tiled_copy_B, tCsB(_, _, i + 1), tCrB_copy_view(_, _, i + 1)); + } + cute::gemm(tiled_mma, tCrA(_, _, i), tCrB(_, _, i), acc); + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Convert acc_layout from (MMA=4, MMA_M, MMA_N) to (nrow=(2, MMA_M), ncol=(2, MMA_N)) +template +__forceinline__ __device__ auto convert_layout_acc_rowcol(Layout acc_layout) { + static_assert(decltype(size<0>(acc_layout))::value == 4); + static_assert(decltype(rank(acc_layout))::value == 3); + auto l = logical_divide(acc_layout, Shape<_2>{}); // ((2, 2), MMA_M, MMA_N) + return make_layout(make_layout(get<0, 1>(l), get<1>(l)), make_layout(get<0, 0>(l), get<2>(l))); +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Convert acc_layout from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) +// if using m16n8k16, or to (4, MMA_M, MMA_N) if using m16n8k8. +template +__forceinline__ __device__ auto convert_layout_acc_Aregs(Layout acc_layout) { + using X = Underscore; + static_assert(decltype(size<0>(acc_layout))::value == 4); + static_assert(decltype(rank(acc_layout))::value == 3); + constexpr int mma_shape_K = get<2>(typename MMA_traits::Shape_MNK{}); + static_assert(mma_shape_K == 8 || mma_shape_K == 16); + if constexpr (mma_shape_K == 8) { + return acc_layout; + } else { + auto l = logical_divide(acc_layout, Shape{}); // (4, MMA_M, (2, MMA_N / 2))) + return make_layout(make_layout(get<0>(l), get<2, 0>(l)), get<1>(l), get<2, 1>(l)); + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Convert acc_layout from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2) +template +__forceinline__ __device__ auto convert_layout_acc_dropout(Layout acc_layout) { + using X = Underscore; + static_assert(decltype(size<0>(acc_layout))::value == 4); + static_assert(decltype(rank(acc_layout))::value == 3); + auto l = logical_divide(acc_layout, Shape{}); // (4, MMA_M, (2, MMA_N / 2))) + return make_layout(make_layout(get<0>(l), get<2, 0>(l)), get<1>(l), get<2, 1>(l)); +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__forceinline__ __device__ auto convert_type(Tensor const& tensor) { + using From_type = typename Engine::value_type; + constexpr int numel = decltype(size(tensor))::value; + cutlass::NumericArrayConverter convert_op; + // HACK: this requires tensor to be "contiguous" + auto frag = convert_op(*reinterpret_cast*>(tensor.data())); + return make_tensor(make_rmem_ptr(&frag), tensor.layout()); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__forceinline__ __device__ void relu_(Tensor& tensor) { + constexpr int numel = decltype(size(tensor))::value; + static_assert(numel % 2 == 0); + using value_t = typename Engine::value_type; + // HACK: this requires tensor to be "contiguous" + Tensor tensor_uint32 = recast(tensor); +#pragma unroll + for (int i = 0; i < size(tensor_uint32); ++i) { + tensor_uint32(i) = relu2(tensor_uint32(i)); + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// On SM80 and above, we can fuse fp32 -> fp16/bf16 conversion and relu into 1 instruction +template +__forceinline__ __device__ auto convert_type_relu(Tensor const& tensor) { + using From_type = typename Engine::value_type; + static_assert(std::is_same_v || std::is_same_v); + static_assert(std::is_same_v); + constexpr int numel = decltype(size(tensor))::value; + static_assert(numel % 2 == 0); +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 + // HACK: this requires tensor to be "contiguous" + Tensor tensor_float2 = recast(tensor); + Tensor out_uint32 = make_tensor(tensor_float2.layout()); +#pragma unroll + for (int i = 0; i < size(out_uint32); ++i) { + out_uint32(i) = convert_relu2(tensor_float2(i)); + } + Tensor out = make_tensor(make_rmem_ptr(out_uint32.data()), tensor.layout()); +#else + Tensor out = flash::convert_type(tensor); + flash::relu_(out); +#endif + return out; +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Blocks until all but N previous cp.async.commit_group operations have committed. +// This differs from cute::cp_async_wait in that when N = 0 we don't call cp.async.wait_all +// (which is equivalent to commit_group then wait_group 0). +// Instead we just call cp.async.wait_group 0, which is slightly faster. +// https://github.com/NVIDIA/cutlass/blob/master/include/cute/arch/copy_sm80.hpp#L113 +template +CUTE_HOST_DEVICE void cp_async_wait() { +#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED) + asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); +#endif +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool Is_even_MN = true, + bool Is_even_K = true, + bool Clear_OOB_MN = false, + bool Clear_OOB_K = true, + typename TiledCopy, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Engine2, + typename Layout2, + typename Engine3, + typename Layout3> +__forceinline__ __device__ void copy( + TiledCopy tiled_copy, + Tensor const& S, + Tensor& D, + Tensor const& identity_MN, + Tensor const& predicate_K, + const int max_MN = 0) { + CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{}); + CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{}); + CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K + // There's no case where !Clear_OOB_K && Clear_OOB_MN + static_assert(!(Clear_OOB_MN && !Clear_OOB_K)); +#pragma unroll + for (int m = 0; m < size<1>(S); ++m) { + if (Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN) { +#pragma unroll + for (int k = 0; k < size<2>(S); ++k) { + if (Is_even_K || predicate_K(k)) { + cute::copy(tiled_copy, S(_, m, k), D(_, m, k)); + } else if (Clear_OOB_K) { + cute::clear(D(_, m, k)); + } + } + } else if (Clear_OOB_MN) { + cute::clear(D(_, m, _)); + } + } + // TD [2023-04-13]: Strange that the code below can cause race condition. + // I think it's because the copies are under an if statement. + // if (Is_even_K) { + // #pragma unroll + // for (int m = 0; m < size<1>(S); ++m) { + // if (Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN) { + // copy(tiled_copy, S(_, m, _), D(_, m, _)); + // } else if (Clear_OOB_MN) { + // clear(D(_, m, _)); + // } + // } + // } else { // It's slightly faster in this case if iterate over K first + // #pragma unroll + // for (int k = 0; k < size<2>(S); ++k) { + // if (predicate_K(k)) { + // #pragma unroll + // for (int m = 0; m < size<1>(S); ++m) { + // if (Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN) { + // copy(tiled_copy, S(_, m, k), D(_, m, k)); + // } else if (Clear_OOB_MN) { + // clear(D(_, m, k)); + // } + // } + // } else if (Clear_OOB_K) { // There's no case where !Clear_OOB_K && Clear_OOB_MN + // if (Clear_OOB_MN || Is_even_MN) { + // clear(D(_, _, k)); + // } else { + // #pragma unroll + // for (int m = 0; m < size<1>(S); ++m) { + // if (!(Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN)) { + // clear(D(_, m, k)); + // } + // } + // } + // } + // } + // } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template < + bool Is_even_K = true, + typename Engine0, + typename Layout0, + typename Engine1, + typename Layout1, + typename Engine2, + typename Layout2, + typename Engine3, + typename Layout3> +__forceinline__ __device__ void copy_w_min_idx( + Tensor const& S, + Tensor& D, + Tensor const& identity_MN, + Tensor const& predicate_K, + const int max_MN = 0, + const int min_MN = 0) { + CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{}); + CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{}); + CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA + CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M + CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K +// if (threadIdx.x == 0 && blockIdx.z == 0) { printf("blockIdx.y = %d, max_MN = %d, min_MN = %d\n", blockIdx.y, max_MN, +// min_MN); } +#pragma unroll + for (int m = 0; m < size<1>(S); ++m) { + // if (threadIdx.x == 0 && blockIdx.z == 0) { printf("blockIdx.y = %d, m = %d\n", blockIdx.y, get<0>(identity_MN(0, + // m, 0))); } + if (get<0>(identity_MN(0, m, 0)) >= min_MN && get<0>(identity_MN(0, m, 0)) < max_MN) { +// if (threadIdx.x == 0 && blockIdx.z == 0) { printf("Inner loop, blockIdx.y = %d, m = %d\n", blockIdx.y, +// get<0>(identity_MN(0, m, 0))); } +#pragma unroll + for (int k = 0; k < size<2>(S); ++k) { + if (Is_even_K || predicate_K(k)) { + cute::copy(S(_, m, k), D(_, m, k)); + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +template +__forceinline__ __device__ void apply_softcap(Tensor& tensor, const float softcap) { +#pragma unroll + for (int i = 0; i < size(tensor); ++i) { + tensor(i) = cutlass::fast_tanh(tensor(i) * softcap); + } +} + +template +__forceinline__ __device__ void +calculate_dtanh(Tensor& src_tensor, Tensor& dst_tensor, const float softcap) { +#pragma unroll + for (int i = 0; i < size(src_tensor); ++i) { + dst_tensor(i) = (1.f - (src_tensor(i) * src_tensor(i))) * softcap; + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace flash diff --git a/sgl-kernel/csrc/infllm_v2/flash_extension.cc b/sgl-kernel/csrc/infllm_v2/flash_extension.cc new file mode 100644 index 000000000..9a6403a69 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/flash_extension.cc @@ -0,0 +1,57 @@ +/* Copyright 2025 SGLang Team. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +// Pybind entry for the InfLLM-V2 FlashAttention backend (vendored from +// 3rdparty/infllmv2_cuda_impl). This builds as a standalone extension module +// `infllm_ops` so its `flash::` symbols stay isolated from sgl-kernel's own +// flash attention (`flash_ops` / `common_ops`). + +#include +#include +#include +#include + +#include + +// Forward declarations of the FlashAttention entry points implemented in +// flash_attn/flash_api.cpp. Signatures must match exactly. +std::vector mha_varlen_fwd_stage1( + at::Tensor& q, + const at::Tensor& k, + const at::Tensor& v, + c10::optional& out_, + const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, + const at::Tensor& cu_seqlens_v, + c10::optional& seqused_k, + c10::optional& leftpad_k_, + c10::optional& block_table_, + c10::optional& alibi_slopes_, + int max_seqlen_q, + const int max_seqlen_k, + const float p_dropout, + const float softmax_scale, + const bool zero_tensors, + bool is_causal, + int window_size_left, + int window_size_right, + const float softcap, + const bool return_softmax, + c10::optional gen_); + +PYBIND11_MODULE(infllm_ops, m) { + m.doc() = "InfLLM V2 FlashAttention backend (vendored into sgl-kernel)"; + m.def("varlen_fwd_stage1", &mha_varlen_fwd_stage1, "Forward pass (variable length) NSA stage 1"); +} diff --git a/sgl-kernel/csrc/infllm_v2/max_pooling.cu b/sgl-kernel/csrc/infllm_v2/max_pooling.cu new file mode 100644 index 000000000..21642c786 --- /dev/null +++ b/sgl-kernel/csrc/infllm_v2/max_pooling.cu @@ -0,0 +1,149 @@ +// InfLLM-V2 variable-length max pooling, AOT build. +// +// Migrated from `3rdparty/infllmv2_cuda_impl/csrc/max_pooling_1d.cuh`. The +// device kernels are kept faithful to the original; only the host-side +// launchers are rewritten from the raw `cudaStream_t` + `data_ptr` pybind +// interface to the sgl-kernel `at::Tensor` + torch.ops convention. +// +// Notes vs. the original implementation: +// * `TypeTraits::inf()` is replaced by `static_cast(INFINITY)`, and the +// pooling max is accumulated in fp32 so we don't rely on half/bf16 +// comparison operators. +// * Outputs are pre-allocated on the Python side and passed in (the original +// wrappers also allocated a zero-filled output before launching). + +#include +#include +#include +#include + +#include "utils.h" + +namespace { + +// input: [num_heads, total_q, max_seqlen_k] +// output: [num_heads, total_q, out_len] +template +__global__ void max_pooling_1d_varlen_kernel( + const T* input, + T* output, + const int* cu_seqlens_q, + const int* cu_seqlens_k, + const int* cache_lens, + int batch_size, + int num_heads, + int max_seqlen_k, + int out_len, + int kernel_size, + int stride, + int padding, + int block_size, + int local_blocks, + int init_blocks) { + const int bidh = blockIdx.y; // head index + const int bidq_global = blockIdx.x; // global query index across all batches + + int batch_idx = 0; + int q_start = 0, q_end = 0, k_start = 0, k_end = 0; + for (int b = 0; b < batch_size; b++) { + q_start = cu_seqlens_q[b]; + q_end = cu_seqlens_q[b + 1]; + k_start = cu_seqlens_k[b]; + k_end = cu_seqlens_k[b + 1]; + if (bidq_global >= q_start && bidq_global < q_end) { + batch_idx = b; + break; + } + } + + const int bidq_local = bidq_global - q_start; + const int seqlen_q = q_end - q_start; + const int seqlen_k = k_end - k_start; + if (bidq_local >= seqlen_q) return; + + const size_t total_q_all = static_cast(cu_seqlens_q[batch_size]); + const size_t in_offset = + static_cast(bidh) * total_q_all * max_seqlen_k + static_cast(bidq_global) * max_seqlen_k; + const T* in = input + in_offset; + const size_t out_offset = + static_cast(bidh) * total_q_all * out_len + static_cast(bidq_global) * out_len; + T* out = output + out_offset; + + const int cache_len = cache_lens[batch_idx]; + const int off_bq = (bidq_local + cache_len) / block_size; + const T pos_inf = static_cast(static_cast(INFINITY)); + + for (int k = threadIdx.x; k < out_len; k += blockDim.x) { + const int off_bk = k; + const bool should_mask_inf = (off_bk < init_blocks) || ((off_bq >= off_bk) && (off_bq <= off_bk + local_blocks)); + + if (should_mask_inf) { + out[k] = pos_inf; + } else { + int start = k * stride - padding; + int end = start + kernel_size; + start = max(start, 0); + end = min(end, seqlen_k); + + float max_val = -INFINITY; + for (int i = start; i < end; i++) { + const float v = static_cast(in[i]); + if (v > max_val) max_val = v; + } + out[k] = static_cast(max_val); + } + } +} + +} // namespace + +void infllm_v2_max_pooling_1d_varlen( + at::Tensor input, + at::Tensor output, + at::Tensor cu_seqlens_q, + at::Tensor cu_seqlens_k, + at::Tensor cache_lens, + int64_t max_seqlen_q, + int64_t max_seqlen_k, + int64_t kernel_size, + int64_t stride, + int64_t padding, + int64_t block_size, + int64_t local_blocks, + int64_t init_blocks, + int64_t total_q) { + TORCH_CHECK(input.dim() == 3, "input must be 3D [num_heads, total_q, max_k]"); + TORCH_CHECK(output.dim() == 3, "output must be 3D [num_heads, total_q, out_len]"); + TORCH_CHECK(cu_seqlens_q.scalar_type() == at::kInt, "cu_seqlens_q must be int32"); + TORCH_CHECK(cu_seqlens_k.scalar_type() == at::kInt, "cu_seqlens_k must be int32"); + TORCH_CHECK(cache_lens.scalar_type() == at::kInt, "cache_lens must be int32"); + + const int batch_size = static_cast(cu_seqlens_q.size(0)) - 1; + const int num_heads = static_cast(input.size(0)); + const int out_len = static_cast(output.size(2)); + const int grid_q = static_cast(total_q > 0 ? total_q : input.size(1)); + + const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + const dim3 grid(grid_q, num_heads); + const dim3 block(256); + + DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input.scalar_type(), c_type, [&] { + max_pooling_1d_varlen_kernel<<>>( + static_cast(input.data_ptr()), + static_cast(output.data_ptr()), + cu_seqlens_q.data_ptr(), + cu_seqlens_k.data_ptr(), + cache_lens.data_ptr(), + batch_size, + num_heads, + static_cast(max_seqlen_k), + out_len, + static_cast(kernel_size), + static_cast(stride), + static_cast(padding), + static_cast(block_size), + static_cast(local_blocks), + static_cast(init_blocks)); + return true; + }); +} diff --git a/sgl-kernel/include/sgl_kernel_ops.h b/sgl-kernel/include/sgl_kernel_ops.h index 2110aa573..efa46041e 100644 --- a/sgl-kernel/include/sgl_kernel_ops.h +++ b/sgl-kernel/include/sgl_kernel_ops.h @@ -112,6 +112,25 @@ int64_t cutlass_mla_get_workspace_size( int64_t sm_count = 0, int64_t num_kv_splits = 1 /* Set to 1 to avoid cuda_graph issue by default. */); +/* + * From csrc/infllm_v2 + */ +void infllm_v2_max_pooling_1d_varlen( + at::Tensor input, + at::Tensor output, + at::Tensor cu_seqlens_q, + at::Tensor cu_seqlens_k, + at::Tensor cache_lens, + int64_t max_seqlen_q, + int64_t max_seqlen_k, + int64_t kernel_size, + int64_t stride, + int64_t padding, + int64_t block_size, + int64_t local_blocks, + int64_t init_blocks, + int64_t total_q); + /* * From csrc/elementwise */ diff --git a/sgl-kernel/python/sgl_kernel/__init__.py b/sgl-kernel/python/sgl_kernel/__init__.py index e179d889c..201d253c5 100644 --- a/sgl-kernel/python/sgl_kernel/__init__.py +++ b/sgl-kernel/python/sgl_kernel/__init__.py @@ -71,6 +71,10 @@ else: shuffle_rows, ) from sgl_kernel.grammar import apply_token_bitmask_inplace_cuda + from sgl_kernel.infllm_v2 import ( + infllmv2_attn_stage1, + max_pooling_1d_varlen, + ) from sgl_kernel.kvcacheio import ( transfer_kv_all_layer, transfer_kv_all_layer_mla, diff --git a/sgl-kernel/python/sgl_kernel/infllm_v2/__init__.py b/sgl-kernel/python/sgl_kernel/infllm_v2/__init__.py new file mode 100644 index 000000000..40189d38f --- /dev/null +++ b/sgl-kernel/python/sgl_kernel/infllm_v2/__init__.py @@ -0,0 +1,7 @@ +from sgl_kernel.infllm_v2.attention import infllmv2_attn_stage1 +from sgl_kernel.infllm_v2.max_pooling import max_pooling_1d_varlen + +__all__ = [ + "infllmv2_attn_stage1", + "max_pooling_1d_varlen", +] diff --git a/sgl-kernel/python/sgl_kernel/infllm_v2/_loader.py b/sgl-kernel/python/sgl_kernel/infllm_v2/_loader.py new file mode 100644 index 000000000..b1f680c6c --- /dev/null +++ b/sgl-kernel/python/sgl_kernel/infllm_v2/_loader.py @@ -0,0 +1,98 @@ +"""Robust loader for the standalone ``infllm_ops`` pybind extension. + +The InfLLM-V2 FlashAttention backend is built as its own module ``infllm_ops`` +(installed into the ``sgl_kernel`` package directory). Under editable installs +the compiled ``.so`` may live in ``site-packages/sgl_kernel`` while the imported +``sgl_kernel`` package resolves to the source tree, so a plain ``from sgl_kernel +import infllm_ops`` is not always sufficient. This loader searches the known +candidate locations and loads the extension by file path. +""" + +import glob +import importlib.util +import site +import sys +from pathlib import Path +from typing import List, Optional + +_infllm_ops = None + + +def _candidate_dirs() -> List[Path]: + dirs: List[Path] = [] + + # 1) The directory of the sgl_kernel package as currently imported. + try: + import sgl_kernel + + dirs.append(Path(sgl_kernel.__file__).parent) + except Exception: + pass + + # 2) This module's parent package directory (source tree). + dirs.append(Path(__file__).resolve().parent.parent) + + # 3) Every ``sgl_kernel`` directory found on the install paths. + search_roots: List[str] = [] + try: + search_roots.extend(site.getsitepackages()) + except Exception: + pass + try: + search_roots.append(site.getusersitepackages()) + except Exception: + pass + search_roots.extend(p for p in sys.path if p) + for root in search_roots: + dirs.append(Path(root) / "sgl_kernel") + + # De-duplicate while preserving order. + seen = set() + unique: List[Path] = [] + for d in dirs: + key = str(d) + if key not in seen: + seen.add(key) + unique.append(d) + return unique + + +def _find_so() -> Optional[Path]: + for d in _candidate_dirs(): + if not d.is_dir(): + continue + matches = sorted(glob.glob(str(d / "infllm_ops*.so"))) + if matches: + return Path(matches[0]) + return None + + +def load_infllm_ops(): + """Import and return the ``infllm_ops`` extension module (cached).""" + global _infllm_ops + if _infllm_ops is not None: + return _infllm_ops + + # Fast path: a normal import may already work. + try: + from sgl_kernel import infllm_ops as _mod # type: ignore + + _infllm_ops = _mod + return _infllm_ops + except Exception: + pass + + so_path = _find_so() + if so_path is None: + raise ImportError( + "[sgl_kernel] Could not locate the 'infllm_ops' extension (infllm_ops*.so). " + "Ensure sgl-kernel was built with the InfLLM-V2 FlashAttention backend." + ) + + spec = importlib.util.spec_from_file_location("infllm_ops", str(so_path)) + if spec is None or spec.loader is None: + raise ImportError(f"[sgl_kernel] Could not create module spec for {so_path}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + _infllm_ops = module + return _infllm_ops diff --git a/sgl-kernel/python/sgl_kernel/infllm_v2/attention.py b/sgl-kernel/python/sgl_kernel/infllm_v2/attention.py new file mode 100644 index 000000000..514e0d7fb --- /dev/null +++ b/sgl-kernel/python/sgl_kernel/infllm_v2/attention.py @@ -0,0 +1,82 @@ +"""InfLLM-V2 sparse FlashAttention public API. + +Ported (drop-in) from ``3rdparty/infllmv2_cuda_impl/infllm_v2/infllmv2_sparse_attention.py``. +The CUDA backend now lives in the standalone ``infllm_ops`` extension. +""" + +import torch +from sgl_kernel.infllm_v2._loader import load_infllm_ops + + +def maybe_contiguous(x): + return x.contiguous() if x is not None and x.stride(-1) != 1 else x + + +def infllmv2_attn_stage1( + q, + k, + v, + cu_seqlens_q, + cu_seqlens_k, + cu_seqlens_v, + max_seqlen_q, + max_seqlen_k, + dropout_p=0.0, + softmax_scale=None, + causal=False, + window_size=(-1, -1), + softcap=0.0, + alibi_slopes=None, + deterministic=False, + return_attn_probs=True, + block_table=None, +): + """Neighborhood Sparse Attention (NSA) Stage 1 with varlen support. + + Drop-in replacement for ``infllm_v2.infllmv2_attn_stage1``. Returns the + attention-score matrix with the NSA sparsity pattern, shape + ``(num_heads_k, total_q, max_seqlen_k)``. + """ + infllm_ops = load_infllm_ops() + if softmax_scale is None: + softmax_scale = q.shape[-1] ** (-0.5) + + q, k, v = [maybe_contiguous(x) for x in (q, k, v)] + + total_q, nheads, head_dim = q.shape + nheads_k = k.shape[1] + nheads_per_group = nheads // nheads_k + + q = q.reshape(total_q, nheads_k, nheads_per_group, head_dim) + q = ( + q.transpose(1, 2) + .reshape(total_q * nheads_per_group, nheads_k, head_dim) + .contiguous() + ) + + result = infllm_ops.varlen_fwd_stage1( + q, + k, + v, + None, + cu_seqlens_q, + cu_seqlens_k, + cu_seqlens_v, + None, + None, + block_table, + alibi_slopes, + max_seqlen_q, + max_seqlen_k, + dropout_p, + softmax_scale, + True, + causal, + window_size[0], + window_size[1], + softcap, + True, + None, + ) + + return result[0] diff --git a/sgl-kernel/python/sgl_kernel/infllm_v2/max_pooling.py b/sgl-kernel/python/sgl_kernel/infllm_v2/max_pooling.py new file mode 100644 index 000000000..39ea31682 --- /dev/null +++ b/sgl-kernel/python/sgl_kernel/infllm_v2/max_pooling.py @@ -0,0 +1,61 @@ +import torch + + +def max_pooling_1d_varlen( + input: torch.Tensor, # num_heads x total_q x max_k + cu_seqlens_q: torch.Tensor, # batch_size + 1 + cu_seqlens_k: torch.Tensor, # batch_size + 1 + cache_lens: torch.Tensor, # batch_size + max_seqlen_q: int, + max_context_len: int, + local_blocks: int, + init_blocks: int, + block_size: int = 64, + stride: int = 16, + total_q: int = -1, +) -> torch.Tensor: + """Variable-length 1D max pooling over packed sequences. + + Drop-in replacement for ``infllm_v2.max_pooling_1d_varlen``. + """ + assert input.dtype in (torch.float16, torch.bfloat16) + assert cu_seqlens_q.dtype == torch.int32 + assert cu_seqlens_k.dtype == torch.int32 + assert cache_lens.dtype == torch.int32 + assert input.dim() == 3, f"Expected 3D input, got {input.dim()}D" + + input = input.contiguous() + cu_seqlens_q = cu_seqlens_q.contiguous() + cu_seqlens_k = cu_seqlens_k.contiguous() + cache_lens = cache_lens.contiguous() + + max_seqlen_k = max_context_len // stride + out_len = (max_context_len + block_size - 1) // block_size + + stride = block_size // stride + kernel_size = stride + 1 + padding = 1 + + num_heads = input.shape[0] + total_q = input.shape[1] + + output = torch.zeros( + num_heads, total_q, out_len, device=input.device, dtype=input.dtype + ) + torch.ops.sgl_kernel.infllm_v2_max_pooling_1d_varlen.default( + input, + output, + cu_seqlens_q, + cu_seqlens_k, + cache_lens, + max_seqlen_q, + max_seqlen_k, + kernel_size, + stride, + padding, + block_size, + local_blocks, + init_blocks, + total_q, + ) + return output diff --git a/sgl-kernel/tests/test_infllm_v2_attention.py b/sgl-kernel/tests/test_infllm_v2_attention.py new file mode 100644 index 000000000..72510fbe5 --- /dev/null +++ b/sgl-kernel/tests/test_infllm_v2_attention.py @@ -0,0 +1,61 @@ +"""Equivalence tests for the migrated InfLLM-V2 FlashAttention API. + +These compare the ``sgl_kernel.infllm_v2`` implementations against the original +``infllm_v2`` package (3rdparty/infllmv2_cuda_impl). Both call the same CUDA +kernels, so outputs are expected to match closely. The whole module is skipped +if the reference ``infllm_v2`` package is not importable. +""" + +import pytest +import torch + +sgl = pytest.importorskip("sgl_kernel.infllm_v2") +ref = pytest.importorskip("infllm_v2") + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), reason="CUDA is required for InfLLM-V2 kernels" +) + + +def _assert_close(a, b, name): + a = a.float() + b = b.float() + assert a.shape == b.shape, f"{name}: shape mismatch {a.shape} vs {b.shape}" + max_diff = (a - b).abs().max().item() + assert torch.allclose(a, b, atol=1e-2, rtol=1e-2), f"{name}: max diff {max_diff}" + + +@pytest.mark.parametrize("head_dim", [64, 128]) +@pytest.mark.parametrize("causal", [False, True]) +@pytest.mark.parametrize("seqlen_q,seqlen_k", [(256, 16), (64, 17)]) +def test_stage1_matches_reference(head_dim, causal, seqlen_q, seqlen_k): + torch.manual_seed(0) + n_heads, n_kv_heads = 32, 2 + dtype = torch.bfloat16 + + q = torch.randn(n_heads, seqlen_q, head_dim, dtype=dtype, device="cuda") + k = torch.randn(n_kv_heads, seqlen_k, head_dim, dtype=dtype, device="cuda") + + cu_seqlens_q = torch.tensor([0, seqlen_q], dtype=torch.int32, device="cuda") + cu_seqlens_k = torch.tensor([0, seqlen_k], dtype=torch.int32, device="cuda") + + q = q.transpose(0, 1).contiguous() + k = k.transpose(0, 1).contiguous() + + common = dict( + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + cu_seqlens_v=cu_seqlens_k, + max_seqlen_q=seqlen_q, + max_seqlen_k=seqlen_k, + causal=causal, + ) + out_ref = ref.infllmv2_attn_stage1(q, k, k, **common) + out_sgl = sgl.infllmv2_attn_stage1(q, k, k, **common) + _assert_close(out_sgl, out_ref, "stage1") + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/sgl-kernel/tests/test_infllm_v2_max_pooling.py b/sgl-kernel/tests/test_infllm_v2_max_pooling.py new file mode 100644 index 000000000..0698b28cb --- /dev/null +++ b/sgl-kernel/tests/test_infllm_v2_max_pooling.py @@ -0,0 +1,109 @@ +import pytest +import torch +from sgl_kernel import max_pooling_1d_varlen + + +def _ref_varlen( + score: torch.Tensor, # [num_heads, total_q, max_k] + cu_seqlens_q: torch.Tensor, + cu_seqlens_k: torch.Tensor, + cache_lens: torch.Tensor, + max_context_len: int, + local_blocks: int, + init_blocks: int, + block_size: int, + kernel_stride: int, +) -> torch.Tensor: + """Pure-torch reference mirroring the CUDA kernel exactly (fp32 math).""" + num_heads, total_q, _ = score.shape + out_len = (max_context_len + block_size - 1) // block_size + stride = block_size // kernel_stride + kernel_size = stride + 1 + padding = 1 + + cu_q = cu_seqlens_q.tolist() + cu_k = cu_seqlens_k.tolist() + cache = cache_lens.tolist() + batch_size = len(cache) + + out = torch.zeros(num_heads, total_q, out_len, dtype=torch.float32) + s = score.float().cpu() + for q in range(total_q): + b = 0 + for bb in range(batch_size): + if cu_q[bb] <= q < cu_q[bb + 1]: + b = bb + break + bidq_local = q - cu_q[b] + seqlen_k = cu_k[b + 1] - cu_k[b] + off_bq = (bidq_local + cache[b]) // block_size + for h in range(num_heads): + for k in range(out_len): + if (k < init_blocks) or (off_bq >= k and off_bq <= k + local_blocks): + out[h, q, k] = float("inf") + else: + start = max(k * stride - padding, 0) + end = min(start + kernel_size, seqlen_k) + if end > start: + out[h, q, k] = s[h, q, start:end].max() + else: + out[h, q, k] = float("-inf") + return out + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("num_heads", [1, 4]) +@pytest.mark.parametrize("seq_lens", [[37], [16, 48], [8, 8, 24]]) +def test_max_pooling_varlen_matches_reference(dtype, num_heads, seq_lens): + torch.manual_seed(0) + block_size = 64 + kernel_stride = 16 + local_blocks = 1 + init_blocks = 1 + max_context_len = 512 + + total_q = sum(seq_lens) + max_k = max_context_len // kernel_stride + cu = [0] + for n in seq_lens: + cu.append(cu[-1] + n) + cu_seqlens_q = torch.tensor(cu, dtype=torch.int32, device="cuda") + cu_seqlens_k = torch.tensor(cu, dtype=torch.int32, device="cuda") + cache_lens = torch.zeros(len(seq_lens), dtype=torch.int32, device="cuda") + + score = torch.randn(num_heads, total_q, max_k, dtype=dtype, device="cuda") + + out = max_pooling_1d_varlen( + score, + cu_seqlens_q, + cu_seqlens_k, + cache_lens, + max_seqlen_q=max(seq_lens), + max_context_len=max_context_len, + local_blocks=local_blocks, + init_blocks=init_blocks, + block_size=block_size, + stride=kernel_stride, + total_q=total_q, + ) + ref = _ref_varlen( + score, + cu_seqlens_q, + cu_seqlens_k, + cache_lens, + max_context_len, + local_blocks, + init_blocks, + block_size, + kernel_stride, + ).to(out.device) + + assert torch.equal(torch.isinf(out) & (out > 0), torch.isinf(ref) & (ref > 0)) + finite = torch.isfinite(ref) + torch.testing.assert_close(out[finite].float(), ref[finite], rtol=1e-2, atol=1e-2) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"]))