mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:20:33 +02:00
ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86 (#29423)
* ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86 * add AVX2 support for masked loading and storing in simd_gemm_ukernel_tail * ggml-cpu: fix FA softcap handling for padded KV tiles
This commit is contained in:
@@ -9037,6 +9037,11 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
|
||||
simd_gemm(KQ, (const float *)Q_q, K_f32, Q_TILE_SZ, DK, KV_TILE_SZ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, scale);
|
||||
|
||||
if (logit_softcap != 0.0f) {
|
||||
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
|
||||
}
|
||||
|
||||
// Set padded KQ entries to -inf so softmax gives them zero weight
|
||||
if (kv_tile < KV_TILE_SZ) {
|
||||
for (int tq = 0; tq < Q_TILE_SZ; tq++) {
|
||||
@@ -9046,11 +9051,6 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
|
||||
}
|
||||
}
|
||||
|
||||
if (logit_softcap != 0.0f) {
|
||||
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
|
||||
}
|
||||
|
||||
if (mask) {
|
||||
ggml_vec_add_f32(tile_rows * KV_TILE_SZ, KQ, KQ, mask32);
|
||||
}
|
||||
@@ -9320,7 +9320,7 @@ static void ggml_compute_forward_flash_attn_ext_f16(
|
||||
kv_is_f32_or_f16 &&
|
||||
k->type == v->type &&
|
||||
neq1 >= Q_TILE_SZ);
|
||||
#ifdef GGML_SIMD
|
||||
#if defined(GGML_SIMD) && !defined(__x86_64__) && !defined(_M_X64)
|
||||
#if defined(__ARM_FEATURE_SVE)
|
||||
const int64_t f32_epr = svcntw();
|
||||
#else
|
||||
|
||||
@@ -56,6 +56,56 @@ static inline void simd_gemm_ukernel(
|
||||
}
|
||||
}
|
||||
|
||||
template <int RM>
|
||||
static inline void simd_gemm_ukernel_tail(
|
||||
float * GGML_RESTRICT C,
|
||||
const float * GGML_RESTRICT A,
|
||||
const float * GGML_RESTRICT B,
|
||||
int K, int N, int cols)
|
||||
{
|
||||
#if defined(__AVX512F__)
|
||||
const __mmask16 mask = (1u << cols) - 1;
|
||||
__m512 acc[RM];
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm512_maskz_loadu_ps(mask, C + i * N);
|
||||
}
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
const __m512 b = _mm512_maskz_loadu_ps(mask, B + kk * N);
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm512_mask3_fmadd_ps(_mm512_set1_ps(A[i * K + kk]), b, acc[i], mask);
|
||||
}
|
||||
}
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
_mm512_mask_storeu_ps(C + i * N, mask, acc[i]);
|
||||
}
|
||||
#elif defined(__AVX2__)
|
||||
const __m256i mask = _mm256_cmpgt_epi32(_mm256_set1_epi32(cols), _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
|
||||
__m256 acc[RM];
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm256_maskload_ps(C + i * N, mask);
|
||||
}
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
const __m256 b = _mm256_maskload_ps(B + kk * N, mask);
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = GGML_F32_VEC_FMA(acc[i], b, _mm256_set1_ps(A[i * K + kk]));
|
||||
}
|
||||
}
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
_mm256_maskstore_ps(C + i * N, mask, acc[i]);
|
||||
}
|
||||
#else
|
||||
for (int64_t j = 0; j < cols; j++) {
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
float a = C[i * N + j];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[i * K + kk] * B[kk * N + j];
|
||||
}
|
||||
C[i * N + j] = a;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// C[M x N] += A[M x K] * B[K x N]
|
||||
static void simd_gemm(
|
||||
float * GGML_RESTRICT C,
|
||||
@@ -74,14 +124,8 @@ static void simd_gemm(
|
||||
for (; jj + KN <= N; jj += KN) {
|
||||
simd_gemm_ukernel<GEMM_RM, 1>(C + jj, A, B + jj, K, N);
|
||||
}
|
||||
for (; jj < N; jj++) {
|
||||
for (int64_t i = 0; i < GEMM_RM; i++) {
|
||||
float a = C[i * N + jj];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[i * K + kk] * B[kk * N + jj];
|
||||
}
|
||||
C[i * N + jj] = a;
|
||||
}
|
||||
if (jj < N) {
|
||||
simd_gemm_ukernel_tail<GEMM_RM>(C + jj, A, B + jj, K, N, N - jj);
|
||||
}
|
||||
|
||||
A += GEMM_RM * K;
|
||||
@@ -97,12 +141,8 @@ static void simd_gemm(
|
||||
for (; jj + KN <= N; jj += KN) {
|
||||
simd_gemm_ukernel<1, 1>(C + jj, A, B + jj, K, N);
|
||||
}
|
||||
for (; jj < N; jj++) {
|
||||
float a = C[jj];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[kk] * B[kk * N + jj];
|
||||
}
|
||||
C[jj] = a;
|
||||
if (jj < N) {
|
||||
simd_gemm_ukernel_tail<1>(C + jj, A, B + jj, K, N, N - jj);
|
||||
}
|
||||
|
||||
A += K;
|
||||
|
||||
@@ -11009,6 +11009,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
||||
// asymmetric head_dim (hsk != hsv) with one or both sides not 64-aligned
|
||||
test_cases.emplace_back(new test_flash_attn_ext(72, 64, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 72, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 113, 75, true, true, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 17, 75, false, false, 0, 1.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 113, 75, false, false, 0, 1.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
|
||||
|
||||
// mixed quant and Q1_0 test cases
|
||||
test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0));
|
||||
|
||||
Reference in New Issue
Block a user