Skip to content

Commit 32fb206

Browse files
authored
Add KleiDI BF16 GEMM and SDPA dispatch (pytorch#22886)
Summary: Add KleiDI NEON and SME2 BF16 GEMM dispatch for Arm64 and route BF16 SDPA prefill through it. Keep single-column decode on BFDOT and retain fallbacks for unsupported CPUs or older KleiDI builds. Scalar NEON packing and an adapter for the existing SME RHS packer allow this integration to land before the new upstream packers. AI-assisted: Codex. Reviewed By: digantdesai, Gasoonjia Differential Revision: D119433385 Pull Request resolved: pytorch#22886
1 parent 4322502 commit 32fb206

10 files changed

Lines changed: 854 additions & 30 deletions

File tree

‎CMakeLists.txt‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1195,6 +1195,22 @@ if(EXECUTORCH_BUILD_KERNELS_TORCHAO)
11951195

11961196
endif()
11971197

1198+
if(TARGET cpublas
1199+
AND TARGET kleidiai
1200+
AND TARGET cpuinfo
1201+
)
1202+
if(CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64|arm64"
1203+
OR ANDROID_ABI STREQUAL "arm64-v8a"
1204+
OR CMAKE_OSX_ARCHITECTURES MATCHES "arm64"
1205+
)
1206+
target_compile_definitions(cpublas PRIVATE ET_BUILD_WITH_KLEIDIAI)
1207+
if(CMAKE_SYSTEM_NAME STREQUAL "Darwin")
1208+
target_compile_definitions(cpublas PRIVATE ET_KLEIDIAI_DISABLE_NEON_BF16)
1209+
endif()
1210+
target_link_libraries(cpublas PRIVATE kleidiai)
1211+
endif()
1212+
endif()
1213+
11981214
# The shared build ships the profiler as one of its libraries, and the Python
11991215
# extension records a hard dependency on it, and the Apple ETDump extension
12001216
# links it, so the target has to exist whenever any of those is being built

‎extension/llm/custom_ops/op_sdpa_impl.h‎

Lines changed: 21 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -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 &&

‎extension/llm/custom_ops/op_sdpa_test.cpp‎

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -839,6 +839,60 @@ TEST(OpScaledDotProductAttentionTest, BFloat16MatchesFloat) {
839839
2e-2, 2e-2);
840840
}
841841

842+
TEST(OpScaledDotProductAttentionTest, BFloat16PrefillTrailingQueryBlocks) {
843+
using executorch::aten::BFloat16;
844+
TensorFactory<executorch::aten::ScalarType::BFloat16> tf_bfloat16;
845+
TensorFactory<executorch::aten::ScalarType::Float> tf_float;
846+
constexpr int32_t kHeadSize = 8;
847+
constexpr int32_t kKeys = 513;
848+
849+
// These sequences use 64-query blocks and leave 1-4 queries in the tail.
850+
// A second KV block also exercises accumulation into the first block's
851+
// result.
852+
for (const int32_t queries : {193, 194, 195, 196}) {
853+
SCOPED_TRACE(::testing::Message() << "queries=" << queries);
854+
auto query = tf_bfloat16.zeros({1, 1, queries, kHeadSize});
855+
auto key = tf_bfloat16.zeros({1, 1, kKeys, kHeadSize});
856+
auto value = tf_bfloat16.zeros({1, 1, kKeys, kHeadSize});
857+
auto out = tf_bfloat16.zeros({1, 1, queries, kHeadSize});
858+
auto query_float = tf_float.zeros({1, 1, queries, kHeadSize});
859+
auto key_float = tf_float.zeros({1, 1, kKeys, kHeadSize});
860+
auto value_float = tf_float.zeros({1, 1, kKeys, kHeadSize});
861+
auto expected = tf_float.zeros({1, 1, queries, kHeadSize});
862+
for (int32_t i = 0; i < queries * kHeadSize; ++i) {
863+
const float v = static_cast<float>((i * 7) % 31 - 15) / 16.0f;
864+
query.mutable_data_ptr<BFloat16>()[i] = BFloat16(v);
865+
query_float.mutable_data_ptr<float>()[i] = v;
866+
}
867+
for (int32_t i = 0; i < kKeys * kHeadSize; ++i) {
868+
const float k = static_cast<float>((i * 11) % 29 - 14) / 16.0f;
869+
const float v = static_cast<float>((i * 13) % 37 - 18) / 16.0f;
870+
key.mutable_data_ptr<BFloat16>()[i] = BFloat16(k);
871+
key_float.mutable_data_ptr<float>()[i] = k;
872+
value.mutable_data_ptr<BFloat16>()[i] = BFloat16(v);
873+
value_float.mutable_data_ptr<float>()[i] = v;
874+
}
875+
op_scaled_dot_product_attention(
876+
query, key, value, std::nullopt, 0.0, false, std::nullopt, out);
877+
op_scaled_dot_product_attention(
878+
query_float,
879+
key_float,
880+
value_float,
881+
std::nullopt,
882+
0.0,
883+
false,
884+
std::nullopt,
885+
expected);
886+
for (int32_t i = 0; i < queries * kHeadSize; ++i) {
887+
EXPECT_NEAR(
888+
static_cast<float>(out.const_data_ptr<BFloat16>()[i]),
889+
expected.const_data_ptr<float>()[i],
890+
2e-3f)
891+
<< "index=" << i;
892+
}
893+
}
894+
}
895+
842896
TEST(OpScaledDotProductAttentionTest, HalfMatchesFloat) {
843897
test_reduced_precision_matches_float<executorch::aten::ScalarType::Half>(
844898
1e-2, 1e-2);

‎kernels/optimized/blas/CPUBlas.cpp‎

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,10 @@
88

99
#include <executorch/kernels/optimized/blas/CPUBlas.h>
1010

11-
#include <limits.h>
11+
#include <cstddef>
12+
#include <cstdint>
13+
14+
#include <executorch/kernels/optimized/blas/KleidiBlas.h>
1215

1316
#ifdef ET_BUILD_WITH_BLAS
1417
#ifdef ET_BUILD_FOR_APPLE
@@ -267,6 +270,12 @@ void gemm(
267270
const float beta,
268271
float *c, int64_t ldc) {
269272
normalize_last_dims(transa, transb, m, n, k, &lda, &ldb, &ldc);
273+
#if defined(ET_KLEIDIAI_HAS_NEON_BF16) || defined(ET_KLEIDIAI_HAS_SME2_BF16)
274+
if (kleidiai_bfloat16_gemm(
275+
transa, transb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc)) {
276+
return;
277+
}
278+
#endif
270279
gemm_impl<BFloat16, float, float>(
271280
transa, transb,
272281
m, n, k,

‎kernels/optimized/blas/CPUBlas.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,9 @@ inline char to_blas(TransposeType trans) {
5050
// this library.
5151
bool gemm_uses_blas();
5252

53+
// Whether the BF16-input, float-output gemm can use KleidiAI for this shape.
54+
bool gemm_uses_kleidiai_bfloat16(TransposeType transb, int64_t n);
55+
5356
// Column-major c = beta * c + alpha * (a @ b), where a is BFloat16 and b and
5457
// c are float.
5558
void gemv(

0 commit comments

Comments
 (0)