[CPU] Support FP8 KV cache (#32733)

Co-authored-by: Zheng, Beilei <beilei.zheng@intel.com>
Co-authored-by: mingfeima <mingfei.ma@intel.com>
This commit is contained in:
Xuan Liao
2026-09-02 10:20:54 +08:00
committed by GitHub
co-authored by Zheng, Beilei mingfeima
parent e874ae64cd
commit 26f760d5c0
14 changed files with 952 additions and 262 deletions
+438 -184
View File
@@ -1,6 +1,7 @@
#include "common.h"
#include "gemm.h"
#include "vec.h"
#include "vec_pack.h"
namespace {
@@ -15,11 +16,12 @@ namespace {
#if defined(CPU_CAPABILITY_AVX512)
// key: from [N, 32] to [32/2, N, 2]
// val: from [N, 32] to [N/2, 32, 2]
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
inline void pack_vnni_Nx32(
scalar_t* __restrict__ dst0,
scalar_t* __restrict__ dst1,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int N,
int ld_src,
@@ -29,7 +31,7 @@ inline void pack_vnni_Nx32(
__m512i vinputs[16];
int n = 0;
for (; n < N; ++n) {
vinputs[n] = _mm512_loadu_si512(src + ind[n] * ld_src);
mm512_load_vec(src, src_scale, ld_src, ind[n], 32, vinputs[n]);
}
// padding with zero to avoid uninitialized vectors
for (; n < 16; ++n) {
@@ -65,11 +67,12 @@ inline void pack_vnni_Nx32(
// * for key: from [N, K/2, 2] to [K/2, N, 2]
// * for value: from [N/2, 2, Kv] to [N/2, Kv, 2]
//
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
void pack_vnni(
scalar_t* __restrict__ dst0,
scalar_t* __restrict__ dst1,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int N,
int K,
@@ -86,10 +89,11 @@ void pack_vnni(
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>(
pack_vnni_Nx32<scalar_t, packed_t, index_t>(
/* dst0 */ dst0 + ((kb * 32) >> 1) * ld_dst0 * 2 + nb * 16 * 2,
/* dst1 */ dst1 + ((nb * 16) >> 1) * ld_dst1 * 2 + kb * 32 * 2,
/* src */ src + kb * 32,
/* src_scale */ src_scale,
/* ind */ ind + nb * 16,
/* N */ nb_size,
/* ld_src */ ld_src,
@@ -101,9 +105,10 @@ void pack_vnni(
#else
for (int n = 0; n < N; ++n) {
index_t index = ind[n];
float scale = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int k = 0; k < K / 2; ++k) {
for (int d = 0; d < 2; ++d) {
dst0[k * ld_dst0 * 2 + n * 2 + d] = src[index * ld_src + k * 2 + d];
dst0[k * ld_dst0 * 2 + n * 2 + d] = src[index * ld_src + k * 2 + d] * scale;
}
}
}
@@ -111,15 +116,18 @@ void pack_vnni(
for (int n = 0; n < (N >> 1) * 2; n += 2) {
index_t index0 = ind[n + 0];
index_t index1 = ind[n + 1];
float scale0 = src_scale != nullptr ? src_scale[0] : 1.0f;
float scale1 = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int k = 0; k < Kv; ++k) {
dst1[(n >> 1) * ld_dst1 * 2 + k * 2 + 0] = src[index0 * ld_src + k];
dst1[(n >> 1) * ld_dst1 * 2 + k * 2 + 1] = src[index1 * ld_src + k];
dst1[(n >> 1) * ld_dst1 * 2 + k * 2 + 0] = src[index0 * ld_src + k] * scale0;
dst1[(n >> 1) * ld_dst1 * 2 + k * 2 + 1] = src[index1 * ld_src + k] * scale1;
}
}
if (N % 2 != 0) {
index_t index = ind[N - 1];
float scale = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int k = 0; k < Kv; ++k) {
dst1[(N >> 1) * ld_dst1 * 2 + k * 2 + 0] = src[index * ld_src + k];
dst1[(N >> 1) * ld_dst1 * 2 + k * 2 + 0] = src[index * ld_src + k] * scale;
dst1[(N >> 1) * ld_dst1 * 2 + k * 2 + 1] = 0;
}
}
@@ -141,6 +149,13 @@ inline void fill_stub(scalar_t* __restrict__ out, float val, int64_t size) {
}
}
template <typename scalar_t, typename packed_t>
inline void copy_stub(scalar_t* __restrict__ out, const packed_t* __restrict__ src, int64_t size) {
for (int64_t i = 0; i < size; ++i) {
out[i] = static_cast<scalar_t>(src[i]);
}
}
template <typename scalar_t>
inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc, float s, int64_t size) {
using bVec = at::vec::Vectorized<scalar_t>;
@@ -200,11 +215,12 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu
// B : [N, K] indexed
// C : [M, N]
//
template <typename scalar_t, typename index_t, int BLOCK_M, int BLOCK_N>
template <typename scalar_t, typename packed_t, typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nt {
static inline void apply(
const scalar_t* __restrict__ A,
const scalar_t* __restrict__ B,
const packed_t* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float scale,
@@ -217,9 +233,11 @@ struct tinygemm_kernel_nt {
for (int64_t n = 0; n < BLOCK_N; ++n) {
float sum = 0.f;
int64_t b_idx = indices[n];
float b_scale = B_scale != nullptr ? B_scale[0] : 1.0f;
float new_scale = scale * b_scale;
TORCH_CHECK(b_idx < max_tokens, "token index out of scope!");
for (int64_t k = 0; k < K; ++k) {
sum += scale * static_cast<float>(A[m * lda + k]) * static_cast<float>(B[b_idx * ldb + k]);
sum += new_scale * static_cast<float>(A[m * lda + k]) * static_cast<float>(B[b_idx * ldb + k]);
}
C[m * ldc + n] = sum;
}
@@ -229,10 +247,11 @@ struct tinygemm_kernel_nt {
#if defined(CPU_CAPABILITY_AVX512)
template <typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nt<at::BFloat16, index_t, BLOCK_M, BLOCK_N> {
struct tinygemm_kernel_nt<at::BFloat16, at::BFloat16, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const at::BFloat16* __restrict__ A,
const at::BFloat16* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float scale,
@@ -306,14 +325,97 @@ struct tinygemm_kernel_nt<at::BFloat16, index_t, BLOCK_M, BLOCK_N> {
Unroll<ROWS * COLS>{}(storec);
}
};
#endif
#if defined(CPU_CAPABILITY_AVX512)
template <typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nt<at::Half, index_t, BLOCK_M, BLOCK_N> {
struct tinygemm_kernel_nt<at::BFloat16, at::Float8_e4m3fn, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const at::BFloat16* __restrict__ A,
const at::Float8_e4m3fn* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float scale,
int64_t lda,
int64_t ldb,
int64_t ldc,
int64_t K,
int64_t max_tokens) {
constexpr int ROWS = BLOCK_M;
constexpr int COLS = BLOCK_N;
__m512bh va;
__m512bh vb[COLS];
__m512 vc[ROWS * COLS];
__m512 vscales[COLS];
auto loadc = [&](auto i) { vc[i] = _mm512_setzero_ps(); };
Unroll<ROWS * COLS>{}(loadc);
// for main loop
auto compute = [&](auto i, int64_t k) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
if constexpr (col == 0) {
va = (__m512bh)(_mm512_loadu_si512(A + row * lda + k));
}
if constexpr (row == 0) {
if constexpr (col + 1 < COLS) {
int64_t b_idx_prefetch = indices[col + 1];
_mm_prefetch(B + b_idx_prefetch * ldb + k, _MM_HINT_T0);
}
int64_t b_idx = indices[col];
TORCH_CHECK(b_idx < max_tokens, "token index out of scope!");
__m256i s8 = _mm256_loadu_si256((__m256i const*)(B + b_idx * ldb + k));
vb[col] = CVT_FP8_TO_BF16_EXT(s8);
vscales[col] = _mm512_mul_ps(_mm512_set1_ps(B_scale[0] * scale), vexp);
}
vc[i] = _mm512_dpbf16_ps(vc[i], va, vb[col]);
};
// for remainder
auto compute2 = [&](auto i, int64_t k, __mmask32 mask) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
if constexpr (col == 0) {
va = (__m512bh)(_mm512_maskz_loadu_epi16(mask, A + row * lda + k));
}
if constexpr (row == 0) {
int64_t b_idx = indices[col];
TORCH_CHECK(b_idx < max_tokens, "token index out of scope!");
__m256i s8 = _mm256_maskz_loadu_epi8(mask, B + b_idx * ldb + k);
vb[col] = CVT_FP8_TO_BF16_EXT(s8);
vscales[col] = _mm512_mul_ps(_mm512_set1_ps(B_scale[0] * scale), vexp);
}
vc[i] = _mm512_dpbf16_ps(vc[i], va, vb[col]);
};
int64_t k = 0;
for (; k <= K - 32; k += 32) {
Unroll<ROWS * COLS>{}(compute, k);
}
int64_t count = K - k;
if (count > 0) {
__mmask32 mask = (1ULL << count) - 1;
Unroll<ROWS * COLS>{}(compute2, k, mask);
}
auto storec = [&](auto i) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
C[row * ldc + col] = _mm512_reduce_add_ps(_mm512_mul_ps(vc[i], vscales[col]));
};
Unroll<ROWS * COLS>{}(storec);
}
};
template <typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nt<at::Half, at::Half, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const at::Half* __restrict__ A,
const at::Half* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float scale,
@@ -395,9 +497,19 @@ struct tinygemm_kernel_nt<at::Half, index_t, BLOCK_M, BLOCK_N> {
};
#endif
#define LAUNCH_TINYGEMM_KERNEL_NT(MB_SIZE, NB_SIZE) \
tinygemm_kernel_nt<scalar_t, index_t, MB_SIZE, NB_SIZE>::apply( \
A + mb_start * lda, B, C + mb_start * ldc + nb_start, indices + nb_start, scale, lda, ldb, ldc, K, max_tokens);
#define LAUNCH_TINYGEMM_KERNEL_NT(MB_SIZE, NB_SIZE) \
tinygemm_kernel_nt<scalar_t, packed_t, index_t, MB_SIZE, NB_SIZE>::apply( \
A + mb_start * lda, \
B, \
B_scale, \
C + mb_start * ldc + nb_start, \
indices + nb_start, \
scale, \
lda, \
ldb, \
ldc, \
K, \
max_tokens);
// this is used when N isn't multiple of 16,
// N corresponds to `head_size_v` which should be 16x
@@ -405,6 +517,7 @@ template <typename scalar_t, typename index_t>
inline void tinygemm_kernel_nn_scalar(
const float* __restrict__ A,
const scalar_t* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
const float* __restrict__ scale,
@@ -420,8 +533,9 @@ inline void tinygemm_kernel_nn_scalar(
C[m * ldc + n] *= scale[m];
for (int64_t k = 0; k < K; ++k) {
int64_t b_idx = indices[k];
float b_scale = B_scale != nullptr ? B_scale[0] : 1.0f;
TORCH_CHECK(b_idx < max_tokens, "token index out of scope!");
C[m * ldc + n] += A[m * lda + k] * static_cast<float>(B[b_idx * ldb + n]);
C[m * ldc + n] += A[m * lda + k] * static_cast<float>(B[b_idx * ldb + n]) * b_scale;
}
}
}
@@ -437,6 +551,7 @@ struct tinygemm_kernel_nn {
static inline void apply(
const float* __restrict__ A,
const scalar_t* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
const float* __restrict__ scale,
@@ -445,7 +560,7 @@ struct tinygemm_kernel_nn {
int64_t ldc,
int64_t K,
int64_t max_tokens) {
tinygemm_kernel_nn_scalar(A, B, C, indices, scale, BLOCK_M, BLOCK_N, K, lda, ldb, ldc, max_tokens);
tinygemm_kernel_nn_scalar(A, B, B_scale, C, indices, scale, BLOCK_M, BLOCK_N, K, lda, ldb, ldc, max_tokens);
}
};
@@ -455,6 +570,7 @@ struct tinygemm_kernel_nn<at::BFloat16, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const float* __restrict__ A,
const at::BFloat16* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
const float* __restrict__ scale,
@@ -528,14 +644,100 @@ struct tinygemm_kernel_nn<at::BFloat16, index_t, BLOCK_M, BLOCK_N> {
Unroll<ROWS * COLS>{}(storec);
}
};
#endif
#if defined(CPU_CAPABILITY_AVX512)
template <typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nn<at::Float8_e4m3fn, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const float* __restrict__ A,
const at::Float8_e4m3fn* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
const float* __restrict__ scale,
int64_t lda,
int64_t ldb,
int64_t ldc,
int64_t K,
int64_t max_tokens) {
constexpr int ROWS = BLOCK_M;
constexpr int COLS = BLOCK_N / 16;
__m512 va;
__m512 vb[COLS];
__m512 vc[ROWS * COLS];
__m512 vscale;
auto loadc = [&](auto i) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
#pragma GCC diagnostic push
#pragma GCC diagnostic ignored "-Warray-bounds"
if constexpr (col == 0) {
vscale = _mm512_set1_ps(scale[row]);
}
#pragma GCC diagnostic pop
vc[i] = _mm512_loadu_ps(C + row * ldc + col * 16);
vc[i] = _mm512_mul_ps(vc[i], vscale);
};
Unroll<ROWS * COLS>{}(loadc);
auto compute = [&](auto i, int64_t k) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
if constexpr (col == 0) {
va = _mm512_set1_ps(A[row * lda + k]);
}
if constexpr (row == 0) {
if (k + 1 < K) {
int64_t b_idx_prefetch = indices[k + 1];
_mm_prefetch(B + b_idx_prefetch * ldb + col * 16, _MM_HINT_T0);
}
int64_t b_idx = indices[k];
TORCH_CHECK(b_idx < max_tokens, "token index out of scope!");
// for COLS = 2, 4, 6, 8 use 512 bit load
// for COLS = 1, 3, 5, 7 use 256 bit load
if constexpr (COLS % 2 == 0) {
if constexpr (col % 2 == 0) {
const __m512 b_scale = _mm512_mul_ps(_mm512_set1_ps(B_scale[0]), vexp);
__m256i s8 = _mm256_loadu_si256((__m256i const*)(B + b_idx * ldb + col * 16));
__m512bh bf16 = CVT_FP8_TO_BF16_EXT(s8);
__m512 f_lo = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32((__m512i)bf16, 0));
__m512 f_hi = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32((__m512i)bf16, 1));
vb[col + 0] = _mm512_mul_ps(f_lo, b_scale);
vb[col + 1] = _mm512_mul_ps(f_hi, b_scale);
}
} else {
const __m512 b_scale = _mm512_mul_ps(_mm512_set1_ps(B_scale[0]), vexp);
__m256i s8 = _mm256_loadu_si256((__m256i const*)(B + b_idx * ldb + col * 16));
__m512bh bf16 = CVT_FP8_TO_BF16_EXT(s8);
__m512 f_lo = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32((__m512i)bf16, 0));
vb[col] = _mm512_mul_ps(f_lo, b_scale);
}
}
vc[i] = _mm512_fmadd_ps(va, vb[col], vc[i]);
};
for (int64_t k = 0; k < K; ++k) {
Unroll<ROWS * COLS>{}(compute, k);
}
auto storec = [&](auto i) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
_mm512_storeu_ps(C + row * ldc + col * 16, vc[i]);
};
Unroll<ROWS * COLS>{}(storec);
}
};
template <typename index_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nn<at::Half, index_t, BLOCK_M, BLOCK_N> {
static inline void apply(
const float* __restrict__ A,
const at::Half* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
const float* __restrict__ scale,
@@ -615,6 +817,7 @@ struct tinygemm_kernel_nn<at::Half, index_t, BLOCK_M, BLOCK_N> {
tinygemm_kernel_nn<scalar_t, index_t, MB_SIZE, NB_SIZE>::apply( \
A + mb_start * lda, \
B + nb_start, \
B_scale, \
C + mb_start * ldc + nb_start, \
indices, \
scale + mb_start, \
@@ -624,10 +827,11 @@ struct tinygemm_kernel_nn<at::Half, index_t, BLOCK_M, BLOCK_N> {
K, \
max_tokens);
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
void index_gemm_kernel_nt(
const scalar_t* __restrict__ A,
const scalar_t* __restrict__ B,
const packed_t* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float scale,
@@ -782,6 +986,7 @@ template <typename scalar_t, typename index_t>
void index_gemm_kernel_nn(
const float* __restrict__ A,
const scalar_t* __restrict__ B,
const float* __restrict__ B_scale,
float* __restrict__ C,
const index_t* __restrict__ indices,
float* __restrict__ scale,
@@ -794,7 +999,7 @@ void index_gemm_kernel_nn(
int64_t max_tokens) {
constexpr int kVecSize = 16;
if ((N & (kVecSize - 1)) != 0) {
tinygemm_kernel_nn_scalar(A, B, C, indices, scale, M, N, K, lda, ldb, ldc, max_tokens);
tinygemm_kernel_nn_scalar(A, B, B_scale, C, indices, scale, M, N, K, lda, ldb, ldc, max_tokens);
return;
}
@@ -936,12 +1141,14 @@ void index_gemm_kernel_nn(
}
}
template <typename scalar_t>
template <typename scalar_t, typename packed_t>
void decode_set_kv_buffer(
scalar_t* __restrict__ k_buffer,
scalar_t* __restrict__ v_buffer,
packed_t* __restrict__ k_buffer,
packed_t* __restrict__ v_buffer,
const scalar_t* __restrict__ key,
const scalar_t* __restrict__ value,
float k_buf_scale,
float v_buf_scale,
const int64_t* __restrict__ loc,
int64_t batches,
int64_t num_heads_kv,
@@ -962,13 +1169,25 @@ void decode_set_kv_buffer(
for (int64_t i = begin; i < end; i++) {
int64_t loc_val = loc[bs];
scalar_t* k_buffer_ptr = k_buffer + loc_val * k_strideN + head_kv_id * k_strideH;
packed_t* k_buffer_ptr = k_buffer + loc_val * k_strideN + head_kv_id * k_strideH;
const scalar_t* new_key_ptr = key + bs * nk_strideN + head_kv_id * nk_strideH;
copy_stub<scalar_t>(k_buffer_ptr, new_key_ptr, head_size);
if constexpr (std::is_same_v<packed_t, at::Float8_e4m3fn>) {
for (int64_t d = 0; d < head_size; ++d) {
k_buffer_ptr[d] = static_cast<packed_t>(static_cast<float>(new_key_ptr[d]) / k_buf_scale);
}
} else {
copy_stub(k_buffer_ptr, new_key_ptr, head_size);
}
if (!is_mla) {
scalar_t* v_buffer_ptr = v_buffer + loc_val * v_strideN + head_kv_id * v_strideH;
packed_t* v_buffer_ptr = v_buffer + loc_val * v_strideN + head_kv_id * v_strideH;
const scalar_t* new_value_ptr = value + bs * nv_strideN + head_kv_id * nv_strideH;
copy_stub<scalar_t>(v_buffer_ptr, new_value_ptr, head_size_v);
if constexpr (std::is_same_v<packed_t, at::Float8_e4m3fn>) {
for (int64_t d = 0; d < head_size_v; ++d) {
v_buffer_ptr[d] = static_cast<packed_t>(static_cast<float>(new_value_ptr[d]) / v_buf_scale);
}
} else {
copy_stub(v_buffer_ptr, new_value_ptr, head_size_v);
}
}
// move to the next index
@@ -1036,13 +1255,15 @@ void decode_accumulate_kv_splits(
});
}
template <typename scalar_t, typename index_t, int64_t BLOCK_N>
template <typename scalar_t, typename packed_t, typename index_t, int64_t BLOCK_N>
void decode_attention_kernel_impl(
scalar_t* __restrict__ output,
float* __restrict__ attn_logits,
const scalar_t* __restrict__ query,
const scalar_t* __restrict__ k_buffer,
const scalar_t* __restrict__ v_buffer,
const packed_t* __restrict__ k_buffer,
const packed_t* __restrict__ v_buffer,
const float* __restrict__ k_scale,
const float* __restrict__ v_scale,
const index_t* __restrict__ req_to_token,
const int64_t* __restrict__ req_pool_indices,
const int64_t* __restrict__ seq_lens,
@@ -1117,9 +1338,10 @@ void decode_attention_kernel_impl(
int64_t n_size = std::min(BLOCK_N, kv_end - n);
// calculate s_i <- scale * Q @ K
index_gemm_kernel_nt<scalar_t, index_t>(
index_gemm_kernel_nt<scalar_t, packed_t, index_t>(
/* A */ q_ptr,
/* B */ k_buffer + head_id * k_strideH,
/* B_scale */ k_scale,
/* C */ s_i,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* scl */ sm_scale,
@@ -1157,9 +1379,10 @@ void decode_attention_kernel_impl(
m_prime = m_i;
// calculate V' <- s_delta @ V + V' * m_delta
index_gemm_kernel_nn<scalar_t, index_t>(
index_gemm_kernel_nn<packed_t, index_t>(
/* A */ s_delta,
/* B */ v_buffer + head_id * v_strideH,
/* B_scale */ v_scale,
/* C */ v_prime,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* scl */ &m_delta,
@@ -1191,13 +1414,15 @@ void decode_attention_kernel_impl(
output, attn_logits, sinks, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2, has_sink);
} // MHA
template <typename scalar_t, typename index_t, int64_t BLOCK_N>
template <typename scalar_t, typename packed_t, typename index_t, int64_t BLOCK_N>
void decode_attention_mla_kernel_impl(
scalar_t* __restrict__ output,
float* __restrict__ attn_logits,
const scalar_t* __restrict__ query,
const scalar_t* __restrict__ k_buffer,
const scalar_t* __restrict__ v_buffer,
const packed_t* __restrict__ k_buffer,
const packed_t* __restrict__ v_buffer,
const float* __restrict__ k_buf_scale,
const float* __restrict__ v_buf_scale,
const index_t* __restrict__ req_to_token,
const int64_t* __restrict__ req_pool_indices,
const int64_t* __restrict__ seq_lens,
@@ -1289,10 +1514,11 @@ void decode_attention_mla_kernel_impl(
const int64_t padded_n_size = div_up(int(n_size), TILE_K) * TILE_K;
// get key and pack
pack_vnni<scalar_t, index_t>(
pack_vnni<scalar_t, packed_t, index_t>(
/* dst0 */ Btmp0,
/* dst1 */ Btmp1,
/* src */ k_buffer + /* head_kv_id */ 0 * k_strideH,
/* src_scale */ k_buf_scale,
/* ind */ req_to_token + req_pool_id * max_context_len + n,
/* N */ n_size,
/* K */ head_size,
@@ -1389,13 +1615,15 @@ void decode_attention_mla_kernel_impl(
output, attn_logits, sinks, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2, has_sink);
} // MLA
template <typename scalar_t, typename index_t, int64_t BLOCK_N>
template <typename scalar_t, typename packed_t, typename index_t, int64_t BLOCK_N>
void decode_attention_grouped_kernel_impl(
scalar_t* __restrict__ output,
float* __restrict__ attn_logits,
const scalar_t* __restrict__ query,
const scalar_t* __restrict__ k_buffer,
const scalar_t* __restrict__ v_buffer,
const packed_t* __restrict__ k_buffer,
const packed_t* __restrict__ v_buffer,
const float* __restrict__ k_scale,
const float* __restrict__ v_scale,
const index_t* __restrict__ req_to_token,
const int64_t* __restrict__ req_pool_indices,
const int64_t* __restrict__ seq_lens,
@@ -1489,9 +1717,10 @@ void decode_attention_grouped_kernel_impl(
int64_t n_size = std::min(BLOCK_N, kv_end - n);
// calculate Q @ K
index_gemm_kernel_nt<scalar_t, index_t>(
index_gemm_kernel_nt<scalar_t, packed_t, index_t>(
/* A */ q_ptr,
/* B */ k_buffer + head_kv_id * k_strideH,
/* B_scale */ k_scale,
/* C */ s_i,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* scl */ sm_scale,
@@ -1533,9 +1762,10 @@ void decode_attention_grouped_kernel_impl(
}
// calculate V' <- s_delta @ V + V' * m_delta
index_gemm_kernel_nn<scalar_t, index_t>(
index_gemm_kernel_nn<packed_t, index_t>(
/* A */ s_delta,
/* B */ v_buffer + head_kv_id * v_strideH,
/* B_scale */ v_scale,
/* C */ v_prime,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* scl */ m_delta,
@@ -1587,6 +1817,8 @@ void decode_attention_cpu(
at::Tensor& query,
at::Tensor& k_buffer,
at::Tensor& v_buffer,
double k_buf_scale,
double v_buf_scale,
at::Tensor& output,
const std::optional<at::Tensor>& key,
const std::optional<at::Tensor>& value,
@@ -1653,13 +1885,19 @@ void decode_attention_cpu(
void* v_buffer_data = v_buffer.data_ptr();
const bool is_mla = (k_buffer_data == v_buffer_data) && (num_heads_kv == 1) && (head_size == head_size_v + 64);
auto kv_dtype = k_buffer.scalar_type();
if (kv_dtype == at::ScalarType::Float8_e4m3fn) {
TORCH_CHECK(v_buffer.scalar_type() == kv_dtype, "k_buffer and v_buffer should have the same dtype");
TORCH_CHECK(k_buf_scale > 0.0 && v_buf_scale > 0.0, "Float8 kv_buffer requires positive static scales");
TORCH_CHECK(!is_mla, "Float8 kv_buffer is only supported for MHA on CPU");
}
// block length for k_buffer and v_buffer
constexpr int BLOCK_N = 256;
// buffer for packing k_cache and v_cache
int num_threads = at::get_num_threads();
int64_t size_per_thread = is_mla ? BLOCK_N * head_size + BLOCK_N * head_size_v : 0;
auto buffer = at::empty({num_threads, size_per_thread}, k_buffer.options());
bool has_encoder_lens = encoder_lens.has_value();
// Since encoder_lens is not used when it is None, encoder_lens_t can be initialized as any tensor of int64_t dtype.
at::Tensor encoder_lens_t = seq_lens;
@@ -1668,147 +1906,163 @@ void decode_attention_cpu(
CHECK_EQ(encoder_lens_t.size(0), num_seqs);
}
bool has_sink = sinks.has_value();
float k_buf_scale_float = static_cast<float>(k_buf_scale);
float v_buf_scale_float = static_cast<float>(v_buf_scale);
const float* k_buf_scale_ptr = kv_dtype == at::ScalarType::Float8_e4m3fn ? &k_buf_scale_float : nullptr;
const float* v_buf_scale_ptr = kv_dtype == at::ScalarType::Float8_e4m3fn ? &v_buf_scale_float : nullptr;
at::Tensor sinks_tensor = has_sink ? sinks.value() : at::empty({num_heads}, query.options());
CHECK_DIM(1, sinks_tensor);
CHECK_EQ(sinks_tensor.size(0), num_heads);
AT_DISPATCH_REDUCED_FLOATING_TYPES(query.scalar_type(), "decode_attention_kernel", [&] {
AT_DISPATCH_INDEX_TYPES(index_dtype, "decode_attention_indices", [&] {
if (key.has_value()) {
TORCH_CHECK(value.has_value(), "key and value should have values at the same time")
CHECK_EQ(loc.numel(), num_seqs);
auto key_tensor = key.value();
auto value_tensor = value.value();
// for MLA, key and value shares the same storage and value could be non-contiguous
CHECK_LAST_DIM_CONTIGUOUS_INPUT(key_tensor);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(value_tensor);
CHECK_DIM(3, key_tensor);
CHECK_DIM(3, value_tensor);
// strides for new key and value
int64_t nk_strideN = key_tensor.stride(0);
int64_t nk_strideH = key_tensor.stride(1);
int64_t nv_strideN = value_tensor.stride(0);
int64_t nv_strideH = value_tensor.stride(1);
// update the kv buffer
decode_set_kv_buffer(
(scalar_t*)k_buffer_data,
(scalar_t*)v_buffer_data,
key_tensor.data_ptr<scalar_t>(),
value_tensor.data_ptr<scalar_t>(),
loc.data_ptr<int64_t>(),
num_seqs,
num_heads_kv,
head_size,
head_size_v,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
nk_strideN,
nk_strideH,
nv_strideN,
nv_strideH,
is_mla);
}
CPU_DISPATCH_PACKED_TYPES(k_buffer.scalar_type(), "decode_attention_packed_types", [&] {
if (key.has_value()) {
TORCH_CHECK(value.has_value(), "key and value should have values at the same time")
CHECK_EQ(loc.numel(), num_seqs);
auto key_tensor = key.value();
auto value_tensor = value.value();
// for MLA, key and value shares the same storage and value could be non-contiguous
CHECK_LAST_DIM_CONTIGUOUS_INPUT(key_tensor);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(value_tensor);
CHECK_DIM(3, key_tensor);
CHECK_DIM(3, value_tensor);
// strides for new key and value
int64_t nk_strideN = key_tensor.stride(0);
int64_t nk_strideH = key_tensor.stride(1);
int64_t nv_strideN = value_tensor.stride(0);
int64_t nv_strideH = value_tensor.stride(1);
// update the kv buffer
decode_set_kv_buffer(
(packed_t*)k_buffer_data,
(packed_t*)v_buffer_data,
key_tensor.data_ptr<scalar_t>(),
value_tensor.data_ptr<scalar_t>(),
k_buf_scale_float,
v_buf_scale_float,
loc.data_ptr<int64_t>(),
num_seqs,
num_heads_kv,
head_size,
head_size_v,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
nk_strideN,
nk_strideH,
nv_strideN,
nv_strideH,
is_mla);
}
if (num_heads == num_heads_kv) {
// MHA
decode_attention_kernel_impl<scalar_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const scalar_t*)k_buffer_data,
(const scalar_t*)v_buffer_data,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
encoder_lens_t.data_ptr<int64_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
sliding_window_size,
is_cross_attn,
has_encoder_lens,
has_sink);
} else if (is_mla) {
// MLA
decode_attention_mla_kernel_impl<scalar_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const scalar_t*)k_buffer_data,
(const scalar_t*)v_buffer_data,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
buffer.data_ptr<scalar_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
size_per_thread,
has_sink);
} else {
// GQA/MQA
decode_attention_grouped_kernel_impl<scalar_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const scalar_t*)k_buffer_data,
(const scalar_t*)v_buffer_data,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
encoder_lens_t.data_ptr<int64_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
num_heads_kv,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
sliding_window_size,
is_cross_attn,
has_encoder_lens,
has_sink);
}
if (num_heads == num_heads_kv) {
// MHA
decode_attention_kernel_impl<scalar_t, packed_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const packed_t*)k_buffer_data,
(const packed_t*)v_buffer_data,
k_buf_scale_ptr,
v_buf_scale_ptr,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
encoder_lens_t.data_ptr<int64_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
sliding_window_size,
is_cross_attn,
has_encoder_lens,
has_sink);
} else if (is_mla) {
// MLA
TORCH_CHECK(key.has_value() && value.has_value(), "MLA requires key and value tensors");
auto buffer = at::empty({num_threads, size_per_thread}, key.value().options());
decode_attention_mla_kernel_impl<scalar_t, packed_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const packed_t*)k_buffer_data,
(const packed_t*)v_buffer_data,
k_buf_scale_ptr,
v_buf_scale_ptr,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
buffer.data_ptr<scalar_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
size_per_thread,
has_sink);
} else {
// GQA/MQA
decode_attention_grouped_kernel_impl<scalar_t, packed_t, index_t, BLOCK_N>(
output.data_ptr<scalar_t>(),
attn_logits.data_ptr<float>(),
query.data_ptr<scalar_t>(),
(const packed_t*)k_buffer_data,
(const packed_t*)v_buffer_data,
k_buf_scale_ptr,
v_buf_scale_ptr,
req_to_token.data_ptr<index_t>(),
req_pool_indices.data_ptr<int64_t>(),
seq_lens.data_ptr<int64_t>(),
encoder_lens_t.data_ptr<int64_t>(),
sinks_tensor.data_ptr<scalar_t>(),
num_seqs,
num_heads,
num_heads_kv,
head_size,
head_size_v,
num_kv_splits,
q_strideM,
q_strideH,
k_strideN,
k_strideH,
v_strideN,
v_strideH,
sm_scale,
logit_cap,
max_num_reqs,
max_context_len,
max_total_num_tokens,
sliding_window_size,
is_cross_attn,
has_encoder_lens,
has_sink);
}
});
});
});
}
+36 -18
View File
@@ -22,14 +22,16 @@ namespace {
// plain causal mask (correct for non-spec extend and topk == 1 chains).
//
template <typename scalar_t, typename index_t, int BLOCK_M, int BLOCK_N>
template <typename scalar_t, typename packed_t, typename index_t, int BLOCK_M, int BLOCK_N>
void extend_attention_kernel_impl(
scalar_t* __restrict__ o_extend,
const scalar_t* __restrict__ q_extend,
const scalar_t* __restrict__ k_extend,
const scalar_t* __restrict__ v_extend,
const scalar_t* __restrict__ k_buffer,
const scalar_t* __restrict__ v_buffer,
const packed_t* __restrict__ k_buffer,
const packed_t* __restrict__ v_buffer,
const float* __restrict__ k_buf_scale,
const float* __restrict__ v_buf_scale,
const index_t* __restrict__ req_to_token,
const int64_t* __restrict__ req_pool_indices,
const int64_t* __restrict__ seq_lens,
@@ -160,9 +162,10 @@ 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, packed_t, index_t>(
/* dst */ Btmp,
/* src */ k_buffer + head_kv_id * k_strideH,
/* src_scale*/ k_buf_scale,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* N */ n_size,
/* K */ head_size,
@@ -214,9 +217,10 @@ void extend_attention_kernel_impl(
}
// get value and pack
pack_vnni2<scalar_t>(
pack_vnni2<scalar_t, packed_t, index_t>(
/* dst */ Btmp,
/* src */ v_buffer + head_kv_id * v_strideH,
/* src_scale*/ v_buf_scale,
/* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset,
/* K */ n_size,
/* N */ head_size_v,
@@ -382,13 +386,15 @@ inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int
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>( \
extend_attention_kernel_impl<scalar_t, packed_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>(), \
k_buffer.data_ptr<packed_t>(), \
v_buffer.data_ptr<packed_t>(), \
k_buf_scale_ptr, \
v_buf_scale_ptr, \
req_to_token.data_ptr<index_t>(), \
req_pool_indices.data_ptr<int64_t>(), \
seq_lens.data_ptr<int64_t>(), \
@@ -453,6 +459,8 @@ void extend_attention_cpu(
at::Tensor& o_extend,
at::Tensor& k_buffer,
at::Tensor& v_buffer,
double k_buf_scale,
double v_buf_scale,
at::Tensor& req_to_token,
at::Tensor& req_pool_indices,
at::Tensor& seq_lens,
@@ -536,8 +544,16 @@ void extend_attention_cpu(
// D and DV need to be 32x as we transpose by 512-bit
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);
auto kv_dtype = k_buffer.scalar_type();
if (kv_dtype == at::ScalarType::Float8_e4m3fn) {
TORCH_CHECK(v_buffer.scalar_type() == kv_dtype, "k_buffer and v_buffer should have same data type");
TORCH_CHECK(k_buf_scale > 0.0 && v_buf_scale > 0.0, "float8 static scales must be positive");
}
int num_threads = at::get_num_threads();
float k_buf_scale_float = static_cast<float>(k_buf_scale);
float v_buf_scale_float = static_cast<float>(v_buf_scale);
const float* k_buf_scale_ptr = kv_dtype == at::ScalarType::Float8_e4m3fn ? &k_buf_scale_float : nullptr;
const float* v_buf_scale_ptr = kv_dtype == at::ScalarType::Float8_e4m3fn ? &v_buf_scale_float : nullptr;
auto buffer = at::empty({}, q_extend.options().dtype(at::kChar));
bool has_encoder_lens = encoder_lens.has_value();
@@ -575,15 +591,17 @@ void extend_attention_cpu(
AT_DISPATCH_REDUCED_FLOATING_TYPES(q_extend.scalar_type(), "extend_attention_kernel", [&] {
AT_DISPATCH_INDEX_TYPES(index_dtype, "extend_attention_indices", [&] {
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);
}
CPU_DISPATCH_PACKED_TYPES(k_buffer.scalar_type(), "extend_attention_packed_types", [&] {
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);
}
});
});
});
}
@@ -211,6 +211,8 @@ void decode_attention_cpu(
at::Tensor& query,
at::Tensor& k_cache,
at::Tensor& v_cache,
double k_cache_scale,
double v_cache_scale,
at::Tensor& output,
const std::optional<at::Tensor>& key,
const std::optional<at::Tensor>& value,
@@ -233,6 +235,8 @@ void extend_attention_cpu(
at::Tensor& o_extend,
at::Tensor& k_buffer,
at::Tensor& v_buffer,
double k_buf_scale,
double v_buf_scale,
at::Tensor& req_to_token,
at::Tensor& req_pool_indices,
at::Tensor& seq_lens,
@@ -708,8 +712,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
// decode
m.def(
"decode_attention_cpu(Tensor query, Tensor k_cache, Tensor v_cahce, Tensor(a!) output, Tensor? key, Tensor? "
"value, "
"decode_attention_cpu(Tensor query, Tensor k_cache, Tensor v_cahce, float k_cache_scale, float "
"v_cache_scale, Tensor(a!) output, Tensor? key, Tensor? value, "
"Tensor loc, Tensor attn_logits, Tensor req_to_token, Tensor req_pool_indices, Tensor seq_lens, float sm_scale, "
"float logit_cap, bool is_cross_attn, int sliding_window_size, Tensor? encoder_lens, Tensor? sinks) -> ()");
m.impl("decode_attention_cpu", torch::kCPU, &decode_attention_cpu);
@@ -717,7 +721,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
// extend
m.def(
"extend_attention_cpu(Tensor q_extend, Tensor? k_extend, Tensor? v_extend, Tensor(a!) o_extend, Tensor k_buffer, "
"Tensor v_buffer, Tensor req_to_token, Tensor req_pool_indices, Tensor seq_lens, Tensor extend_seq_lens, Tensor "
"Tensor v_buffer, float k_buf_scale, float v_buf_scale, Tensor req_to_token, Tensor req_pool_indices, Tensor "
"seq_lens, Tensor extend_seq_lens, Tensor "
"extend_start_loc, int max_len_extend, float sm_scale, float logit_cap, bool is_cross_attn, int "
"sliding_window_size, Tensor? "
"encoder_lens, Tensor? sinks, Tensor? tree_mask=None, bool is_causal=True) -> ()");
+103 -35
View File
@@ -13,11 +13,44 @@ inline index_t get_index(index_t* ind, int i) {
}
#if defined(CPU_CAPABILITY_AVX512)
inline __mmask32 elem_mask(int n) {
return n >= 32 ? 0xFFFFFFFFu : static_cast<__mmask32>((1u << n) - 1);
}
template <typename scalar_t>
inline void mm512_load_vec(
const scalar_t* __restrict__ src, const float* /*src_scale*/, int64_t ld, int64_t index, int n, __m512i& dst) {
const scalar_t* p = src + index * ld;
dst = n >= 32 ? _mm512_loadu_si512(p) : _mm512_maskz_loadu_epi16(elem_mask(n), p);
}
const __m512 vexp = _mm512_castsi512_ps(_mm512_set1_epi32(kFP8_BIAS));
inline void mm512_load_vec(
const at::Float8_e4m3fn* __restrict__ src,
const float* __restrict__ scale,
int64_t ld,
int64_t index,
int n,
__m512i& dst) {
const __m512 s = _mm512_set1_ps(scale[0]);
const auto* p = src + index * ld;
__m256i s8 =
n >= 32 ? _mm256_loadu_si256(reinterpret_cast<const __m256i*>(p)) : _mm256_maskz_loadu_epi8(elem_mask(n), p);
// TODO: optimize the process of converting fp8 to bf16: fp8 -> bf16 -> fp32 -> bf16
__m512bh bf16 = cvt_e4m3_bf16_intrinsic_with_denorm(s8);
__m512 f_lo = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32((__m512i)bf16, 0));
__m512 f_hi = CVT_BF16_TO_FP32(_mm512_extracti32x8_epi32((__m512i)bf16, 1));
dst = (__m512i)_mm512_cvtne2ps_pbh(_mm512_mul_ps(f_hi, s), _mm512_mul_ps(f_lo, s));
}
// key: from [N, 32] to [32/2, N, 2]
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
inline void pack_vnni_Nx32(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int N,
int ld_src,
@@ -27,7 +60,7 @@ inline void pack_vnni_Nx32(
int n = 0;
for (; n < N; ++n) {
index_t index = get_index(ind, n);
vinputs[n] = _mm512_loadu_si512(src + index * ld_src);
mm512_load_vec(src, src_scale, ld_src, index, 32, vinputs[n]);
}
// padding with zero to avoid uninitialized vectors
for (; n < 16; ++n) {
@@ -43,10 +76,11 @@ inline void pack_vnni_Nx32(
}
}
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
inline void pack_vnni_N_remainder(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int N,
int K,
@@ -55,12 +89,10 @@ inline void pack_vnni_N_remainder(
__m512i vinputs[16];
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);
mm512_load_vec(src, src_scale, ld_src, index, K, vinputs[n]);
}
// padding with zero to avoid uninitialized vectors
for (; n < 16; ++n) {
@@ -77,10 +109,11 @@ inline void pack_vnni_N_remainder(
}
// value: from [K, 32] to [K/2, 32, 2]
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
inline void pack_vnni_Kx32(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int K,
int ld_src,
@@ -90,7 +123,7 @@ inline void pack_vnni_Kx32(
int k = 0;
for (; k < K; ++k) {
index_t index = get_index(ind, k);
vinputs[k] = _mm512_loadu_si512(src + index * ld_src);
mm512_load_vec(src, src_scale, ld_src, index, 32, vinputs[k]);
}
// padding with zero to avoid uninitialized vectors
for (; k < 2; ++k) {
@@ -104,10 +137,11 @@ inline void pack_vnni_Kx32(
_mm512_storeu_si512(dst + 0 * ld_dst * 2 + 32, d1);
}
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t, typename index_t>
inline void pack_vnni_K_remainder(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int K,
int N,
@@ -115,12 +149,10 @@ inline void pack_vnni_K_remainder(
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);
mm512_load_vec(src, src_scale, ld_src, index, N, vinputs[k]);
}
// padding with zero to avoid uninitialized vectors
for (; k < 2; ++k) {
@@ -146,10 +178,11 @@ inline void pack_vnni_K_remainder(
// 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>
template <typename scalar_t, typename packed_t, typename index_t, bool is_indexed>
void pack_vnni(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int N,
int K,
@@ -164,18 +197,20 @@ void pack_vnni(
int nb_size = std::min(N - nb * 16, 16);
for (int kb = 0; kb < KB; ++kb) {
// handle 16x512bits each block
pack_vnni_Nx32<scalar_t, index_t>(
pack_vnni_Nx32<scalar_t, packed_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),
/* src_scale*/ src_scale,
/* ind */ is_indexed ? ind + nb * 16 : nullptr,
/* N */ nb_size,
/* ld_src */ ld_src,
/* ld_dst */ ld_dst);
}
if (K_remainder > 0) {
pack_vnni_N_remainder<scalar_t, index_t>(
pack_vnni_N_remainder<scalar_t, packed_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),
/* src_scale */ src_scale,
/* ind */ is_indexed ? ind + nb * 16 : nullptr,
/* N */ nb_size,
/* K */ K_remainder,
@@ -186,9 +221,10 @@ void pack_vnni(
#else
for (int n = 0; n < N; ++n) {
index_t index = get_index(ind, n);
float scale = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int k = 0; k < K / 2; ++k) {
for (int d = 0; d < 2; ++d) {
dst[k * ld_dst * 2 + n * 2 + d] = src[index * ld_src + k * 2 + d];
dst[k * ld_dst * 2 + n * 2 + d] = src[index * ld_src + k * 2 + d] * scale;
}
}
}
@@ -197,28 +233,42 @@ void pack_vnni(
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) {
pack_vnni<scalar_t, int32_t, false>(dst, src, nullptr, N, K, ld_src, ld_dst);
pack_vnni<scalar_t, scalar_t, int32_t, false>(dst, src, nullptr, nullptr, N, K, ld_src, ld_dst);
}
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t>
void pack_vnni(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
int N,
int K,
int ld_src,
int ld_dst) {
pack_vnni<scalar_t, packed_t, int32_t, false>(dst, src, src_scale, nullptr, N, K, ld_src, ld_dst);
}
template <typename scalar_t, typename packed_t, typename index_t>
void pack_vnni(
scalar_t* __restrict__ dst,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
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);
pack_vnni<scalar_t, packed_t, index_t, true>(dst, src, src_scale, 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, bool is_indexed>
template <typename scalar_t, typename packed_t, typename index_t, bool is_indexed>
void pack_vnni2(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
const index_t* __restrict__ ind,
int K,
int N,
@@ -233,9 +283,10 @@ void pack_vnni2(
int kb_size = std::min(K - kb * 2, 2);
for (int nb = 0; nb < NB; ++nb) {
// handle 2x512bits each block
pack_vnni_Kx32<scalar_t, index_t>(
pack_vnni_Kx32<scalar_t, packed_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,
/* src_scale */ src_scale,
/* ind */ is_indexed ? ind + kb * 2 : nullptr,
/* K */ kb_size,
/* ld_src */ ld_src,
@@ -245,6 +296,7 @@ void pack_vnni2(
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,
/* src_scale */ src_scale,
/* ind */ is_indexed ? ind + kb * 2 : nullptr,
/* K */ kb_size,
/* N */ N_remainder,
@@ -257,15 +309,18 @@ void pack_vnni2(
for (; k < (K >> 1) * 2; k += 2) {
index_t index0 = get_index(ind, k + 0);
index_t index1 = get_index(ind, k + 1);
float scale0 = src_scale != nullptr ? src_scale[0] : 1.0f;
float scale1 = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int n = 0; n < N; ++n) {
dst[(k >> 1) * ld_dst * 2 + n * 2 + 0] = src[index0 * ld_src + n];
dst[(k >> 1) * ld_dst * 2 + n * 2 + 1] = src[index1 * ld_src + n];
dst[(k >> 1) * ld_dst * 2 + n * 2 + 0] = src[index0 * ld_src + n] * scale0;
dst[(k >> 1) * ld_dst * 2 + n * 2 + 1] = src[index1 * ld_src + n] * scale1;
}
}
if (K % 2 != 0) {
index_t index = get_index(ind, K - 1);
float scale = src_scale != nullptr ? src_scale[0] : 1.0f;
for (int n = 0; n < N; ++n) {
dst[(K >> 1) * ld_dst * 2 + n * 2 + 0] = src[index * ld_src + n];
dst[(K >> 1) * ld_dst * 2 + n * 2 + 0] = src[index * ld_src + n] * scale;
dst[(K >> 1) * ld_dst * 2 + n * 2 + 1] = 0;
}
k += 2;
@@ -275,20 +330,33 @@ void pack_vnni2(
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);
pack_vnni2<scalar_t, scalar_t, int32_t, false>(dst, src, nullptr, nullptr, K, N, ld_src, ld_dst);
}
template <typename scalar_t, typename index_t>
template <typename scalar_t, typename packed_t>
void pack_vnni2(
scalar_t* __restrict__ dst,
const scalar_t* __restrict__ src,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
int K,
int N,
int ld_src,
int ld_dst) {
pack_vnni2<scalar_t, packed_t, int32_t, false>(dst, src, src_scale, nullptr, K, N, ld_src, ld_dst);
}
template <typename scalar_t, typename packed_t, typename index_t>
void pack_vnni2(
scalar_t* __restrict__ dst,
const packed_t* __restrict__ src,
const float* __restrict__ src_scale,
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);
pack_vnni2<scalar_t, packed_t, index_t, true>(dst, src, src_scale, ind, K, N, ld_src, ld_dst);
}
} // anonymous namespace
@@ -29,6 +29,7 @@ class IntelAMXAttnBackend(AttentionBackend):
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.use_mla = model_runner.use_mla_backend
self.max_context_len = model_runner.model_config.context_len
# full->SWA translated out_cache_loc, computed once per forward (the only
@@ -203,13 +204,24 @@ class IntelAMXAttnBackend(AttentionBackend):
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
key_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
value_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
if save_kv_cache and k is not None and v is not None:
# Cross-attention never writes to the SWA pool, so only thread the
# full->SWA location for non-cross-attention layers.
swa_loc = None if layer.is_cross_attention else self.swa_out_cache_loc
self.token_to_kv_pool.set_kv_buffer(
layer, KVWriteLoc(cache_loc, swa_loc), k, v
)
write_loc = KVWriteLoc(cache_loc, swa_loc)
if not self.use_mla and key_buffer.dtype == torch.float8_e4m3fn:
self.token_to_kv_pool.set_kv_buffer(
layer,
write_loc,
k,
v,
k_scale=layer.k_scale_float,
v_scale=layer.v_scale_float,
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, write_loc, k, v)
# Precomputed once per forward pass in init_forward_metadata (spec
# verify batches carry no extend_* fields; see _build_extend_metadata).
@@ -219,6 +231,8 @@ class IntelAMXAttnBackend(AttentionBackend):
if seq_lens.dtype != torch.int64:
seq_lens = seq_lens.to(torch.int64)
key_scale = layer.k_scale_float or 1.0
value_scale = layer.v_scale_float or 1.0
is_causal = True
if layer.is_cross_attention or layer.attn_type == AttentionType.ENCODER_ONLY:
is_causal = False
@@ -230,8 +244,10 @@ class IntelAMXAttnBackend(AttentionBackend):
k,
v,
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
key_buffer,
value_buffer,
key_scale,
value_scale,
self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices,
seq_lens,
@@ -283,15 +299,38 @@ class IntelAMXAttnBackend(AttentionBackend):
o = q.new_empty((q.shape[0], layer.tp_q_head_num * layer.v_head_dim))
else:
o = torch.empty_like(q)
key_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
value_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
key_scale = layer.k_scale_float or 1.0
value_scale = layer.v_scale_float or 1.0
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
if (
save_kv_cache
and k is not None
and v is not None
and key_buffer.dtype == torch.float8_e4m3fn
):
swa_loc = None if layer.is_cross_attention else self.swa_out_cache_loc
self.token_to_kv_pool.set_kv_buffer(
layer,
KVWriteLoc(cache_loc, swa_loc),
k,
v,
k_scale=layer.k_scale_float,
v_scale=layer.v_scale_float,
)
k = None
v = None
self.decode_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
key_buffer,
value_buffer,
key_scale,
value_scale,
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
k,
v,
@@ -330,6 +330,58 @@ class UnquantizedKVCacheMethod(KVCacheQuantMethodBase):
)
class CPUFP8KVCacheMethod(KVCacheQuantMethodBase):
name = "cpu_fp8_e4m3"
SCALE_BLOCK_SIZE = 1
def create_buffers(self, size, head_num, head_dim, layer_num, device) -> dict:
buffer_shape = (size, head_num, head_dim)
return {
"k_buffer": [
torch.zeros(buffer_shape, dtype=torch.float8_e4m3fn, device=device)
for _ in range(layer_num)
],
"v_buffer": [
torch.zeros(buffer_shape, dtype=torch.float8_e4m3fn, device=device)
for _ in range(layer_num)
],
"k_scale_buffer": None,
"v_scale_buffer": None,
"dq_k_buffer": None,
"dq_v_buffer": None,
"store_dtype": torch.float8_e4m3fn,
}
def quantize_and_store(
self,
k_buffer,
v_buffer,
k_scale_buffer,
v_scale_buffer,
loc,
cache_k,
cache_v,
k_scale=None,
v_scale=None,
) -> None:
k_scale = 1.0 if k_scale is None else k_scale
v_scale = 1.0 if v_scale is None else v_scale
k_buffer[loc] = (cache_k / k_scale).to(torch.float8_e4m3fn)
v_buffer[loc] = (cache_v / v_scale).to(torch.float8_e4m3fn)
def dequantize_prev_kv(
self, k_fp8, k_scales, v_fp8, v_scales, layer_id
) -> tuple[Tensor, Tensor]:
raise NotImplementedError(
"CPU FP8 KV cache is consumed directly by the CPU attention kernels."
)
def compute_cell_size(
self, head_num: int, head_dim: int, num_layers: int, kv_size: int
) -> int:
return head_num * head_dim * num_layers * kv_size * 2
class NVFP4KVCacheMethod(KVCacheQuantMethodBase):
"""NVFP4 two-level scaling: global FP32 + per-block FP8 E4M3.
@@ -717,6 +769,7 @@ _FP4_MX_MHA_BACKENDS = frozenset(
{"triton", "torch_native", "flex_attention", "trtllm_mha"}
)
_FP4_MX_PREFILL_BACKENDS = _FP4_MX_MHA_BACKENDS | frozenset({"fa4"})
_CPU_FP8_BACKENDS = frozenset({"intel_amx"})
def _backend_matcher(backends) -> KVCacheBackendMatcher:
@@ -779,6 +832,10 @@ KV_CACHE_ATTENTION_ACCESS_REGISTRY: dict[str, tuple[KVCacheAttentionAccess, ...]
_plain(_PREFILL, _ANY_BACKEND),
_plain(_DECODE, _ANY_BACKEND),
),
CPUFP8KVCacheMethod.name: (
_plain(_PREFILL, _CPU_FP8_BACKENDS),
_plain(_DECODE, _CPU_FP8_BACKENDS),
),
NVFP4KVCacheMethod.name: (
_dq_workspace(_PREFILL, _NVFP4_PREFILL_BACKENDS, _NVFP4_SCALE, _FP8_E4M3),
_native_fp4(_DECODE, _NVFP4_DECODE_BACKENDS, _NVFP4_SCALE, _TORCH_FP4),
@@ -792,6 +849,7 @@ KV_CACHE_ATTENTION_ACCESS_REGISTRY: dict[str, tuple[KVCacheAttentionAccess, ...]
# Registry: explicit --kv-cache-dtype value -> method class.
KV_CACHE_QUANT_REGISTRY: dict[str, type[KVCacheQuantMethodBase]] = {
"cpu_fp8_e4m3": CPUFP8KVCacheMethod,
"nvfp4": NVFP4KVCacheMethod,
"fp4_mx_block16": FP4MXBlock16KVCacheMethod,
}
@@ -67,6 +67,7 @@ from sglang.srt.mem_cache.memory_pool import (
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
attention_backends,
get_context,
get_disagg,
get_exec,
@@ -83,6 +84,7 @@ from sglang.srt.runtime_context import (
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import (
cpu_has_amx_support,
get_available_gpu_memory,
get_device_memory_capacity,
is_float4_e2m1fn_x2,
@@ -293,8 +295,24 @@ class KVCacheConfigurator:
quant_method.load_scales_from_model(self.model)
return quant_method
def _build_mha_quant_method(self, *, num_layers: int):
if current_platform.is_cpu() and self.kv_cache_dtype == torch.float8_e4m3fn:
return get_kv_cache_quant_method("cpu_fp8_e4m3")
return self._build_fp4_quant_method(num_layers=num_layers)
def configure(self, *, pre_model_load_memory: int) -> KVCacheConfigResult:
"""Apply a resolved MemoryPoolConfig and initialize pools."""
if current_platform.is_cpu() and self.kv_cache_dtype == torch.float8_e4m3fn:
if self.use_mla_backend:
raise ValueError("CPU FP8 KV cache is only supported for MHA.")
if not cpu_has_amx_support():
raise ValueError("CPU FP8 KV cache requires Intel AMX support.")
configured_backends = set(attention_backends())
if configured_backends - {"intel_amx"}:
raise ValueError(
"CPU FP8 KV cache requires the intel_amx attention backend."
)
if not self.spec_algorithm.is_none() and self.is_draft_worker:
assert (
self.memory_pool_config is not None
@@ -1228,14 +1246,15 @@ class KVCacheConfigurator:
mha_pool_class=mha_pool_class,
)
else:
quant_method = None
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
quant_method = self._build_mha_quant_method(
num_layers=self.layer_info.num_effective_layers
)
if quant_method is not None and is_float4_e2m1fn_x2(
self.kv_cache_dtype
):
assert (
not enable_page_major
), "page-major KV layout is not supported with fp4 KV cache"
quant_method = self._build_fp4_quant_method(
num_layers=self.layer_info.num_effective_layers
)
token_to_kv_pool = self._build_mha_kv_pool(
max_total_num_tokens=sizes.max_total_num_tokens,
mha_pool_class=mha_pool_class,
@@ -1740,7 +1759,7 @@ class KVCacheConfigurator:
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
)
quant_method = self._build_fp4_quant_method(
quant_method = self._build_mha_quant_method(
num_layers=len(full_attention_layer_ids)
)
# MXFP8 KV cache needs the block-scaled pool (data + UE8M0 scale
@@ -5,6 +5,7 @@ import torch
from torch import nn
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
from sglang.srt.platforms import current_platform
from sglang.srt.utils import is_hip
logger = logging.getLogger(__name__)
@@ -46,6 +47,8 @@ def configure_kv_cache_dtype(
else:
kv_cache_dtype = model_dtype
elif server_args_kv_cache_dtype == "fp8_e5m2":
if current_platform.is_cpu():
raise ValueError("--kv-cache-dtype fp8_e5m2 is not supported on CPU.")
if _is_hip: # Using natively supported format
kv_cache_dtype = fp8_dtype
else: