From 96127976f29c59563e52607c6c3a88680b163ff5 Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Sat, 9 May 2026 08:32:20 +0300 Subject: [PATCH] Use AVX2 when available for greedy speculative sampling (#1761) * Use AVX2 when available for greedy speculative sampling * Avoid some code duplication --- common/sampling.cpp | 148 ++++++++++++++++++++++++++++++++++++++--- common/speculative.cpp | 3 + 2 files changed, 143 insertions(+), 8 deletions(-) diff --git a/common/sampling.cpp b/common/sampling.cpp index 3cc7aec7..5c0a860d 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -4,7 +4,11 @@ #include "common.h" #include "reasoning-budget.cpp" +#include #include +#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__)) +#include +#endif #include using json = nlohmann::ordered_json; @@ -88,9 +92,6 @@ struct common_sampler * common_sampler_init(const struct llama_model * model, co result->grammar_str = grammar_str; result->grammar_root = "root"; } - - - // Feed generation prompt tokens to the grammar sampler so it advances past // tokens the template already placed in the prompt. @@ -884,6 +885,127 @@ common_grammar_trigger common_grammar_trigger::from_json(const json& in) { return out; } +#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__)) +__attribute__((target("avx2"))) +static bool common_sampler_speculative_top1_avx2(const float * logits, const int n_vocab, int & best_id, float & max_val) { + if (n_vocab < 8) { + return false; + } + + __m256 max_v = _mm256_loadu_ps(logits); + __m256i id_v = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + + const __m256i step = _mm256_set1_epi32(8); + __m256i cur_id = _mm256_add_epi32(id_v, step); + + int i = 8; + for (; i + 7 < n_vocab; i += 8) { + const __m256 x = _mm256_loadu_ps(logits + i); + const __m256 gt_max = _mm256_cmp_ps(x, max_v, _CMP_GT_OQ); + + max_v = _mm256_blendv_ps(max_v, x, gt_max); + id_v = _mm256_blendv_epi8(id_v, cur_id, _mm256_castps_si256(gt_max)); + cur_id = _mm256_add_epi32(cur_id, step); + } + + alignas(32) float max_buf[8]; + alignas(32) int id_buf[8]; + _mm256_store_ps(max_buf, max_v); + _mm256_store_si256((__m256i *) id_buf, id_v); + + best_id = id_buf[0]; + max_val = max_buf[0]; + + for (int j = 1; j < 8; ++j) { + if (max_buf[j] > max_val) { + max_val = max_buf[j]; best_id = id_buf[j]; + } + } + for (; i < n_vocab; ++i) { + + if (logits[i] > max_val) { + max_val = logits[i]; best_id = i; + } + } + return true; +} +__attribute__((target("avx2,fma"))) +static inline __m256 v_expf(__m256 x) { + const __m256 r = _mm256_set1_ps(0x1.8p23f); + const __m256 z = _mm256_fmadd_ps(x, _mm256_set1_ps(0x1.715476p+0f), r); + const __m256 n = _mm256_sub_ps(z, r); + const __m256 b = _mm256_fnmadd_ps(n, _mm256_set1_ps(0x1.7f7d1cp-20f), + _mm256_fnmadd_ps(n, _mm256_set1_ps(0x1.62e4p-1f), x)); + const __m256i e = _mm256_slli_epi32(_mm256_castps_si256(z), 23); + const __m256 k = _mm256_castsi256_ps( + _mm256_add_epi32(e, _mm256_castps_si256(_mm256_set1_ps(1)))); + const __m256i c = _mm256_castps_si256( + _mm256_cmp_ps(_mm256_andnot_ps(_mm256_set1_ps(-0.f), n), + _mm256_set1_ps(126), _CMP_GT_OQ)); + const __m256 u = _mm256_mul_ps(b, b); + const __m256 j = _mm256_fmadd_ps(_mm256_fmadd_ps(_mm256_fmadd_ps(_mm256_set1_ps(0x1.0e4020p-7f), b, + _mm256_set1_ps(0x1.573e2ep-5f)), u, + _mm256_fmadd_ps(_mm256_set1_ps(0x1.555e66p-3f), b, + _mm256_set1_ps(0x1.fffdb6p-2f))), + u, _mm256_mul_ps(_mm256_set1_ps(0x1.ffffecp-1f), b)); + if (!_mm256_movemask_ps(_mm256_castsi256_ps(c))) + return _mm256_fmadd_ps(j, k, k); + const __m256i g = _mm256_and_si256( + _mm256_castps_si256(_mm256_cmp_ps(n, _mm256_setzero_ps(), _CMP_LE_OQ)), + _mm256_set1_epi32(0x82000000u)); + const __m256 s1 = + _mm256_castsi256_ps(_mm256_add_epi32(g, _mm256_set1_epi32(0x7f000000u))); + const __m256 s2 = _mm256_castsi256_ps(_mm256_sub_epi32(e, g)); + const __m256i d = _mm256_castps_si256( + _mm256_cmp_ps(_mm256_andnot_ps(_mm256_set1_ps(-0.f), n), + _mm256_set1_ps(192), _CMP_GT_OQ)); + return _mm256_or_ps( + _mm256_and_ps(_mm256_castsi256_ps(d), _mm256_mul_ps(s1, s1)), + _mm256_andnot_ps( + _mm256_castsi256_ps(d), + _mm256_or_ps( + _mm256_and_ps(_mm256_castsi256_ps(c), + _mm256_mul_ps(_mm256_fmadd_ps(s2, j, s2), s1)), + _mm256_andnot_ps(_mm256_castsi256_ps(c), _mm256_fmadd_ps(k, j, k))))); +} +__attribute__((target("avx2"))) +static inline float hsum_float_4(__m128 x) { + x = _mm_add_ps(x, _mm_movehl_ps(x, x)); + x = _mm_add_ss(x, _mm_movehdup_ps(x)); + return _mm_cvtss_f32(x); +} +__attribute__((target("avx2"))) +static inline float hsum_float_8(__m256 x) { + return hsum_float_4(_mm_add_ps(_mm256_castps256_ps128(x), _mm256_extractf128_ps(x, 1))); +} +__attribute__((target("avx2,fma"))) +static float prob_avx2(int n, const float * logits, float max_val) { + float sumf = 0; + int i = 0; + if (n >= 8) { + auto sum_v = _mm256_setzero_ps(); + auto max_v = _mm256_set1_ps(max_val); + for (; i < n - 7; i += 8) { + auto x = _mm256_loadu_ps(logits + i); + auto exp_x = v_expf(_mm256_sub_ps(x, max_v)); + sum_v = _mm256_add_ps(sum_v, exp_x); + } + sumf = hsum_float_8(sum_v); + } + for (; i < n; ++i) { + sumf += expf(logits[i] - max_val); + } + return 1.0f/sumf; +} +#endif +static float prob_scalar(int n, const float * logits, float max_val) { + double sum_exp = 0.0; + for (int i = 0; i < n; ++i) { + sum_exp += exp((double)(logits[i] - max_val)); + } + return (float)(1./sum_exp); +} + llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, float * out_prob) { GGML_UNUSED(gsmpl); @@ -892,6 +1014,20 @@ llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, str int best_id = 0; float max_val = logits[0]; +#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__)) + static const bool has_avx2 = __builtin_cpu_supports("avx2"); + if (has_avx2 && common_sampler_speculative_top1_avx2(logits, n_vocab, best_id, max_val)) { + if (out_prob) { + static const bool has_fma = __builtin_cpu_supports("fma"); + if (has_fma) { + *out_prob = prob_avx2(n_vocab, logits, max_val); + } else { + *out_prob = prob_scalar(n_vocab, logits, max_val); + } + } + return best_id; + } +#endif for (int i = 1; i < n_vocab; ++i) { if (logits[i] > max_val) { max_val = logits[i]; @@ -900,11 +1036,7 @@ llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, str } if (out_prob) { - double sum_exp = 0.0; - for (int i = 0; i < n_vocab; ++i) { - sum_exp += exp((double)(logits[i] - max_val)); - } - *out_prob = (float)(1.0 / sum_exp); + *out_prob = prob_scalar(n_vocab, logits, max_val); } return best_id; diff --git a/common/speculative.cpp b/common/speculative.cpp index d2685067..5abb7467 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -1400,6 +1400,8 @@ std::vector mtp_speculative_gen_draft( int i0 = 0; if (last.last_id >= 0) { if (last.prob < p_min) { + llama_batch_free(mtp_batch); + llama_set_mtp_op_type(ctx, MTP_OP_NONE); return drafts; } current_input_id = last.last_id; @@ -1421,6 +1423,7 @@ std::vector mtp_speculative_gen_draft( } llama_token id_next = common_sampler_sample_speculative(smpl, ctx, 0, prob_ptr); + if (i > 0 && prob_ptr && prob < p_min) { break; }