@@ -955,6 +955,12 @@ void cpu_flash_attention(
955955 int64_t qSplitSize = q_split_size > qSize ? qSize : q_split_size;
956956 int64_t kvSplitSize = kv_split_size > kvSize ? kvSize : kv_split_size;
957957 int64_t qSlice = (qSize - 1 ) / qSplitSize + 1 ;
958+ const auto can_use_kleidiai_bfloat16_prefill = [&](int64_t q_block_size) {
959+ return seq_dim == SeqDim::TWO && qSize > 1 &&
960+ std::is_same<scalar_t , ::executorch::aten::BFloat16>::value &&
961+ ::executorch::cpublas::gemm_uses_kleidiai_bfloat16 (
962+ ::executorch::cpublas::TransposeType::NoTranspose, q_block_size);
963+ };
958964#ifdef ET_USE_THREADPOOL
959965 int64_t num_thread =
960966 ::executorch::extension::threadpool::get_threadpool ()->get_thread_count();
@@ -1005,11 +1011,15 @@ void cpu_flash_attention(
10051011 // Scratch for widening q@K.T to fp32 (see _q_at_k_gemm): one K block plus one
10061012 // q block. qBlockSize cannot exceed qSplitSize, so include the runtime bounds
10071013 // that determine whether any block can use the widened path.
1014+ const auto can_widen_qk = [&](int64_t q_block_size) {
1015+ return !can_use_kleidiai_bfloat16_prefill (q_block_size) &&
1016+ std::is_same<scalar_t , ::executorch::aten::BFloat16>::value &&
1017+ ::executorch::cpublas::gemm_uses_blas () &&
1018+ headSize <= kMaxHeadSizeForWidenedQK &&
1019+ q_block_size >= kMinQBlockForWidenedQK;
1020+ };
10081021 const bool widen_reduced_qk =
1009- std::is_same<scalar_t , ::executorch::aten::BFloat16>::value &&
1010- ::executorch::cpublas::gemm_uses_blas () &&
1011- headSize <= kMaxHeadSizeForWidenedQK &&
1012- qSplitSize >= kMinQBlockForWidenedQK;
1022+ can_widen_qk (qSplitSize) || can_widen_qk (qSize % qSplitSize);
10131023#if defined(__APPLE__)
10141024 const bool dequantize_qk = is_quantized_sdpa &&
10151025 ::executorch::cpublas::gemm_uses_blas () && qSplitSize > 4;
@@ -1082,6 +1092,9 @@ void cpu_flash_attention(
10821092 for (int64_t z = begin; z < end; z++) {
10831093 int64_t m = k * qSplitSize;
10841094 int64_t qBlockSize = std::min (qSplitSize, qSize - m);
1095+ const bool use_kleidiai_bfloat16_prefill =
1096+ can_use_kleidiai_bfloat16_prefill (qBlockSize);
1097+ const bool widen_qk_block = can_widen_qk (qBlockSize);
10851098 // Initialize max and sum
10861099 fill_stub (
10871100 qk_max_data, -std::numeric_limits<accum_t >::infinity (), qBlockSize);
@@ -1181,10 +1194,8 @@ void cpu_flash_attention(
11811194 k_sub_matrix_data,
11821195 kStrideN ,
11831196 qk_data,
1184- ((widen_reduced_qk && qBlockSize >= kMinQBlockForWidenedQK ) ||
1185- (dequantize_qk && qBlockSize > 4 ))
1186- ? widen_ptr
1187- : nullptr );
1197+ (widen_qk_block || (dequantize_qk && qBlockSize > 4 )) ? widen_ptr
1198+ : nullptr );
11881199
11891200 // Update coefficients with scaling, attention mask, and softmax.
11901201 accum_t tmp_max = 0 , tmp_sum = 0 , exp_tmp = 0 ;
@@ -1298,7 +1309,8 @@ void cpu_flash_attention(
12981309 // them in accum_t, widen V and let BLAS multiply -- also one rounding
12991310 // step fewer. Below it the widening stops amortizing.
13001311 constexpr int64_t kMinQBlockForWidenedAV = 64 ;
1301- const bool widen_v = is_reduced_type && !is_quantized_sdpa &&
1312+ const bool widen_v = !use_kleidiai_bfloat16_prefill &&
1313+ is_reduced_type && !is_quantized_sdpa &&
13021314 std::is_same<scalar_t , ::executorch::aten::BFloat16>::value &&
13031315 qBlockSize >= kMinQBlockForWidenedAV ;
13041316 const bool use_fp32_qk_weights = is_reduced_type &&
0 commit comments