From cd8e4fe2879253c63b46746d494ae032c1b5bbd1 Mon Sep 17 00:00:00 2001 From: Jake Stevens Date: Mon, 24 Aug 2026 10:50:31 -0700 Subject: [PATCH] Add BF16 support to sampler Summary: As title. Convert to FP32 and keep the compute in that dtype for numerics and speed per local benchmarking Differential Revision: D117223300 --- extension/llm/sampler/sampler.cpp | 13 +++++++++++-- extension/llm/sampler/sampler.h | 6 ++++++ extension/llm/sampler/test/test_sampler.cpp | 19 +++++++++++++++++++ 3 files changed, 36 insertions(+), 2 deletions(-) diff --git a/extension/llm/sampler/sampler.cpp b/extension/llm/sampler/sampler.cpp index d41da96f07e..c4af9ecc350 100644 --- a/extension/llm/sampler/sampler.cpp +++ b/extension/llm/sampler/sampler.cpp @@ -250,12 +250,21 @@ int32_t Sampler::sample(T* logits) { return next; } +template <> +int32_t Sampler::sample( + executorch::aten::BFloat16* logits) { + float_logits_buffer_.resize(vocab_size_); + for (int i = 0; i < vocab_size_; i++) { + float_logits_buffer_[i] = static_cast(logits[i]); + } + + return sample(float_logits_buffer_.data()); +} + template int32_t Sampler::sample(float* logits); template int32_t Sampler::sample(uint16_t* logits); template int32_t Sampler::sample( executorch::aten::Half* logits); -template int32_t Sampler::sample( - executorch::aten::BFloat16* logits); } // namespace llm } // namespace extension diff --git a/extension/llm/sampler/sampler.h b/extension/llm/sampler/sampler.h index e340c12c83d..13e964dcb27 100644 --- a/extension/llm/sampler/sampler.h +++ b/extension/llm/sampler/sampler.h @@ -14,6 +14,7 @@ #include #include #include +#include #ifdef USE_ATEN_LIB #include #endif @@ -80,8 +81,13 @@ class ET_EXPERIMENTAL Sampler { // 0 (or >= vocab_size_) means top-k is disabled. int32_t topk_ = 0; unsigned long long rng_state_; + std::vector float_logits_buffer_; }; +template <> +int32_t Sampler::sample( + executorch::aten::BFloat16* logits); + } // namespace llm } // namespace extension } // namespace executorch diff --git a/extension/llm/sampler/test/test_sampler.cpp b/extension/llm/sampler/test/test_sampler.cpp index 8463c2e9678..617d17c5559 100644 --- a/extension/llm/sampler/test/test_sampler.cpp +++ b/extension/llm/sampler/test/test_sampler.cpp @@ -9,6 +9,7 @@ #include #include +#include #include #include @@ -42,6 +43,24 @@ TEST(SamplerTest, TestArgMaxWithFP16) { EXPECT_EQ(sampler.sample(input.data_ptr()), 396); } +TEST(SamplerTest, TestBFloat16SamplingUsesFloatArithmetic) { + constexpr int kVocabSize = 32768; + constexpr unsigned long long kSeed = 42; + Sampler float_sampler{kVocabSize, 1.0f, 1.0f, kSeed}; + Sampler bf16_sampler{kVocabSize, 1.0f, 1.0f, kSeed}; + + std::vector bf16_logits( + kVocabSize, executorch::aten::BFloat16(0.0f)); + std::vector float_logits(kVocabSize); + for (int i = 0; i < kVocabSize; i++) { + float_logits[i] = static_cast(bf16_logits[i]); + } + + EXPECT_EQ( + bf16_sampler.sample(bf16_logits.data()), + float_sampler.sample(float_logits.data())); +} + TEST(SamplerTest, TestTopKRestrictsToCandidates) { // With topk=3, sampling must always return one of the top-3 indices, // regardless of the random draw.