diff --git a/sgl-kernel/csrc/cpu/decode.cpp b/sgl-kernel/csrc/cpu/decode.cpp index 3de2708e7..85b69258f 100644 --- a/sgl-kernel/csrc/cpu/decode.cpp +++ b/sgl-kernel/csrc/cpu/decode.cpp @@ -1049,7 +1049,7 @@ void decode_attention_kernel_impl( int64_t k_strideH, int64_t v_strideN, int64_t v_strideH, - float scaling, + float sm_scale, float logit_cap, int64_t max_num_reqs, int64_t max_context_len, @@ -1103,7 +1103,7 @@ void decode_attention_kernel_impl( /* B */ k_buffer + head_id * k_strideH, /* C */ s_i, /* ind */ req_to_token + req_pool_id * max_context_len + n, - /* scl */ scaling, + /* scl */ sm_scale, /* M */ 1, /* N */ n_size, /* K */ head_size, @@ -1192,7 +1192,7 @@ void decode_attention_mla_kernel_impl( int64_t k_strideH, int64_t v_strideN, int64_t v_strideH, - float scaling, + float sm_scale, float logit_cap, int64_t max_num_reqs, int64_t max_context_len, @@ -1291,7 +1291,7 @@ void decode_attention_mla_kernel_impl( /* B */ Btmp0, /* C */ s_i); - const Vec scale_vec = Vec(scaling); + const Vec scale_vec = Vec(sm_scale); for (int64_t h = 0; h < h_size; ++h) { // s_i <- s_i * scale at::vec::map( @@ -1384,7 +1384,7 @@ void decode_attention_grouped_kernel_impl( int64_t k_strideH, int64_t v_strideN, int64_t v_strideH, - float scaling, + float sm_scale, float logit_cap, int64_t max_num_reqs, int64_t max_context_len, @@ -1457,7 +1457,7 @@ void decode_attention_grouped_kernel_impl( /* B */ k_buffer + head_kv_id * k_strideH, /* C */ s_i, /* ind */ req_to_token + req_pool_id * max_context_len + n, - /* scl */ scaling, + /* scl */ sm_scale, /* M */ h_size, /* N */ n_size, /* K */ head_size, @@ -1474,7 +1474,7 @@ void decode_attention_grouped_kernel_impl( BLOCK_H * BLOCK_N); } - // update the scaling coefficients + // update the sm_scale coefficients for (int64_t h = 0; h < h_size; ++h) { // m_i: max value per row float m_i = at::vec::reduce_all( diff --git a/sgl-kernel/csrc/cpu/extend.cpp b/sgl-kernel/csrc/cpu/extend.cpp index 55ca9ca30..63da654e4 100644 --- a/sgl-kernel/csrc/cpu/extend.cpp +++ b/sgl-kernel/csrc/cpu/extend.cpp @@ -1,71 +1,16 @@ #include "common.h" +#include "flash_attn.h" #include "gemm.h" -#include "vec.h" -#include "vec_pack.h" namespace { // [NOTE]: extend attention for CPU -// 1. tune BLOCK_M and BLOCK_N -// 2. can handle non-contiguous k_exttend and v_extend +// 1. BLOCK_M and BLOCK_N tuned for various seq lengths +// 2. can handle non-contiguous k_extend and v_extend // 3. computes attention for prefix and extend separately -// 4. TODO: vectorize `pack_vnni` and `pack_vnni2` +// 4. TODO: apply head dimension blocking to optimize GQA // -template -inline void fill_stub(scalar_t* __restrict__ out, float val, int size) { - using Vec = at::vec::Vectorized; - constexpr int kVecSize = Vec::size(); - const Vec data_vec = Vec(static_cast(val)); - int d = 0; -#pragma GCC unroll 4 - for (; d <= size - kVecSize; d += kVecSize) { - data_vec.store(out + d); - } - if (size - d > 0) { - data_vec.store(out + d, size - d); - } -} - -template -inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ input) { - static_assert(BLOCK_N % 32 == 0); - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - - constexpr int COLS = BLOCK_N / 16; - auto store = [&](auto i) { - constexpr int col = i % COLS; - // for COLS = 2, 4 use 512bit store - if constexpr (col % 2 == 0) { - fVec a_fvec0 = fVec::loadu(input + col * 16); - fVec a_fvec1 = fVec::loadu(input + col * 16 + 16); - bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); - out_bvec.store(out + col * 16); - } - }; - Unroll{}(store); -} - -template -inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, float s, int size) { - using bVec = at::vec::Vectorized; - using fVec = at::vec::Vectorized; - constexpr int kVecSize = bVec::size(); - const fVec s_fvec = fVec(s); - int d = 0; -#pragma GCC unroll 4 - for (; d <= size - kVecSize; d += kVecSize) { - fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec; - fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec; - bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); - out_bvec.store(out + d); - } - for (; d < size; ++d) { - out[d] = static_cast(acc[d] * s); - } -} - template void extend_attention_kernel_impl( scalar_t* __restrict__ o_extend, @@ -95,16 +40,13 @@ void extend_attention_kernel_impl( int k_strideH, int v_strideN, int v_strideH, - float scaling, - float logit_cap, + float sm_scale, int max_num_reqs, int max_context_len, int max_total_num_tokens, int max_len_extend, int buffer_size_per_thread, bool is_prefix_skipped) { - using Vec = at::vec::Vectorized; - // strides const int o_strideM = num_heads * head_size_v; const int o_strideH = head_size_v; @@ -112,9 +54,6 @@ void extend_attention_kernel_impl( // we use same buffer for packed key and value const int ldb_tmp = std::max(head_size, head_size_v); - const bool has_logit_cap = logit_cap > 0; - float rlogit_cap = has_logit_cap ? 1 / logit_cap : 0.f; - const int num_groups = num_heads / num_heads_kv; TORCH_CHECK(num_groups * num_heads_kv == num_heads); @@ -129,16 +68,13 @@ void extend_attention_kernel_impl( int tid = at::get_thread_num(); // s_i and s_delta: [BLOCK_M, BLOCK_N] float* __restrict__ s_i = reinterpret_cast((char*)(buffer) + tid * buffer_size_per_thread); - float* __restrict__ s_delta = s_i; + scalar_t* __restrict__ s_delta = reinterpret_cast(s_i); // v_prime: [BLOCK_M, head_size_v] float* __restrict__ v_prime = s_i + BLOCK_M * BLOCK_N; - // s_delta2: [BLOCK_M, BLOCK_N]; copy of s_delta in scalar_t - scalar_t* __restrict__ s_delta2 = reinterpret_cast(v_prime + BLOCK_N * head_size_v); - // Btmp: [BLOCK_N, max(head_size, head_size_v)] - scalar_t* __restrict__ Btmp = s_delta2 + BLOCK_M * BLOCK_N; + scalar_t* __restrict__ Btmp = reinterpret_cast(v_prime + BLOCK_M * head_size_v); // init Btmp just once for each thread to prevent NaN fill_stub(Btmp, 0.f, BLOCK_N * ldb_tmp); @@ -164,7 +100,7 @@ void extend_attention_kernel_impl( } // offset and size in MB - int m = mb * BLOCK_N; + int m = mb * BLOCK_M; int m_size = std::min(BLOCK_M, seq_len_extend - m); if (m_size <= 0) { @@ -210,51 +146,8 @@ void extend_attention_kernel_impl( /* B */ Btmp, /* C */ s_i); - const Vec scale_vec = Vec(scaling); - for (int row = 0; row < m_size; ++row) { - // s_i <- s_i * scale - at::vec::map( - [scale_vec](Vec x) { return x * scale_vec; }, s_i + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - - // TODO: `tanh` from torch uses sleef u10, going to be slow - if (has_logit_cap) { - at::vec::map( - [logit_cap, rlogit_cap](Vec x) { return Vec(logit_cap) * (x * Vec(rlogit_cap)).tanh(); }, - s_i + row * BLOCK_N, - s_i + row * BLOCK_N, - n_size); - } - - // m_i: max value per row - float m_i = at::vec::reduce_all( - [](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, s_i + row * BLOCK_N, n_size); - m_i = std::max(m_i, m_prime[row]); - - // m_delta <- exp(m' - m_i) - float m_delta = std::exp(m_prime[row] - m_i); - - // s_delta <- exp(s_i - m_i) - at::vec::map( - [m_i](Vec x) { return (x - Vec(m_i)).exp_u20(); }, s_delta + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - - // s' <- s' * m_delta + sum(s_delta) - s_prime[row] *= m_delta; - s_prime[row] += - at::vec::reduce_all([](Vec& x, Vec& y) { return x + y; }, s_delta + row * BLOCK_N, n_size); - - m_prime[row] = m_i; - - // v' <- v' * m_delta - at::vec::map( - [m_delta](Vec x) { return x * Vec(m_delta); }, - v_prime + row * head_size_v, - v_prime + row * head_size_v, - head_size_v); - - // pad s_delta with 0 first and then convert to scalar_t - fill_stub(s_delta + row * BLOCK_N + n_size, 0.f, padded_n_size - n_size); - copy_stub(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N); - } + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); // get value and pack pack_vnni2( @@ -275,7 +168,7 @@ void extend_attention_kernel_impl( /* ldb */ head_size_v, /* ldc */ head_size_v, /* add_C */ true, - /* A */ s_delta2, + /* A */ s_delta, /* B */ Btmp, /* C */ v_prime); } // loop with seq_len_prefix @@ -289,10 +182,9 @@ void extend_attention_kernel_impl( const int padded_n_size = div_up(n_size, TILE_K) * TILE_K; // get key and pack - pack_vnni( + pack_vnni( /* dst */ Btmp, /* src */ k_extend + (seq_extend_start_loc + n) * ke_strideN + head_kv_id * ke_strideH, - /* ind */ nullptr, /* N */ n_size, /* K */ head_size, /* ld_src */ ke_strideN, @@ -321,57 +213,13 @@ void extend_attention_kernel_impl( } } - const Vec scale_vec = Vec(scaling); - for (int row = 0; row < m_size; ++row) { - // s_i <- s_i * scale - at::vec::map( - [scale_vec](Vec x) { return x * scale_vec; }, s_i + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - - // TODO: `tanh` from torch uses sleef u10, going to be slow - if (has_logit_cap) { - at::vec::map( - [logit_cap, rlogit_cap](Vec x) { return Vec(logit_cap) * (x * Vec(rlogit_cap)).tanh(); }, - s_i + row * BLOCK_N, - s_i + row * BLOCK_N, - n_size); - } - - // m_i: max value per row - float m_i = at::vec::reduce_all( - [](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, s_i + row * BLOCK_N, n_size); - m_i = std::max(m_i, m_prime[row]); - - // m_delta <- exp(m' - m_i) - float m_delta = std::exp(m_prime[row] - m_i); - - // s_delta <- exp(s_i - m_i) - at::vec::map( - [m_i](Vec x) { return (x - Vec(m_i)).exp_u20(); }, s_delta + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - - // s' <- s' * m_delta + sum(s_delta) - s_prime[row] *= m_delta; - s_prime[row] += - at::vec::reduce_all([](Vec& x, Vec& y) { return x + y; }, s_delta + row * BLOCK_N, n_size); - - m_prime[row] = m_i; - - // v' <- v' * m_delta - at::vec::map( - [m_delta](Vec x) { return x * Vec(m_delta); }, - v_prime + row * head_size_v, - v_prime + row * head_size_v, - head_size_v); - - // pad s_delta with 0 first and then convert to scalar_t - fill_stub(s_delta + row * BLOCK_N + n_size, 0.f, padded_n_size - n_size); - copy_stub(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N); - } + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); // get value and pack - pack_vnni2( + pack_vnni2( /* dst */ Btmp, /* src */ v_extend + (seq_extend_start_loc + n) * ve_strideN + head_kv_id * ve_strideH, - /* ind */ nullptr, /* K */ n_size, /* N */ head_size_v, /* ld_src */ ve_strideN, @@ -386,7 +234,7 @@ void extend_attention_kernel_impl( /* ldb */ head_size_v, /* ldc */ head_size_v, /* add_C */ true, - /* A */ s_delta2, + /* A */ s_delta, /* B */ Btmp, /* C */ v_prime); } // loop with seq_len_extend @@ -406,6 +254,59 @@ void extend_attention_kernel_impl( } // anonymous namespace +template +inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int head_size_v) { + static_assert(BLOCK_M <= BLOCK_N, "Make sure BLOCK_M <= BLOCK_N to prevent buffer overflows during causal masking"); + const int size_per_thread = + /* s_i */ BLOCK_M * BLOCK_N * sizeof(float) + + /* v_prime */ BLOCK_M * head_size_v * sizeof(float) + + /* Btmp */ BLOCK_N * std::max(head_size, head_size_v) * sizeof(uint16_t); + + buffer.resize_({num_threads, size_per_thread}); + return size_per_thread; +} + +#define LAUNCH_EXTEND_ATTENTION_KERNEL(BLOCK_M, BLOCK_N) \ + do { \ + int sz = resize_buffer(buffer, num_threads, head_size, head_size_v); \ + \ + extend_attention_kernel_impl( \ + o_extend.data_ptr(), \ + q_extend.data_ptr(), \ + k_extend.data_ptr(), \ + v_extend.data_ptr(), \ + k_buffer.data_ptr(), \ + v_buffer.data_ptr(), \ + req_to_token.data_ptr(), \ + req_pool_indices.data_ptr(), \ + seq_lens.data_ptr(), \ + extend_seq_lens.data_ptr(), \ + extend_start_loc.data_ptr(), \ + buffer.data_ptr(), \ + num_seqs, \ + num_heads, \ + num_heads_kv, \ + head_size, \ + head_size_v, \ + q_strideM, \ + q_strideH, \ + ke_strideN, \ + ke_strideH, \ + ve_strideN, \ + ve_strideH, \ + k_strideN, \ + k_strideH, \ + v_strideN, \ + v_strideH, \ + sm_scale, \ + max_num_reqs, \ + max_context_len, \ + max_total_num_tokens, \ + max_len_extend, \ + sz, \ + is_prefix_skipped); \ + } while (0) + // q_extend, k_extend, v_extend, o_extend: contiguous tensors // k_buffer, v_buffer: (prefix + extend) tensors in mem_manager // @@ -449,7 +350,8 @@ void extend_attention_cpu( req_pool_indices, seq_lens, extend_seq_lens, - extend_start_loc})); + extend_start_loc, + max_len_extend})); CHECK_LAST_DIM_CONTIGUOUS_INPUT(q_extend); CHECK_INPUT(o_extend); @@ -511,58 +413,20 @@ void extend_attention_cpu( TORCH_CHECK(head_size % 32 == 0, "invalid head_size ", head_size); TORCH_CHECK(head_size_v % 32 == 0, "invalid head_size_v ", head_size_v); - // block size for query seq length - constexpr int BLOCK_M = 32; - // block size for key/value seq length - constexpr int BLOCK_N = 32; - - const int size_per_thread = - /* s_i */ BLOCK_M * BLOCK_N * sizeof(float) + - /* v_prime */ BLOCK_M * head_size_v * sizeof(float) + - /* s_delta */ BLOCK_M * BLOCK_N * sizeof(uint16_t) + - /* Btmp */ BLOCK_N * std::max(head_size, head_size_v) * sizeof(uint16_t); - int num_threads = at::get_num_threads(); - auto buffer = at::empty({num_threads, size_per_thread}, q_extend.options().dtype(at::kChar)); + auto buffer = at::empty({}, q_extend.options().dtype(at::kChar)); AT_DISPATCH_REDUCED_FLOATING_TYPES(q_extend.scalar_type(), "extend_attention_kernel", [&] { AT_DISPATCH_INDEX_TYPES(index_dtype, "extend_attention_indices", [&] { - extend_attention_kernel_impl( - o_extend.data_ptr(), - q_extend.data_ptr(), - k_extend.data_ptr(), - v_extend.data_ptr(), - k_buffer.data_ptr(), - v_buffer.data_ptr(), - req_to_token.data_ptr(), - req_pool_indices.data_ptr(), - seq_lens.data_ptr(), - extend_seq_lens.data_ptr(), - extend_start_loc.data_ptr(), - buffer.data_ptr(), - num_seqs, - num_heads, - num_heads_kv, - head_size, - head_size_v, - q_strideM, - q_strideH, - ke_strideN, - ke_strideH, - ve_strideN, - ve_strideH, - k_strideN, - k_strideH, - v_strideN, - v_strideH, - sm_scale, - logit_cap, - max_num_reqs, - max_context_len, - max_total_num_tokens, - max_len_extend, - size_per_thread, - is_prefix_skipped); + if (max_len_extend <= 256) { + LAUNCH_EXTEND_ATTENTION_KERNEL(32, 64); + } else if (max_len_extend <= 1024) { + LAUNCH_EXTEND_ATTENTION_KERNEL(128, 256); + } else if (max_len_extend <= 4096) { + LAUNCH_EXTEND_ATTENTION_KERNEL(256, 768); + } else { // max_len_extend > 4096 + LAUNCH_EXTEND_ATTENTION_KERNEL(512, 768); + } }); }); } diff --git a/sgl-kernel/csrc/cpu/flash_attn.cpp b/sgl-kernel/csrc/cpu/flash_attn.cpp new file mode 100644 index 000000000..de521a980 --- /dev/null +++ b/sgl-kernel/csrc/cpu/flash_attn.cpp @@ -0,0 +1,544 @@ +/***************************************************************************************** + * Copyright (c) 2025 - 2025 Codeplay Software Ltd. All rights reserved. + * Copyright (C) 2025 Intel Corporation, All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + ****************************************************************************************/ +#include "flash_attn.h" + +#include "common.h" +#include "gemm.h" + +// [NOTE]: flash attention interface for CPU + +namespace { + +template +void flash_attn_kernel_impl( + scalar_t* __restrict__ out, + const scalar_t* __restrict__ q, + const scalar_t* __restrict__ k, + const scalar_t* __restrict__ v, + void* __restrict__ buffer, + int seqlen_q, + int seqlen_k, + int batches, + int num_heads, + int num_heads_kv, + int head_size, + int head_size_v, + int q_strideM, + int q_strideH, + int k_strideN, + int k_strideH, + int v_strideN, + int v_strideH, + float sm_scale, + int buffer_size_per_thread, + bool causal) { + // strides + const int o_strideM = num_heads * head_size_v; + const int o_strideH = head_size_v; + + // we use same buffer for packed key and value + const int ldb_tmp = std::max(head_size, head_size_v); + + const int num_groups = num_heads / num_heads_kv; + TORCH_CHECK(num_groups * num_heads_kv == num_heads); + + // number of super locks along M + int MB = div_up(seqlen_q, BLOCK_M); + + // parallel on [batches, num_heads, MB] + parallel_for(batches * num_heads * MB, [&](int begin, int end) { + int bs{0}, head_id{0}, mb{0}; + data_index_init(begin, bs, batches, head_id, num_heads, mb, MB); + + int tid = get_thread_num(); + // s_i and s_delta: [BLOCK_M, BLOCK_N] + float* __restrict__ s_i = reinterpret_cast((char*)(buffer) + tid * buffer_size_per_thread); + scalar_t* __restrict__ s_delta = reinterpret_cast(s_i); + + // v_prime: [BLOCK_M, head_size_v] + float* __restrict__ v_prime = s_i + BLOCK_M * BLOCK_N; + + // Btmp: [BLOCK_N, max(head_size, head_size_v)] + scalar_t* __restrict__ Btmp = reinterpret_cast(v_prime + BLOCK_M * head_size_v); + + // init Btmp and Btmp2 just once for each thread to prevent NaN + fill_stub(Btmp, 0.f, BLOCK_N * ldb_tmp); + + alignas(64) float s_prime[BLOCK_M]; + alignas(64) float m_prime[BLOCK_M]; + + for (int i = begin; i < end; ++i) { + int seq_q_start_loc = bs * seqlen_q; + int seq_k_start_loc = bs * seqlen_k; + + // offset and size in MB + int m = mb * BLOCK_M; + int m_size = std::min(BLOCK_M, seqlen_q - m); + + assert(m_size > 0); + + int head_kv_id = head_id / num_groups; + + // get query + const scalar_t* __restrict__ q_ptr = q + (seq_q_start_loc + m) * q_strideM + head_id * q_strideH; + + // init v', s' and m' + fill_stub(v_prime, 0.f, m_size * head_size_v); + fill_stub(s_prime, 0.f, m_size); + fill_stub(m_prime, -std::numeric_limits::infinity(), m_size); + + int num_keys = causal ? std::min(m + m_size, seqlen_k) : seqlen_k; + for (int n = 0; n < num_keys; n += BLOCK_N) { + int n_size = std::min(BLOCK_N, num_keys - n); + + // `n_size` is K in 2nd gemm, pad to TILE_K; + const int padded_n_size = div_up(n_size, TILE_K) * TILE_K; + + // get key and pack + pack_vnni( + /* dst */ Btmp, + /* src */ k + (seq_k_start_loc + n) * k_strideN + head_kv_id * k_strideH, + /* N */ n_size, + /* K */ head_size, + /* ld_src */ k_strideN, + /* ld_dst */ BLOCK_N); + + // calculate s_i <- Q @ K + at::native::cpublas::brgemm( + /* M */ m_size, + /* N */ n_size, + /* K */ head_size, + /* lda */ q_strideM, + /* ldb */ BLOCK_N, + /* ldc */ BLOCK_N, + /* add_C */ false, + /* A */ q_ptr, + /* B */ Btmp, + /* C */ s_i); + + // apply causal mask + if (causal && num_keys - n <= BLOCK_N) { + for (int row = 0; row < m_size; ++row) { + int last_col = m + row - n; + // fill [last_col + 1, n_size) to -inf + float* row_ptr = s_i + row * BLOCK_N; + fill_stub(row_ptr + last_col + 1, -std::numeric_limits::infinity(), n_size - last_col - 1); + } + } + + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); + + // get value and pack + pack_vnni2( + /* dst */ Btmp, + /* src */ v + (seq_k_start_loc + n) * v_strideN + head_kv_id * v_strideH, + /* K */ n_size, + /* N */ head_size_v, + /* ld_src */ v_strideN, + /* ld_dst */ head_size_v); + + // calculate V' <- s_delta @ V + V' + at::native::cpublas::brgemm( + /* M */ m_size, + /* N */ head_size_v, + /* K */ padded_n_size, // n_size + /* lda */ BLOCK_N, + /* ldb */ head_size_v, + /* ldc */ head_size_v, + /* add_C */ true, + /* A */ s_delta, + /* B */ Btmp, + /* C */ v_prime); + } // loop with seqlen_k + + scalar_t* __restrict__ out_ptr = out + (seq_q_start_loc + m) * o_strideM + head_id * o_strideH; + for (int row = 0; row < m_size; ++row) { + float s = 1 / s_prime[row]; + copy_stub(out_ptr + row * o_strideM, v_prime + row * head_size_v, s, head_size_v); + } + + // move to the next index + data_index_step(bs, batches, head_id, num_heads, mb, MB); + } + at::native::cpublas::brgemm_release(); + }); +} + +template +void flash_attn_varlen_kernel_impl( + scalar_t* __restrict__ out, + const scalar_t* __restrict__ q, + const scalar_t* __restrict__ k, + const scalar_t* __restrict__ v, + const int32_t* __restrict__ cu_seqlens_q, + const int32_t* __restrict__ cu_seqlens_k, + void* __restrict__ buffer, + int32_t* __restrict__ indices, + int max_seqlen_q, + int max_seqlen_k, + int batches, + int num_heads, + int num_heads_kv, + int head_size, + int head_size_v, + int q_strideM, + int q_strideH, + int k_strideN, + int k_strideH, + int v_strideN, + int v_strideH, + float sm_scale, + int buffer_size_per_thread, + bool causal) { + // strides + const int o_strideM = num_heads * head_size_v; + const int o_strideH = head_size_v; + + // compute index (bs, mb_offset) for Query blocks + // do this sequentially as usually problem size won't be big + int idx = 0; + for (int32_t bs = 0; bs < batches; ++bs) { + int32_t seqlen_q = cu_seqlens_q[bs + 1] - cu_seqlens_q[bs]; + int32_t seqlen_k = cu_seqlens_k[bs + 1] - cu_seqlens_k[bs]; + TORCH_CHECK(seqlen_q <= max_seqlen_q && seqlen_k <= max_seqlen_k); + + int32_t blocks = div_up(seqlen_q, BLOCK_M); + for (int32_t offset = 0; offset < blocks; ++offset) { + indices[idx * 2 + 0] = bs; + indices[idx * 2 + 1] = offset; + idx++; + } + } + // number of query blocks + int MB = idx; + + // we use same buffer for packed key and value + const int ldb_tmp = std::max(head_size, head_size_v); + + const int num_groups = num_heads / num_heads_kv; + TORCH_CHECK(num_groups * num_heads_kv == num_heads); + + // parallel on [MB, num_heads] + parallel_for(num_heads * MB, [&](int begin, int end) { + int head_id{0}, mb{0}; + data_index_init(begin, head_id, num_heads, mb, MB); + + int tid = get_thread_num(); + // s_i and s_delta: [BLOCK_M, BLOCK_N] + float* __restrict__ s_i = reinterpret_cast((char*)(buffer) + tid * buffer_size_per_thread); + scalar_t* __restrict__ s_delta = reinterpret_cast(s_i); + + // v_prime: [BLOCK_M, head_size_v] + float* __restrict__ v_prime = s_i + BLOCK_M * BLOCK_N; + + // Btmp: [BLOCK_N, max(head_size, head_size_v)] + scalar_t* __restrict__ Btmp = reinterpret_cast(v_prime + BLOCK_M * head_size_v); + + // init Btmp just once for each thread to prevent NaN + fill_stub(Btmp, 0.f, BLOCK_N * ldb_tmp); + + alignas(64) float s_prime[BLOCK_M]; + alignas(64) float m_prime[BLOCK_M]; + + for (int i = begin; i < end; ++i) { + int32_t bs = indices[mb * 2 + 0]; + int32_t seq_q_start_loc = cu_seqlens_q[bs]; + int32_t seq_k_start_loc = cu_seqlens_k[bs]; + int32_t seqlen_q = cu_seqlens_q[bs + 1] - cu_seqlens_q[bs]; + + // offset and size in MB + int m = indices[mb * 2 + 1] * BLOCK_M; + int m_size = std::min(BLOCK_M, seqlen_q - m); + + assert(m_size > 0); + + int head_kv_id = head_id / num_groups; + + // get query + const scalar_t* __restrict__ q_ptr = q + (seq_q_start_loc + m) * q_strideM + head_id * q_strideH; + + // init v', s' and m' + fill_stub(v_prime, 0.f, m_size * head_size_v); + fill_stub(s_prime, 0.f, m_size); + fill_stub(m_prime, -std::numeric_limits::infinity(), m_size); + + int seqlen_k = cu_seqlens_k[bs + 1] - cu_seqlens_k[bs]; + int num_keys = causal ? std::min(m + m_size, seqlen_k) : seqlen_k; + for (int n = 0; n < num_keys; n += BLOCK_N) { + int n_size = std::min(BLOCK_N, num_keys - n); + + // `n_size` is K in 2nd gemm, pad to TILE_K; + const int padded_n_size = div_up(n_size, TILE_K) * TILE_K; + + // get key and pack + pack_vnni( + /* dst */ Btmp, + /* src */ k + (seq_k_start_loc + n) * k_strideN + head_kv_id * k_strideH, + /* N */ n_size, + /* K */ head_size, + /* ld_src */ k_strideN, + /* ld_dst */ BLOCK_N); + + // calculate s_i <- Q @ K + at::native::cpublas::brgemm( + /* M */ m_size, + /* N */ n_size, + /* K */ head_size, + /* lda */ q_strideM, + /* ldb */ BLOCK_N, + /* ldc */ BLOCK_N, + /* add_C */ false, + /* A */ q_ptr, + /* B */ Btmp, + /* C */ s_i); + + // apply causal mask + if (causal && num_keys - n <= BLOCK_N) { + for (int row = 0; row < m_size; ++row) { + int last_col = m + row - n; + // fill [last_col + 1, n_size) to -inf + float* row_ptr = s_i + row * BLOCK_N; + fill_stub(row_ptr + last_col + 1, -std::numeric_limits::infinity(), n_size - last_col - 1); + } + } + + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); + + // get value and pack + pack_vnni2( + /* dst */ Btmp, + /* src */ v + (seq_k_start_loc + n) * v_strideN + head_kv_id * v_strideH, + /* K */ n_size, + /* N */ head_size_v, + /* ld_src */ v_strideN, + /* ld_dst */ head_size_v); + + // calculate V' <- s_delta @ V + V' + at::native::cpublas::brgemm( + /* M */ m_size, + /* N */ head_size_v, + /* K */ padded_n_size, // n_size + /* lda */ BLOCK_N, + /* ldb */ head_size_v, + /* ldc */ head_size_v, + /* add_C */ true, + /* A */ s_delta, + /* B */ Btmp, + /* C */ v_prime); + } // loop with seqlen_k + + scalar_t* __restrict__ out_ptr = out + (seq_q_start_loc + m) * o_strideM + head_id * o_strideH; + for (int row = 0; row < m_size; ++row) { + float s = 1 / s_prime[row]; + copy_stub(out_ptr + row * o_strideM, v_prime + row * head_size_v, s, head_size_v); + } + + // move to the next index + data_index_step(head_id, num_heads, mb, MB); + } + at::native::cpublas::brgemm_release(); + }); +} + +} // anonymous namespace + +template +inline bool has_varlen_sequences( + const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, + int batches, + index_t max_seqlen_q, + index_t max_seqlen_k) { + const index_t* cu_seqlens_q_data = cu_seqlens_q.data_ptr(); + const index_t* cu_seqlens_k_data = cu_seqlens_k.data_ptr(); + + for (int bs = 0; bs < batches; ++bs) { + index_t seqlen_q = cu_seqlens_q_data[bs + 1] - cu_seqlens_q_data[bs]; + index_t seqlen_k = cu_seqlens_k_data[bs + 1] - cu_seqlens_k_data[bs]; + if (seqlen_q != max_seqlen_q || seqlen_k != max_seqlen_k) { + return true; + } + } + return false; +} + +template +inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int head_size_v) { + static_assert(BLOCK_M <= BLOCK_N, "Make sure BLOCK_M <= BLOCK_N to prevent buffer overflows during causal masking"); + const int size_per_thread = + /* s_i */ BLOCK_M * BLOCK_N * sizeof(float) + + /* v_prime */ BLOCK_M * head_size_v * sizeof(float) + + /* Btmp */ BLOCK_N * std::max(head_size, head_size_v) * sizeof(uint16_t); + + buffer.resize_({num_threads, size_per_thread}); + return size_per_thread; +} + +template +inline void resize_indices(at::Tensor& indices, int num_seqs, int max_seqlen_q) { + // we allocate memory based on max seqlen + indices.resize_({num_seqs, div_up(max_seqlen_q, BLOCK_M), 2}); +} + +// [NOTE]: `flash_attn_varlen_func` AMX kernel +// +// q: [num_tokens, num_heads, head_size] +// k: [num_tokens, num_heads_kv, head_size] +// v: [num_tokens, num_heads_kv, head_size_v] +// cu_seqlens_q: [num_seqs + 1] +// cu_seqlens_k: [num_seqs + 1] +// out: [num_tokens, num_heads, head_size_v] +// +at::Tensor flash_attn_varlen_func( + const at::Tensor& q, + const at::Tensor& k, + const at::Tensor& v, + const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, + int64_t max_seqlen_q, + int64_t max_seqlen_k, + bool causal) { + RECORD_FUNCTION( + "sgl_kernel::flash_attn_varlen_func", + std::vector({q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, causal})); + + CHECK_LAST_DIM_CONTIGUOUS_INPUT(q); + CHECK_LAST_DIM_CONTIGUOUS_INPUT(k); + CHECK_LAST_DIM_CONTIGUOUS_INPUT(v); + CHECK_DIM(3, q); + CHECK_DIM(3, k); + CHECK_DIM(3, v); + CHECK_INPUT(cu_seqlens_q); + CHECK_INPUT(cu_seqlens_k); + CHECK_EQ(cu_seqlens_q.scalar_type(), at::kInt); + CHECK_EQ(cu_seqlens_k.scalar_type(), at::kInt); + + int num_seqs = cu_seqlens_q.size(0) - 1; + int num_tokens = q.size(0); + int num_heads = q.size(1); + int num_heads_kv = k.size(1); + int head_size = q.size(2); + int head_size_v = v.size(2); + + // strides for q, k and v + int q_strideM = q.stride(0); + int q_strideH = q.stride(1); + int k_strideN = k.stride(0); + int k_strideH = k.stride(1); + int v_strideN = v.stride(0); + int v_strideH = v.stride(1); + + // check sizes + CHECK_EQ(k.size(2), head_size); + CHECK_EQ(v.size(1), num_heads_kv); + CHECK_EQ(cu_seqlens_k.size(0), num_seqs + 1); + + // D and DV need to be even as we transpose by 512-bit + TORCH_CHECK(head_size % 2 == 0, "invalid head_size ", head_size); + TORCH_CHECK(head_size_v % 2 == 0, "invalid head_size_v ", head_size_v); + + // softmax scale + double sm_scale = 1.0 / std::sqrt(static_cast(head_size)); + + // check whether the batch has variant lengths + const bool is_varlen = + has_varlen_sequences(cu_seqlens_q, cu_seqlens_k, num_seqs, max_seqlen_q, max_seqlen_k); + + int num_threads = at::get_num_threads(); + at::Tensor buffer = at::empty({}, q.options().dtype(at::kChar)); + at::Tensor indices = at::empty({}, q.options().dtype(at::kInt)); + at::Tensor out = at::empty({num_tokens, num_heads, head_size_v}, q.options()); + + // TODO: tune the block size + constexpr int BLOCK_M = 512; + constexpr int BLOCK_N = 768; + + AT_DISPATCH_REDUCED_FLOATING_TYPES(q.scalar_type(), "flash_attn_varlen_func", [&] { + int sz = resize_buffer(buffer, num_threads, head_size, head_size_v); + + if (is_varlen) { + resize_indices(indices, num_seqs, max_seqlen_q); + flash_attn_varlen_kernel_impl( + out.data_ptr(), + q.data_ptr(), + k.data_ptr(), + v.data_ptr(), + cu_seqlens_q.data_ptr(), + cu_seqlens_k.data_ptr(), + buffer.data_ptr(), + indices.data_ptr(), + max_seqlen_q, + max_seqlen_k, + num_seqs, + num_heads, + num_heads_kv, + head_size, + head_size_v, + q_strideM, + q_strideH, + k_strideN, + k_strideH, + v_strideN, + v_strideH, + sm_scale, + sz, + causal); + } else { + flash_attn_kernel_impl( + out.data_ptr(), + q.data_ptr(), + k.data_ptr(), + v.data_ptr(), + buffer.data_ptr(), + max_seqlen_q, + max_seqlen_k, + num_seqs, + num_heads, + num_heads_kv, + head_size, + head_size_v, + q_strideM, + q_strideH, + k_strideN, + k_strideH, + v_strideN, + v_strideH, + sm_scale, + sz, + causal); + } + }); + + return out; +} diff --git a/sgl-kernel/csrc/cpu/flash_attn.h b/sgl-kernel/csrc/cpu/flash_attn.h new file mode 100644 index 000000000..a500d11a4 --- /dev/null +++ b/sgl-kernel/csrc/cpu/flash_attn.h @@ -0,0 +1,245 @@ +#pragma once +#include "common.h" +#include "vec.h" +#include "vec_pack.h" + +template +inline void fill_stub(scalar_t* __restrict__ out, float val, int size) { + using Vec = at::vec::Vectorized; + constexpr int kVecSize = Vec::size(); + const Vec data_vec = Vec(static_cast(val)); + int d = 0; +#pragma GCC unroll 4 + for (; d <= size - kVecSize; d += kVecSize) { + data_vec.store(out + d); + } + if (size - d > 0) { + data_vec.store(out + d, size - d); + } +} + +template +inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ input) { + static_assert(BLOCK_N % 32 == 0); + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + + constexpr int COLS = BLOCK_N / 16; + auto store = [&](auto i) { + constexpr int col = i % COLS; + // for COLS = 2, 4 use 512bit store + if constexpr (col % 2 == 0) { + fVec a_fvec0 = fVec::loadu(input + col * 16); + fVec a_fvec1 = fVec::loadu(input + col * 16 + 16); + bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); + out_bvec.store(out + col * 16); + } + }; + Unroll{}(store); +} + +template +inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, float s, int size) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + constexpr int kVecSize = bVec::size(); + const fVec s_fvec = fVec(s); + int d = 0; +#pragma GCC unroll 4 + for (; d <= size - kVecSize; d += kVecSize) { + fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec; + fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec; + bVec out_bvec = convert_from_float_ext(a_fvec0, a_fvec1); + out_bvec.store(out + d); + } + for (; d < size; ++d) { + out[d] = static_cast(acc[d] * s); + } +} + +#if defined(CPU_CAPABILITY_AVX512) +template <> +inline void copy_stub(at::BFloat16* __restrict__ out, const float* __restrict__ acc, float s, int size) { + const __m512 vscale = _mm512_set1_ps(s); + int d = 0; +#pragma GCC unroll 4 + for (; d <= size - 32; d += 32) { + __m512 va0 = _mm512_mul_ps(_mm512_loadu_ps(acc + d), vscale); + __m512 va1 = _mm512_mul_ps(_mm512_loadu_ps(acc + d + 16), vscale); + __m512i vb = (__m512i)(_mm512_cvtne2ps_pbh(va1, va0)); + _mm512_storeu_si512(out + d, vb); + } + int remainder = size - d; + if (remainder > 0) { + if (remainder <= 16) { + const __mmask16 vmask = (1ULL << remainder) - 1; + __m512 va = _mm512_mul_ps(_mm512_maskz_loadu_ps(vmask, acc + d), vscale); + __m256i vb = (__m256i)(_mm512_cvtneps_pbh(va)); + _mm256_mask_storeu_epi16(reinterpret_cast<__m256i*>(out + d), vmask, vb); + } else { // remainder > 16 + const __mmask16 vmask = (1ULL << (remainder - 16)) - 1; + __m512 va0 = _mm512_mul_ps(_mm512_loadu_ps(acc + d), vscale); + __m512 va1 = _mm512_mul_ps(_mm512_maskz_loadu_ps(vmask, acc + d + 16), vscale); + __m512i vb = (__m512i)(_mm512_cvtne2ps_pbh(va1, va0)); + const __mmask32 vmask2 = (1ULL << remainder) - 1; + _mm512_mask_storeu_epi16(reinterpret_cast<__m512i*>(out + d), vmask2, vb); + } + } +} +#endif + +template +struct flash_attn_softmax { + static inline void apply( + float* __restrict__ s_i, + scalar_t* __restrict__ s_delta2, + float* __restrict__ v_prime, + float* __restrict__ s_prime, + float* __restrict__ m_prime, + int m_size, + int n_size, + int padded_n_size, + int head_size_v, + const float sm_scale) { + using Vec = at::vec::Vectorized; + const Vec scale_vec = Vec(sm_scale); + float* s_delta = s_i; + for (int row = 0; row < m_size; ++row) { + // s_i <- s_i * scale + at::vec::map( + [scale_vec](Vec x) { return x * scale_vec; }, s_i + row * BLOCK_N, s_i + row * BLOCK_N, n_size); + + // m_i: max value per row + float m_i = at::vec::reduce_all( + [](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, s_i + row * BLOCK_N, n_size); + m_i = std::max(m_i, m_prime[row]); + + // m_delta <- exp(m' - m_i) + float m_delta = std::exp(m_prime[row] - m_i); + + // s_delta <- exp(s_i - m_i) + at::vec::map( + [m_i](Vec x) { return (x - Vec(m_i)).fexp_u20(); }, s_delta + row * BLOCK_N, s_i + row * BLOCK_N, n_size); + + // s' <- s' * m_delta + sum(s_delta) + s_prime[row] *= m_delta; + s_prime[row] += at::vec::reduce_all([](Vec& x, Vec& y) { return x + y; }, s_delta + row * BLOCK_N, n_size); + + m_prime[row] = m_i; + + // v' <- v' * m_delta + at::vec::map( + [m_delta](Vec x) { return x * Vec(m_delta); }, + v_prime + row * head_size_v, + v_prime + row * head_size_v, + head_size_v); + + // pad s_delta with 0 first and then convert to scalar_t + fill_stub(s_delta + row * BLOCK_N + n_size, 0.f, padded_n_size - n_size); + copy_stub(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N); + } + } +}; + +#if defined(CPU_CAPABILITY_AVX512) +template +struct flash_attn_softmax { + static inline void apply( + float* __restrict__ s_i, + at::BFloat16* __restrict__ s_delta2, + float* __restrict__ v_prime, + float* __restrict__ s_prime, + float* __restrict__ m_prime, + int m_size, + int n_size, + int padded_n_size, + int head_size_v, + const float sm_scale) { + float* s_delta = s_i; + const __m512 vscale = _mm512_set1_ps(sm_scale); + + int n_remainder = n_size & 15; // 0xF + const __mmask16 vmask = (1ULL << n_remainder) - 1; + + int v_remainder = head_size_v & 15; // 0xF + const __mmask16 vmask1 = (1ULL << v_remainder) - 1; + + constexpr float NEG_INF = -std::numeric_limits::infinity(); + + __m512 va; + __m256i vb; + __m512 vmax; + __m512 vsum; + __m512 vmdelta; + + const __m512 vneg_inf = _mm512_set1_ps(NEG_INF); + + for (int m = 0; m < m_size; ++m) { + vmax = vneg_inf; + + // s_i <- s_i * scale + int n = 0; + for (; n <= n_size - 16; n += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); + vmax = _mm512_max_ps(va, vmax); + } + if (n_remainder > 0) { + va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); + vmax = _mm512_max_ps(va, vmax); + } + + // m_i: max value per row + float m_i = _mm512_reduce_max_ps(vmax); + vmax = _mm512_set1_ps(m_i); + + // m_delta <- exp(m' - m_i) + float m_delta = std::exp(m_prime[m] - m_i); + + // s_delta <- exp(s_i - m_i) + vsum = _mm512_setzero_ps(); + for (n = 0; n <= n_size - 16; n += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); + va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); + vsum = _mm512_add_ps(vsum, va); + + vb = (__m256i)(_mm512_cvtneps_pbh(va)); + _mm256_storeu_si256(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vb); + } + if (n_remainder > 0) { + va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); + va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); + vsum = _mm512_add_ps(vsum, va); + + vb = (__m256i)(_mm512_cvtneps_pbh(va)); + _mm256_mask_storeu_epi16(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vmask, vb); + } + + // s' <- s' * m_delta + sum(s_delta) + s_prime[m] *= m_delta; + s_prime[m] += _mm512_reduce_add_ps(vsum); + + m_prime[m] = m_i; + + // pad s_delta with 0, pad_size range from [0, 32) + int pad_size = padded_n_size - n_size; + if (pad_size > 0) { + const __m512i vzero = _mm512_setzero_si512(); + __mmask32 vmask2 = (1ULL << pad_size) - 1; + _mm512_mask_storeu_epi16(reinterpret_cast<__m512i*>(s_delta2 + m * BLOCK_N + n_size), vmask2, vzero); + } + + // v' <- v' * m_delta + vmdelta = _mm512_set1_ps(m_delta); + int k = 0; + for (; k <= head_size_v - 16; k += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(v_prime + m * head_size_v + k), vmdelta); + _mm512_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), va); + } + if (v_remainder > 0) { + va = _mm512_mul_ps(_mm512_maskz_loadu_ps(vmask1, v_prime + m * head_size_v + k), vmdelta); + _mm512_mask_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), vmask1, va); + } + } + } +}; +#endif diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index 428fa090d..8c66e1d2b 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -109,6 +109,17 @@ void extend_attention_cpu( double sm_scale, double logit_cap); +// flash attention +at::Tensor flash_attn_varlen_func( + const at::Tensor& q, + const at::Tensor& k, + const at::Tensor& v, + const at::Tensor& cu_seqlens_q, + const at::Tensor& cu_seqlens_k, + int64_t max_seqlen_q, + int64_t max_seqlen_k, + bool causal); + // linear attention std::tuple chunk_gated_delta_rule_cpu( const at::Tensor& query, @@ -382,6 +393,12 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "extend_start_loc, int max_len_extend, float sm_scale, float logit_cap) -> ()"); m.impl("extend_attention_cpu", torch::kCPU, &extend_attention_cpu); + // flash attn + m.def( + "flash_attn_varlen_func(Tensor q, Tensor k, Tensor v, Tensor cu_seqlens_q, Tensor cu_seqlens_k, " + "int max_seqlen_q, int max_seqlen_k, bool causal) -> Tensor"); + m.impl("flash_attn_varlen_func", torch::kCPU, &flash_attn_varlen_func); + // linear attn m.def( "chunk_gated_delta_rule_cpu(Tensor query, Tensor key, Tensor value, Tensor g, Tensor beta, " diff --git a/sgl-kernel/csrc/cpu/vec.h b/sgl-kernel/csrc/cpu/vec.h index 486ff260a..107022ffd 100644 --- a/sgl-kernel/csrc/cpu/vec.h +++ b/sgl-kernel/csrc/cpu/vec.h @@ -148,7 +148,7 @@ inline __attribute__((always_inline)) __m512bh CVT_FP8_TO_BF16_EXT(__m256i a) { #endif // vector to scalar reduction -#if defined(CPU_CAPABILITY_AVX512) && 0 +#if defined(CPU_CAPABILITY_AVX512) inline float vec_reduce_sum(const Vectorized& a) { return _mm512_reduce_add_ps(__m512(a)); } @@ -323,6 +323,56 @@ inline std::tuple<__m512i, __m512i> transpose_2x32_16bit(__m512i r0, __m512i r1) } #pragma GCC diagnostic pop +inline __attribute__((always_inline)) __m512 _mm512_fexp_u20_ps(const __m512 values) { + const __m512 vec_c0 = _mm512_set1_ps(0.00010703434948458272f); + const __m512 vec_c1 = _mm512_set1_ps(0.30354260500649682f); + const __m512 vec_c2 = _mm512_set1_ps(-0.22433836478672356); + const __m512 vec_c3 = _mm512_set1_ps(-0.079204240219773236); + + const __m512 vec_exp_log2ef = _mm512_castsi512_ps(_mm512_set1_epi32(0x3fb8aa3b)); // log2(e) + + const __m512 vec_a = _mm512_set1_ps(std::pow(2, 23) / std::log2(2)); + const __m512 vec_b = _mm512_set1_ps(std::pow(2, 23) * 127.f); + + const __m512 vec_ln_flt_min = _mm512_castsi512_ps(_mm512_set1_epi32(0xc2aeac50)); + const __m512 vec_ln_flt_max = _mm512_castsi512_ps(_mm512_set1_epi32(0x42b17218)); + __m512i vec_infinity = _mm512_set1_epi32(0x7F800000); + __m512i vec_zero = _mm512_setzero_epi32(); + + // Fast Exponential Computation on SIMD Architectures + // A. Cristiano I. Malossi, Yves Ineichen, Costas Bekas, and Alessandro + // Curioni exp(x) = 2**(x * log2(e)) + // = 2**xi * 2**xf - TIPS we are using the EEEE floating point + // representation with identification to the exponent and the + // mentissa + // 2**xf will be approximated to a polynomial of degree 3 computed with + // Horner method + // mask for the boundary condition + auto min_mask = _mm512_cmp_ps_mask(values, vec_ln_flt_min, _CMP_LT_OS); + auto max_mask = _mm512_cmp_ps_mask(values, vec_ln_flt_max, _CMP_GT_OS); + + // transformation with log2(e) + auto vec_src = _mm512_mul_ps(values, vec_exp_log2ef); + auto vec_fractional = _mm512_sub_ps(vec_src, _mm512_floor_ps(vec_src)); + + // compute polynomial using Horner Scheme, for superscalar processor + auto vec_res = _mm512_fmadd_ps(vec_fractional, vec_c3, vec_c2); + vec_res = _mm512_fmadd_ps(vec_fractional, vec_res, vec_c1); + vec_res = _mm512_fmadd_ps(vec_fractional, vec_res, vec_c0); + + vec_src = _mm512_sub_ps(vec_src, vec_res); + // the tips is here, headache in perspective + auto tmp = _mm512_fmadd_ps(vec_a, vec_src, vec_b); + // headache bis - we loose precision with the cast but it "fits", but ok + // after f32 -> f16 later + __m512i casted_integer = _mm512_cvttps_epi32(tmp); + // boundary condition, lower than the min -> 0 + casted_integer = _mm512_mask_mov_epi32(casted_integer, min_mask, vec_zero); + // boundary condition, larger than the max -> +oo + casted_integer = _mm512_mask_mov_epi32(casted_integer, max_mask, vec_infinity); + // final interpretation to float + return _mm512_castsi512_ps(casted_integer); +} #endif } // anonymous namespace diff --git a/sgl-kernel/csrc/cpu/vec_pack.h b/sgl-kernel/csrc/cpu/vec_pack.h index bf7433dfd..4a166111c 100644 --- a/sgl-kernel/csrc/cpu/vec_pack.h +++ b/sgl-kernel/csrc/cpu/vec_pack.h @@ -13,30 +13,6 @@ inline index_t get_index(index_t* ind, int i) { } #if defined(CPU_CAPABILITY_AVX512) -// key: from [N, 32] to [32/2, N, 2] -template -inline void -pack_vnni_Nx32(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int N, int ld_src, int ld_dst) { - __m512i vinputs[16]; - - int n = 0; - for (; n < N; ++n) { - vinputs[n] = _mm512_loadu_si512(src + n * ld_src); - } - // padding with zero to avoid uninitialized vectors - for (; n < 16; ++n) { - vinputs[n] = _mm512_set1_epi32(0); - } - - // pack key - transpose_16x16_32bit(vinputs); - - const __mmask16 vmask = (1 << N) - 1; - for (int k = 0; k < 16; ++k) { - _mm512_mask_storeu_epi32(dst + k * ld_dst * 2, vmask, vinputs[k]); - } -} - // key: from [N, 32] to [32/2, N, 2] template inline void pack_vnni_Nx32( @@ -67,26 +43,37 @@ inline void pack_vnni_Nx32( } } -// value: from [K, 32] to [K/2, 32, 2] -template -inline void -pack_vnni_Kx32(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int K, int ld_src, int ld_dst) { - __m512i vinputs[2]; +template +inline void pack_vnni_N_remainder( + scalar_t* __restrict__ dst, + const scalar_t* __restrict__ src, + const index_t* __restrict__ ind, + int N, + int K, + int ld_src, + int ld_dst) { + __m512i vinputs[16]; - int k = 0; - for (; k < K; ++k) { - vinputs[k] = _mm512_loadu_si512(src + k * ld_src); + int K2 = K >> 1; + const __mmask16 vmask = (1 << K2) - 1; + + int n = 0; + for (; n < N; ++n) { + index_t index = get_index(ind, n); + vinputs[n] = _mm512_maskz_loadu_epi32(vmask, src + index * ld_src); } // padding with zero to avoid uninitialized vectors - for (; k < 2; ++k) { - vinputs[k] = _mm512_set1_epi32(0); + for (; n < 16; ++n) { + vinputs[n] = _mm512_set1_epi32(0); } - // pack value - __m512i d0, d1; - std::tie(d0, d1) = transpose_2x32_16bit(vinputs[0], vinputs[1]); - _mm512_storeu_si512(dst + 0 * ld_dst * 2, d0); - _mm512_storeu_si512(dst + 0 * ld_dst * 2 + 32, d1); + // pack key + transpose_16x16_32bit(vinputs); + + const __mmask16 vmask2 = (1 << N) - 1; + for (int k = 0; k < K2; ++k) { + _mm512_mask_storeu_epi32(dst + k * ld_dst * 2, vmask2, vinputs[k]); + } } // value: from [K, 32] to [K/2, 32, 2] @@ -116,42 +103,50 @@ inline void pack_vnni_Kx32( _mm512_storeu_si512(dst + 0 * ld_dst * 2, d0); _mm512_storeu_si512(dst + 0 * ld_dst * 2 + 32, d1); } -#endif -// convert to vnni format -// from [N, K/2, 2] to [K/2, N, 2] for bfloat16 and float16 -template -void pack_vnni(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int N, int K, int ld_src, int ld_dst) { -#if defined(CPU_CAPABILITY_AVX512) - const int NB = div_up(N, 16); - const int KB = K / 32; // no remainder - - for (int nb = 0; nb < NB; ++nb) { - for (int kb = 0; kb < KB; ++kb) { - // handle 16x512bits each block - int nb_size = std::min(N - nb * 16, 16); - pack_vnni_Nx32( - /* dst */ dst + ((kb * 32) >> 1) * ld_dst * 2 + nb * 16 * 2, - /* src */ src + kb * 32 + nb * 16 * ld_src, - /* N */ nb_size, - /* ld_src */ ld_src, - /* ld_dst */ ld_dst); - } - } -#else - for (int n = 0; n < N; ++n) { - for (int k = 0; k < K / 2; ++k) { - for (int d = 0; d < 2; ++d) { - dst[k * ld_dst * 2 + n * 2 + d] = src[n * ld_src + k * 2 + d]; - } - } - } -#endif -} - -// convert to vnni format -// from [N, K/2, 2] to [K/2, N, 2] for bfloat16 and float16 template +inline void pack_vnni_K_remainder( + scalar_t* __restrict__ dst, + const scalar_t* __restrict__ src, + const index_t* __restrict__ ind, + int K, + int N, + int ld_src, + int ld_dst) { + __m512i vinputs[2]; + + const __mmask32 vmask = (1 << N) - 1; + + int k = 0; + for (; k < K; ++k) { + index_t index = get_index(ind, k); + vinputs[k] = _mm512_maskz_loadu_epi16(vmask, src + index * ld_src); + } + // padding with zero to avoid uninitialized vectors + for (; k < 2; ++k) { + vinputs[k] = _mm512_set1_epi32(0); + } + + // pack value + __m512i d0, d1; + std::tie(d0, d1) = transpose_2x32_16bit(vinputs[0], vinputs[1]); + + if (N <= 16) { + // 2N * 16bits: N * 32bits + const __mmask16 vmask2 = (1 << N) - 1; + _mm512_mask_storeu_epi32(dst + 0 * ld_dst * 2, vmask2, d0); + } else { + // 2(N-16) * 16bits: (N-16) * 32bits + const __mmask16 vmask2 = (1 << (N - 16)) - 1; + _mm512_storeu_epi32(dst + 0 * ld_dst * 2, d0); + _mm512_mask_storeu_epi32(dst + 0 * ld_dst * 2 + 32, vmask2, d1); + } +} +#endif + +// convert to vnni format +// from [N, K/2, 2] to [K/2, N, 2] for bfloat16 and float16 +template void pack_vnni( scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, @@ -162,13 +157,13 @@ void pack_vnni( int ld_dst) { #if defined(CPU_CAPABILITY_AVX512) const int NB = div_up(N, 16); - const int KB = K / 32; // no remainder - const bool is_indexed = ind != nullptr; + const int KB = K / 32; + const int K_remainder = K - KB * 32; for (int nb = 0; nb < NB; ++nb) { + int nb_size = std::min(N - nb * 16, 16); for (int kb = 0; kb < KB; ++kb) { // handle 16x512bits each block - int nb_size = std::min(N - nb * 16, 16); pack_vnni_Nx32( /* dst */ dst + ((kb * 32) >> 1) * ld_dst * 2 + nb * 16 * 2, /* src */ src + kb * 32 + (is_indexed ? 0 : nb * 16 * ld_src), @@ -177,6 +172,16 @@ void pack_vnni( /* ld_src */ ld_src, /* ld_dst */ ld_dst); } + if (K_remainder > 0) { + pack_vnni_N_remainder( + /* dst */ dst + ((KB * 32) >> 1) * ld_dst * 2 + nb * 16 * 2, + /* src */ src + KB * 32 + (is_indexed ? 0 : nb * 16 * ld_src), + /* ind */ is_indexed ? ind + nb * 16 : nullptr, + /* N */ nb_size, + /* K */ K_remainder, + /* ld_src */ ld_src, + /* ld_dst */ ld_dst); + } } #else for (int n = 0; n < N; ++n) { @@ -190,47 +195,27 @@ void pack_vnni( #endif } -// convert to vnni format -// from [K/2, 2, N] to [K/2, N, 2] for bfloat16 and float16 template -void pack_vnni2(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int K, int N, int ld_src, int ld_dst) { -#if defined(CPU_CAPABILITY_AVX512) - const int KB = div_up(K, 2); - const int NB = N / 32; // no remainder +void pack_vnni(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int N, int K, int ld_src, int ld_dst) { + pack_vnni(dst, src, nullptr, N, K, ld_src, ld_dst); +} - for (int kb = 0; kb < KB; ++kb) { - for (int nb = 0; nb < NB; ++nb) { - // handle 2x512bits each block - int kb_size = std::min(K - kb * 2, 2); - pack_vnni_Kx32( - /* dst */ dst + ((kb * 2) >> 1) * ld_dst * 2 + nb * 32 * 2, - /* src */ src + kb * 2 * ld_src + nb * 32, - /* K */ kb_size, - /* ld_src */ ld_src, - /* ld_dst */ ld_dst); - } - } -#else - int k = 0; - for (; k < (K >> 1) * 2; k += 2) { - for (int n = 0; n < N; ++n) { - dst[(k >> 1) * ld_dst * 2 + n * 2 + 0] = src[k * ld_src + n]; - dst[(k >> 1) * ld_dst * 2 + n * 2 + 1] = src[(k + 1) * ld_src + n]; - } - } - if (K % 2 != 0) { - for (int n = 0; n < N; ++n) { - dst[(K >> 1) * ld_dst * 2 + n * 2 + 0] = src[(K - 1) * ld_src + n]; - dst[(K >> 1) * ld_dst * 2 + n * 2 + 1] = 0; - } - k += 2; - } -#endif +template +void pack_vnni( + scalar_t* __restrict__ dst, + const scalar_t* __restrict__ src, + const index_t* __restrict__ ind, + int N, + int K, + int ld_src, + int ld_dst) { + assert(ind != nullptr); + pack_vnni(dst, src, ind, N, K, ld_src, ld_dst); } // convert to vnni format // from [K/2, 2, N] to [K/2, N, 2] for bfloat16 and float16 -template +template void pack_vnni2( scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, @@ -241,13 +226,13 @@ void pack_vnni2( int ld_dst) { #if defined(CPU_CAPABILITY_AVX512) const int KB = div_up(K, 2); - const int NB = N / 32; // no remainder - const bool is_indexed = ind != nullptr; + const int NB = N / 32; + const int N_remainder = N - NB * 32; for (int kb = 0; kb < KB; ++kb) { + int kb_size = std::min(K - kb * 2, 2); for (int nb = 0; nb < NB; ++nb) { // handle 2x512bits each block - int kb_size = std::min(K - kb * 2, 2); pack_vnni_Kx32( /* dst */ dst + ((kb * 2) >> 1) * ld_dst * 2 + nb * 32 * 2, /* src */ src + (is_indexed ? 0 : kb * 2 * ld_src) + nb * 32, @@ -256,6 +241,16 @@ void pack_vnni2( /* ld_src */ ld_src, /* ld_dst */ ld_dst); } + if (N_remainder > 0) { + pack_vnni_K_remainder( + /* dst */ dst + ((kb * 2) >> 1) * ld_dst * 2 + NB * 32 * 2, + /* src */ src + (is_indexed ? 0 : kb * 2 * ld_src) + NB * 32, + /* ind */ is_indexed ? ind + kb * 2 : nullptr, + /* K */ kb_size, + /* N */ N_remainder, + /* ld_src */ ld_src, + /* ld_dst */ ld_dst); + } } #else int k = 0; @@ -278,4 +273,22 @@ void pack_vnni2( #endif } +template +void pack_vnni2(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int K, int N, int ld_src, int ld_dst) { + pack_vnni2(dst, src, nullptr, K, N, ld_src, ld_dst); +} + +template +void pack_vnni2( + scalar_t* __restrict__ dst, + const scalar_t* __restrict__ src, + const index_t* __restrict__ ind, + int K, + int N, + int ld_src, + int ld_dst) { + assert(ind != nullptr); + pack_vnni2(dst, src, ind, K, N, ld_src, ld_dst); +} + } // anonymous namespace diff --git a/test/srt/cpu/test_extend.py b/test/srt/cpu/test_extend.py index 7277050c2..fc315c8da 100644 --- a/test/srt/cpu/test_extend.py +++ b/test/srt/cpu/test_extend.py @@ -183,6 +183,7 @@ class TestExtendAttention(CustomTestCase): self._test_extend_attention_once(1, 123, 1, 1, 128, 96, is_mla) self._test_extend_attention_once(1, 123, 16, 1, 128, 96, is_mla) self._test_extend_attention_once(4, 1230, 16, 4, 128, 96, is_mla) + self._test_extend_attention_once(1, 9000, 16, 1, 32, 32, is_mla) if __name__ == "__main__": diff --git a/test/srt/cpu/test_flash_attn.py b/test/srt/cpu/test_flash_attn.py new file mode 100644 index 000000000..e84e12b1c --- /dev/null +++ b/test/srt/cpu/test_flash_attn.py @@ -0,0 +1,115 @@ +import unittest + +import sgl_kernel # noqa: F401 +import torch +import torch.nn.functional as F +from utils import parametrize, precision + +from sglang.test.test_utils import CustomTestCase + +flash_attn_varlen_func = torch.ops.sgl_kernel.flash_attn_varlen_func + + +torch.manual_seed(1234) + + +def flash_attn_varlen_ref( + q, + k, + v, + cu_seqlens_q, + cu_seqlens_k, + is_causal, + enable_gqa, +): + cu_q = cu_seqlens_q.tolist() + cu_k = cu_seqlens_k.tolist() + batch = len(cu_k) - 1 + + # [T, H, D] -> [1, H, T, D] + q, k, v = [x.unsqueeze(0).transpose(1, 2) for x in [q, k, v]] + + B, H, T, D = q.shape + out = torch.empty(B, H, T, v.size(-1), dtype=q.dtype) + for b in range(batch): + start_q, end_q = cu_q[b], cu_q[b + 1] + start_k, end_k = cu_k[b], cu_k[b + 1] + + out[:, :, start_q:end_q, :] = F.scaled_dot_product_attention( + q[:, :, start_q:end_q, :], + k[:, :, start_k:end_k, :], + v[:, :, start_k:end_k, :], + is_causal=is_causal, + enable_gqa=enable_gqa, + ) + + # [1, H, T, D] -> [T, H, D] + return out.transpose(1, 2).squeeze(0) + + +class TestFlashAttn(CustomTestCase): + + @parametrize( + batch=[4], + max_seqlen_q=[35, 96], + max_seqlen_k=[35, 96], + num_heads=[16], + num_heads_kv=[16, 2], + head_dim=[32, 48], # test when D is not 32x + head_dim_v=[32], + is_causal=[True, False], + ) + def test_flash_attn_varlen( + self, + batch, + max_seqlen_q, + max_seqlen_k, + num_heads, + num_heads_kv, + head_dim, + head_dim_v, + is_causal, + ): + dtype = torch.bfloat16 + + # random seqlens for k and kv + seqlens_q = torch.randint(1, max_seqlen_q, (batch,), dtype=torch.int32) + seqlens_k = torch.randint(1, max_seqlen_k, (batch,), dtype=torch.int32) + cu_seqlens_q = torch.zeros((batch + 1,), dtype=torch.int32) + cu_seqlens_k = torch.zeros((batch + 1,), dtype=torch.int32) + cu_seqlens_q[1:] = torch.cumsum(seqlens_q, 0) + cu_seqlens_k[1:] = torch.cumsum(seqlens_k, 0) + + sum_seqlen_q = seqlens_q.sum().item() + sum_seqlen_k = seqlens_k.sum().item() + q = torch.randn(sum_seqlen_q, num_heads, head_dim).to(dtype) + k = torch.randn(sum_seqlen_k, num_heads_kv, head_dim).to(dtype) + v = torch.randn(sum_seqlen_k, num_heads_kv, head_dim_v).to(dtype) + + out_ref = flash_attn_varlen_ref( + q, + k, + v, + cu_seqlens_q, + cu_seqlens_k, + is_causal=is_causal, + enable_gqa=num_heads != num_heads_kv, + ) + + out = flash_attn_varlen_func( + q, + k, + v, + cu_seqlens_q, + cu_seqlens_k, + seqlens_q.max().item(), + seqlens_k.max().item(), + is_causal, + ) + + atol = rtol = precision[dtype] + torch.testing.assert_close(out_ref, out, atol=atol, rtol=rtol) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index 3e98d0cca..0c1bd240e 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -106,6 +106,7 @@ suite_xeon = { TestFile("cpu/test_cpu_graph.py"), TestFile("cpu/test_decode.py"), TestFile("cpu/test_extend.py"), + TestFile("cpu/test_flash_attn.py"), TestFile("cpu/test_gemm.py"), TestFile("cpu/test_intel_amx_attention_backend_a.py"), TestFile("cpu/test_intel_amx_attention_backend_b.py"),