[CPU] optimize flash_attn_varlen_func (#15708)
This commit is contained in:
@@ -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<float>(
|
||||
@@ -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<float>(
|
||||
|
||||
+81
-217
@@ -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 <typename scalar_t>
|
||||
inline void fill_stub(scalar_t* __restrict__ out, float val, int size) {
|
||||
using Vec = at::vec::Vectorized<scalar_t>;
|
||||
constexpr int kVecSize = Vec::size();
|
||||
const Vec data_vec = Vec(static_cast<scalar_t>(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 <typename scalar_t, int BLOCK_N>
|
||||
inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ input) {
|
||||
static_assert(BLOCK_N % 32 == 0);
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
|
||||
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<scalar_t>(a_fvec0, a_fvec1);
|
||||
out_bvec.store(out + col * 16);
|
||||
}
|
||||
};
|
||||
Unroll<COLS>{}(store);
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, float s, int size) {
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
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<scalar_t>(a_fvec0, a_fvec1);
|
||||
out_bvec.store(out + d);
|
||||
}
|
||||
for (; d < size; ++d) {
|
||||
out[d] = static_cast<scalar_t>(acc[d] * s);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename scalar_t, typename index_t, int BLOCK_M, int BLOCK_N>
|
||||
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<float>;
|
||||
|
||||
// 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<float*>((char*)(buffer) + tid * buffer_size_per_thread);
|
||||
float* __restrict__ s_delta = s_i;
|
||||
scalar_t* __restrict__ s_delta = reinterpret_cast<scalar_t*>(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<scalar_t*>(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<scalar_t*>(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<float>(
|
||||
[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<float>(
|
||||
[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<float>(
|
||||
[](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<float>(
|
||||
[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<float>([](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<float>(
|
||||
[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<scalar_t, BLOCK_N>(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N);
|
||||
}
|
||||
flash_attn_softmax<scalar_t, BLOCK_M, BLOCK_N>::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<scalar_t, index_t>(
|
||||
@@ -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<scalar_t, index_t>(
|
||||
pack_vnni<scalar_t>(
|
||||
/* 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<float>(
|
||||
[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<float>(
|
||||
[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<float>(
|
||||
[](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<float>(
|
||||
[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<float>([](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<float>(
|
||||
[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<scalar_t, BLOCK_N>(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N);
|
||||
}
|
||||
flash_attn_softmax<scalar_t, BLOCK_M, BLOCK_N>::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<scalar_t, index_t>(
|
||||
pack_vnni2<scalar_t>(
|
||||
/* 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 <int BLOCK_M, int BLOCK_N>
|
||||
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<BLOCK_M, BLOCK_N>(buffer, num_threads, head_size, head_size_v); \
|
||||
\
|
||||
extend_attention_kernel_impl<scalar_t, index_t, BLOCK_M, BLOCK_N>( \
|
||||
o_extend.data_ptr<scalar_t>(), \
|
||||
q_extend.data_ptr<scalar_t>(), \
|
||||
k_extend.data_ptr<scalar_t>(), \
|
||||
v_extend.data_ptr<scalar_t>(), \
|
||||
k_buffer.data_ptr<scalar_t>(), \
|
||||
v_buffer.data_ptr<scalar_t>(), \
|
||||
req_to_token.data_ptr<index_t>(), \
|
||||
req_pool_indices.data_ptr<int64_t>(), \
|
||||
seq_lens.data_ptr<int64_t>(), \
|
||||
extend_seq_lens.data_ptr<index_t>(), \
|
||||
extend_start_loc.data_ptr<index_t>(), \
|
||||
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<scalar_t, index_t, BLOCK_M, BLOCK_N>(
|
||||
o_extend.data_ptr<scalar_t>(),
|
||||
q_extend.data_ptr<scalar_t>(),
|
||||
k_extend.data_ptr<scalar_t>(),
|
||||
v_extend.data_ptr<scalar_t>(),
|
||||
k_buffer.data_ptr<scalar_t>(),
|
||||
v_buffer.data_ptr<scalar_t>(),
|
||||
req_to_token.data_ptr<index_t>(),
|
||||
req_pool_indices.data_ptr<int64_t>(),
|
||||
seq_lens.data_ptr<int64_t>(),
|
||||
extend_seq_lens.data_ptr<index_t>(),
|
||||
extend_start_loc.data_ptr<index_t>(),
|
||||
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);
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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 <typename scalar_t, int BLOCK_M, int BLOCK_N>
|
||||
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<float*>((char*)(buffer) + tid * buffer_size_per_thread);
|
||||
scalar_t* __restrict__ s_delta = reinterpret_cast<scalar_t*>(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<scalar_t*>(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<scalar_t>::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<scalar_t>(
|
||||
/* 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<float>::infinity(), n_size - last_col - 1);
|
||||
}
|
||||
}
|
||||
|
||||
flash_attn_softmax<scalar_t, BLOCK_M, BLOCK_N>::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<scalar_t>(
|
||||
/* 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<scalar_t>(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 <typename scalar_t, int BLOCK_M, int BLOCK_N>
|
||||
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<float*>((char*)(buffer) + tid * buffer_size_per_thread);
|
||||
scalar_t* __restrict__ s_delta = reinterpret_cast<scalar_t*>(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<scalar_t*>(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<scalar_t>::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<scalar_t>(
|
||||
/* 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<float>::infinity(), n_size - last_col - 1);
|
||||
}
|
||||
}
|
||||
|
||||
flash_attn_softmax<scalar_t, BLOCK_M, BLOCK_N>::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<scalar_t>(
|
||||
/* 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<scalar_t>(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 <typename index_t>
|
||||
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<index_t>();
|
||||
const index_t* cu_seqlens_k_data = cu_seqlens_k.data_ptr<index_t>();
|
||||
|
||||
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 <int BLOCK_M, int BLOCK_N>
|
||||
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 <int BLOCK_M>
|
||||
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<c10::IValue>({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<double>(head_size));
|
||||
|
||||
// check whether the batch has variant lengths
|
||||
const bool is_varlen =
|
||||
has_varlen_sequences<int32_t>(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<BLOCK_M, BLOCK_N>(buffer, num_threads, head_size, head_size_v);
|
||||
|
||||
if (is_varlen) {
|
||||
resize_indices<BLOCK_M>(indices, num_seqs, max_seqlen_q);
|
||||
flash_attn_varlen_kernel_impl<scalar_t, BLOCK_M, BLOCK_N>(
|
||||
out.data_ptr<scalar_t>(),
|
||||
q.data_ptr<scalar_t>(),
|
||||
k.data_ptr<scalar_t>(),
|
||||
v.data_ptr<scalar_t>(),
|
||||
cu_seqlens_q.data_ptr<int32_t>(),
|
||||
cu_seqlens_k.data_ptr<int32_t>(),
|
||||
buffer.data_ptr(),
|
||||
indices.data_ptr<int32_t>(),
|
||||
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<scalar_t, BLOCK_M, BLOCK_N>(
|
||||
out.data_ptr<scalar_t>(),
|
||||
q.data_ptr<scalar_t>(),
|
||||
k.data_ptr<scalar_t>(),
|
||||
v.data_ptr<scalar_t>(),
|
||||
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;
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
#pragma once
|
||||
#include "common.h"
|
||||
#include "vec.h"
|
||||
#include "vec_pack.h"
|
||||
|
||||
template <typename scalar_t>
|
||||
inline void fill_stub(scalar_t* __restrict__ out, float val, int size) {
|
||||
using Vec = at::vec::Vectorized<scalar_t>;
|
||||
constexpr int kVecSize = Vec::size();
|
||||
const Vec data_vec = Vec(static_cast<scalar_t>(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 <typename scalar_t, int BLOCK_N>
|
||||
inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ input) {
|
||||
static_assert(BLOCK_N % 32 == 0);
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
|
||||
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<scalar_t>(a_fvec0, a_fvec1);
|
||||
out_bvec.store(out + col * 16);
|
||||
}
|
||||
};
|
||||
Unroll<COLS>{}(store);
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, float s, int size) {
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
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<scalar_t>(a_fvec0, a_fvec1);
|
||||
out_bvec.store(out + d);
|
||||
}
|
||||
for (; d < size; ++d) {
|
||||
out[d] = static_cast<scalar_t>(acc[d] * s);
|
||||
}
|
||||
}
|
||||
|
||||
#if defined(CPU_CAPABILITY_AVX512)
|
||||
template <>
|
||||
inline void copy_stub<at::BFloat16>(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 <typename scalar_t, int BLOCK_M, int BLOCK_N>
|
||||
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<float>;
|
||||
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<float>(
|
||||
[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<float>(
|
||||
[](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<float>(
|
||||
[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<float>([](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<float>(
|
||||
[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<scalar_t, BLOCK_N>(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
#if defined(CPU_CAPABILITY_AVX512)
|
||||
template <int BLOCK_M, int BLOCK_N>
|
||||
struct flash_attn_softmax<at::BFloat16, BLOCK_M, BLOCK_N> {
|
||||
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<float>::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
|
||||
@@ -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<at::Tensor, at::Tensor> 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, "
|
||||
|
||||
@@ -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<float>& 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
|
||||
|
||||
+127
-114
@@ -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 <typename scalar_t>
|
||||
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 <typename scalar_t, typename index_t>
|
||||
inline void pack_vnni_Nx32(
|
||||
@@ -67,26 +43,37 @@ inline void pack_vnni_Nx32(
|
||||
}
|
||||
}
|
||||
|
||||
// value: from [K, 32] to [K/2, 32, 2]
|
||||
template <typename scalar_t>
|
||||
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 <typename scalar_t, typename index_t>
|
||||
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 <typename scalar_t>
|
||||
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<scalar_t>(
|
||||
/* 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 <typename scalar_t, typename index_t>
|
||||
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 <typename scalar_t, typename index_t, bool is_indexed>
|
||||
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<scalar_t, index_t>(
|
||||
/* 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<scalar_t, index_t>(
|
||||
/* 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 <typename scalar_t>
|
||||
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<scalar_t, int32_t, false>(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<scalar_t>(
|
||||
/* 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 <typename scalar_t, typename index_t>
|
||||
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<scalar_t, index_t, true>(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 <typename scalar_t, typename index_t>
|
||||
template <typename scalar_t, typename index_t, bool is_indexed>
|
||||
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<scalar_t, index_t>(
|
||||
/* 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 <typename scalar_t>
|
||||
void pack_vnni2(scalar_t* __restrict__ dst, const scalar_t* __restrict__ src, int K, int N, int ld_src, int ld_dst) {
|
||||
pack_vnni2<scalar_t, int32_t, false>(dst, src, nullptr, K, N, ld_src, ld_dst);
|
||||
}
|
||||
|
||||
template <typename scalar_t, typename index_t>
|
||||
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<scalar_t, index_t, true>(dst, src, ind, K, N, ld_src, ld_dst);
|
||||
}
|
||||
|
||||
} // anonymous namespace
|
||||
|
||||
@@ -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__":
|
||||
|
||||
@@ -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()
|
||||
@@ -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"),
|
||||
|
||||
Reference in New Issue
Block a user