test1111111 / native /kernels_avx2.cpp
spitfire4794's picture
Serve SurjoLabs/Surjo-50m-SFT-Only (int8, 2T) on 7860
92dcc4e verified
Raw History Blame Contribute Delete
55.2 kB
#include "kernels.hpp"
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <cstdlib>
#include <cstring>
#include <immintrin.h>
#include <limits>
namespace cism {
// All kernels require AVX2 plus FMA (checked together in has_avx2()); every
// FMA here is an explicit intrinsic, never compiler contraction.
static float reduce_sum(__m256 value) {
__m128 sum = _mm_add_ps(_mm256_castps256_ps128(value), _mm256_extractf128_ps(value, 1));
sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
sum = _mm_add_ss(sum, _mm_shuffle_ps(sum, sum, _MM_SHUFFLE(1, 1, 1, 1)));
return _mm_cvtss_f32(sum);
}
static __m256 int8_values(const std::int8_t* weights) {
const __m128i bytes = _mm_loadl_epi64(reinterpret_cast<const __m128i*>(weights));
return _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(bytes));
}
#ifdef CISM_HAVE_F16C
// FP16 dot via F16C widening (1 uop per 8) + FMA: same shape as int8 above
// (8 accumulators, 64-wide), half the weight bytes of fp32. Scalar tail via
// the portable helper (bit-exact RNE twin of vcvtph).
float dot_fp16_avx2(const std::uint16_t* weights, const float* input, std::size_t n) {
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
__m256 sum4 = _mm256_setzero_ps(), sum5 = _mm256_setzero_ps();
__m256 sum6 = _mm256_setzero_ps(), sum7 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, weights += 64, input += 64) {
// +2048 B runway (1024 fp16 elems), matching dot_int8's tuned distance.
_mm_prefetch(reinterpret_cast<const char*>(weights + 1024), _MM_HINT_T0);
sum0 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights))),
_mm256_loadu_ps(input), sum0);
sum1 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 8))),
_mm256_loadu_ps(input + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16))),
_mm256_loadu_ps(input + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 24))),
_mm256_loadu_ps(input + 24), sum3);
sum4 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 32))),
_mm256_loadu_ps(input + 32), sum4);
sum5 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 40))),
_mm256_loadu_ps(input + 40), sum5);
sum6 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 48))),
_mm256_loadu_ps(input + 48), sum6);
sum7 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 56))),
_mm256_loadu_ps(input + 56), sum7);
}
for (; i + 8 <= n; i += 8, weights += 8, input += 8)
sum0 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights))),
_mm256_loadu_ps(input), sum0);
float result = reduce_sum(_mm256_add_ps(_mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3)),
_mm256_add_ps(_mm256_add_ps(sum4, sum5), _mm256_add_ps(sum6, sum7))));
for (; i < n; ++i, ++weights, ++input) result += fp16_to_fp32(*weights) * *input;
return result;
}
#endif
float dot_int8_avx2(const std::int8_t* weights, const float* input, std::size_t n) {
// Eight independent FMA/convert chains hide Zen3's ~4-cycle FMA latency
// across its 2 FMA units (8 chains ideal); each chain owns its
// int8->fp32 converts so widen latency hides too. 64-wide blocking keeps
// one FMA per accumulator per iteration (no intra-iteration dependency,
// unlike 2x reuse under 4 accums). Weights T0 prefetch at +2048 B
// (measured optimum on Zen3: 256 -> 1.00x, 1024 -> 1.03x, 2048 -> 1.07x,
// 4096 -> 1.03x; longer runways overlap more DRAM stalls until cache
// pollution dominates). Activations are L1-resident/shared across rows
// so their prefetch was pure uop overhead. Override via
// CISM_PREFETCH_DIST for other machines. Tail/merge order below only
// matters for <64 leftovers (Surjo-50m cols are multiples of 64).
static const std::intptr_t prefetch_dist = []() -> std::intptr_t {
if (const char* env = std::getenv("CISM_PREFETCH_DIST")) {
long v = std::atol(env);
if (v >= 0 && v <= 4096) return static_cast<std::intptr_t>(v);
}
return 2048;
}();
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
__m256 sum4 = _mm256_setzero_ps(), sum5 = _mm256_setzero_ps();
__m256 sum6 = _mm256_setzero_ps(), sum7 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, weights += 64, input += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights) + prefetch_dist, _MM_HINT_T0);
sum0 = _mm256_fmadd_ps(int8_values(weights), _mm256_loadu_ps(input), sum0);
sum1 = _mm256_fmadd_ps(int8_values(weights + 8), _mm256_loadu_ps(input + 8), sum1);
sum2 = _mm256_fmadd_ps(int8_values(weights + 16), _mm256_loadu_ps(input + 16), sum2);
sum3 = _mm256_fmadd_ps(int8_values(weights + 24), _mm256_loadu_ps(input + 24), sum3);
sum4 = _mm256_fmadd_ps(int8_values(weights + 32), _mm256_loadu_ps(input + 32), sum4);
sum5 = _mm256_fmadd_ps(int8_values(weights + 40), _mm256_loadu_ps(input + 40), sum5);
sum6 = _mm256_fmadd_ps(int8_values(weights + 48), _mm256_loadu_ps(input + 48), sum6);
sum7 = _mm256_fmadd_ps(int8_values(weights + 56), _mm256_loadu_ps(input + 56), sum7);
}
for (; i + 32 <= n; i += 32, weights += 32, input += 32) {
sum0 = _mm256_fmadd_ps(int8_values(weights), _mm256_loadu_ps(input), sum0);
sum1 = _mm256_fmadd_ps(int8_values(weights + 8), _mm256_loadu_ps(input + 8), sum1);
sum2 = _mm256_fmadd_ps(int8_values(weights + 16), _mm256_loadu_ps(input + 16), sum2);
sum3 = _mm256_fmadd_ps(int8_values(weights + 24), _mm256_loadu_ps(input + 24), sum3);
}
const __m256 s01 = _mm256_add_ps(sum0, sum1);
const __m256 s23 = _mm256_add_ps(sum2, sum3);
const __m256 s45 = _mm256_add_ps(sum4, sum5);
const __m256 s67 = _mm256_add_ps(sum6, sum7);
__m256 sum = _mm256_add_ps(_mm256_add_ps(s01, s23), _mm256_add_ps(s45, s67));
for (; i + 8 <= n; i += 8, weights += 8, input += 8)
sum = _mm256_fmadd_ps(int8_values(weights), _mm256_loadu_ps(input), sum);
float result = reduce_sum(sum);
for (; i < n; ++i, ++weights, ++input)
result += static_cast<float>(*weights) * *input;
return result;
}
float dot_int4_avx2(const std::uint8_t* weights, const float* act_perm, const float* act_orig, const float* scales, std::size_t n) {
(void)act_perm; // split-half layout dots against linear acts; no permute.
// Split-half nibbles: byte j holds w[j] (low) and w[j+16] (high), so one
// 16B load + and/srli yields both linear halves with no LUT and no
// shuffle: sub-8 (signed codes) + widen + FMA vs linear activations.
// 4 accumulators stay independent across blocks (scaled per block, single
// reduce at the end).
const __m128i mask = _mm_set1_epi8(15);
const __m128i eight = _mm_set1_epi8(8);
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
// 64-wide inner blocking: two 32-blocks per iteration in order, same 4
// accumulators (bitwise identical to the 32-wide loop). Prefetches target
// the next chunk; x86 prefetches never fault.
for (; i + 64 <= n; i += 64, weights += 32, act_orig += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 32), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act_orig + 64), _MM_HINT_T0);
{
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i lo = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i hi = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
const __m256 scale = _mm256_broadcast_ss(scales + i / 32);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale),
_mm256_loadu_ps(act_orig), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(lo, 8))), scale),
_mm256_loadu_ps(act_orig + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale),
_mm256_loadu_ps(act_orig + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(hi, 8))), scale),
_mm256_loadu_ps(act_orig + 24), sum3);
}
{
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16));
const __m128i lo = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i hi = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
const __m256 scale = _mm256_broadcast_ss(scales + i / 32 + 1);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale),
_mm256_loadu_ps(act_orig + 32), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(lo, 8))), scale),
_mm256_loadu_ps(act_orig + 40), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale),
_mm256_loadu_ps(act_orig + 48), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(hi, 8))), scale),
_mm256_loadu_ps(act_orig + 56), sum3);
}
}
for (; i + 32 <= n; i += 32, weights += 16, act_orig += 32) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i lo = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i hi = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
const __m256 scale = _mm256_broadcast_ss(scales + i / 32);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale),
_mm256_loadu_ps(act_orig), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(lo, 8))), scale),
_mm256_loadu_ps(act_orig + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale),
_mm256_loadu_ps(act_orig + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(
_mm_srli_si128(hi, 8))), scale),
_mm256_loadu_ps(act_orig + 24), sum3);
}
float result = reduce_sum(_mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3)));
// Partial tail block, read only to its logical end. weights points at
// the current 16-byte split-half block; tail elements use block-relative
// split-half indexing against linear activations.
if (i < n) {
float block_sum = 0;
const float scale = scales[i / 32];
const float* tail_act = act_orig;
const std::size_t m = n - i;
for (std::size_t t = 0; t < m; ++t) {
const int nibble = t < 16 ? (weights[t] & 15) : ((weights[t - 16] >> 4) & 15);
block_sum += static_cast<float>(nibble - 8) * tail_act[t];
}
result += block_sum * scale;
}
return result;
}
float dot_fp4_avx2(const std::uint8_t* weights, const float* act_perm, const float* act_orig, const std::uint8_t* scales, std::size_t n) {
// E2M1 elements: raw nibbles go through a pshufb 16-entry table (half
// values, exact in int8), widen to FP32, then scale with the halved E4M3
// table entry — no sign-fix chain, no per-element convert beyond the
// widen. Permuted activations: low nibbles dot PERM[i..i+15], high
// nibbles dot PERM[i+16..i+31], no unpacklo/hi interleave. Chains stay
// independent across blocks (scaled per block, single reduce at the end).
// Scale mapping: lo[0..7]+hi[0..7] are row elements 0-15 (scale0),
// lo[8..15]+hi[8..15] are elements 16-31 (scale1): the 16-element scale
// boundary cuts across the even/odd split, not along it.
const __m128i mask = _mm_set1_epi8(15);
const __m128i lut = _mm_loadu_si128(reinterpret_cast<const __m128i*>(fp4_element_lut()));
const float* scale_lut = fp4_scale_lut();
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
// 64-wide inner blocking (two 32-groups per iteration, 4 scales in order;
// bitwise identical to the 32-wide form). Prefetches target next chunk.
for (; i + 64 <= n; i += 64, weights += 32, act_perm += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 32), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act_perm + 64), _MM_HINT_T0);
{
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m128i lo = _mm_shuffle_epi8(lut, low);
const __m128i hi = _mm_shuffle_epi8(lut, high);
const std::size_t block = i / 16;
const __m256 scale0 = _mm256_broadcast_ss(scale_lut + scales[block]);
const __m256 scale1 = _mm256_broadcast_ss(scale_lut + scales[block + 1]);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale0),
_mm256_loadu_ps(act_perm), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(lo, 8))), scale1),
_mm256_loadu_ps(act_perm + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale0),
_mm256_loadu_ps(act_perm + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(hi, 8))), scale1),
_mm256_loadu_ps(act_perm + 24), sum3);
}
{
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m128i lo = _mm_shuffle_epi8(lut, low);
const __m128i hi = _mm_shuffle_epi8(lut, high);
const std::size_t block = i / 16 + 2;
const __m256 scale0 = _mm256_broadcast_ss(scale_lut + scales[block]);
const __m256 scale1 = _mm256_broadcast_ss(scale_lut + scales[block + 1]);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale0),
_mm256_loadu_ps(act_perm + 32), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(lo, 8))), scale1),
_mm256_loadu_ps(act_perm + 40), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale0),
_mm256_loadu_ps(act_perm + 48), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(hi, 8))), scale1),
_mm256_loadu_ps(act_perm + 56), sum3);
}
}
for (; i + 32 <= n; i += 32, weights += 16, act_perm += 32) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m128i lo = _mm_shuffle_epi8(lut, low);
const __m128i hi = _mm_shuffle_epi8(lut, high);
const std::size_t block = i / 16;
const __m256 scale0 = _mm256_broadcast_ss(scale_lut + scales[block]);
const __m256 scale1 = _mm256_broadcast_ss(scale_lut + scales[block + 1]);
sum0 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(lo)), scale0),
_mm256_loadu_ps(act_perm), sum0);
sum1 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(lo, 8))), scale1),
_mm256_loadu_ps(act_perm + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(hi)), scale0),
_mm256_loadu_ps(act_perm + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_mul_ps(_mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(hi, 8))), scale1),
_mm256_loadu_ps(act_perm + 24), sum3);
}
float result = reduce_sum(_mm256_add_ps(_mm256_add_ps(sum0, sum2), _mm256_add_ps(sum1, sum3)));
if (i < n) {
// Partial tail block, read only to the row's logical end; uses the
// ORIGINAL-order activation pointer with the scalar nibble code.
const auto* elements = fp4_element_lut();
const float* tail_act = act_orig + i;
float block_sum = 0;
for (std::size_t j = 0; i < n; ++i, ++j) {
const int nibble = (weights[j / 2] >> (4 * (i % 2))) & 15;
block_sum += static_cast<float>(elements[nibble]) * scale_lut[scales[i / 16]] * tail_act[j];
}
result += block_sum;
}
return result;
}
float dot_avx2(const float* a, const float* b, std::size_t n) {
// 64-wide inner blocking, same 4 accumulators in order (bitwise identical
// to the 32-wide form); T0 prefetches help spill sizes, free otherwise.
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, a += 64, b += 64) {
_mm_prefetch(reinterpret_cast<const char*>(a + 128), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(b + 128), _MM_HINT_T0);
sum0 = _mm256_fmadd_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b), sum0);
sum1 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 8), _mm256_loadu_ps(b + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 16), _mm256_loadu_ps(b + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 24), _mm256_loadu_ps(b + 24), sum3);
sum0 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 32), _mm256_loadu_ps(b + 32), sum0);
sum1 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 40), _mm256_loadu_ps(b + 40), sum1);
sum2 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 48), _mm256_loadu_ps(b + 48), sum2);
sum3 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 56), _mm256_loadu_ps(b + 56), sum3);
}
for (; i + 32 <= n; i += 32, a += 32, b += 32) {
sum0 = _mm256_fmadd_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b), sum0);
sum1 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 8), _mm256_loadu_ps(b + 8), sum1);
sum2 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 16), _mm256_loadu_ps(b + 16), sum2);
sum3 = _mm256_fmadd_ps(_mm256_loadu_ps(a + 24), _mm256_loadu_ps(b + 24), sum3);
}
sum0 = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
for (; i + 8 <= n; i += 8, a += 8, b += 8)
sum0 = _mm256_fmadd_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b), sum0);
alignas(32) float lanes[8];
_mm256_store_ps(lanes, sum0);
float result = lanes[0] + lanes[1] + lanes[2] + lanes[3] + lanes[4] + lanes[5] + lanes[6] + lanes[7];
for (; i < n; ++i, ++a, ++b)
result += *a * *b;
return result;
}
// ---- Quantized-activation kernels ------------------------------------------
// The caller pre-quantizes the shared activation to int16 with one dequant
// scale per 32-element block (blockwise absmax keeps outlier features from
// crushing the resolution of the rest of the row). Weights decode to int16
// through pshufb LUTs and multiply-accumulate in integer pmaddwd lanes; fp32
// work is one convert+FMA per weight block. Four independent chains keep the
// FMA/madd pipelines busy. Weight scales combine with activation scales in
// the block epilogue: int8 storage keeps its per-row scale on the caller
// side, int4/fp4 multiply their block scale by the activation block scale.
float dot_int8_q8_avx2(const std::int8_t* weights, const std::int16_t* act,
const float* act_scales, std::size_t n) {
__m256i acc0 = _mm256_setzero_si256(), acc1 = _mm256_setzero_si256();
__m256i acc2 = _mm256_setzero_si256(), acc3 = _mm256_setzero_si256();
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 128 <= n; i += 128, weights += 128, act += 128) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 256), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act + 256), _MM_HINT_T0);
for (std::size_t k = 0; k < 4; ++k) {
__m256i* acc = k == 0 ? &acc0 : k == 1 ? &acc1 : k == 2 ? &acc2 : &acc3;
__m256* sum = k == 0 ? &sum0 : k == 1 ? &sum1 : k == 2 ? &sum2 : &sum3;
*acc = _mm256_add_epi32(*acc, _mm256_madd_epi16(
_mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 32 * k))),
_mm256_loadu_si256(reinterpret_cast<const __m256i*>(act + 32 * k))));
*acc = _mm256_add_epi32(*acc, _mm256_madd_epi16(
_mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 32 * k + 16))),
_mm256_loadu_si256(reinterpret_cast<const __m256i*>(act + 32 * k + 16))));
*sum = _mm256_fmadd_ps(_mm256_cvtepi32_ps(*acc),
_mm256_broadcast_ss(act_scales + i / 32 + k), *sum);
*acc = _mm256_setzero_si256();
}
}
sum0 = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
for (; i + 32 <= n; i += 32, weights += 32, act += 32) {
acc0 = _mm256_add_epi32(acc0, _mm256_madd_epi16(
_mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights))),
_mm256_loadu_si256(reinterpret_cast<const __m256i*>(act))));
acc0 = _mm256_add_epi32(acc0, _mm256_madd_epi16(
_mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16))),
_mm256_loadu_si256(reinterpret_cast<const __m256i*>(act + 16))));
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(acc0),
_mm256_broadcast_ss(act_scales + i / 32), sum0);
acc0 = _mm256_setzero_si256();
}
float result = reduce_sum(sum0);
if (i < n) {
// Partial tail block, read only to the row's logical end.
float block_sum = 0;
const float scale = act_scales[i / 32];
for (std::size_t j = 0; i < n; ++i, ++j)
block_sum += static_cast<float>(weights[j]) * static_cast<float>(act[j]);
result += block_sum * scale;
}
return result;
}
// Vectorized int8 activation quantizer (127/absmax per 32-element block).
// Same math as quantize_row_i8: absmax scale, round-half-to-even (sroundps
// follows MXCSR, the default nearest-even, matching lrint), packs saturation
// to int8. Scalar tail for partial blocks.
void quantize_row_i8_avx2(const float* input, std::size_t n, std::int8_t* values, float* scales) {
const __m256i sign = _mm256_set1_epi32(0x7fffffff);
std::size_t start = 0;
for (; start + 32 <= n; start += 32) {
__m256 a0 = _mm256_loadu_ps(input + start);
__m256 a1 = _mm256_loadu_ps(input + start + 8);
__m256 a2 = _mm256_loadu_ps(input + start + 16);
__m256 a3 = _mm256_loadu_ps(input + start + 24);
__m256 m0 = _mm256_and_ps(a0, _mm256_castsi256_ps(sign));
__m256 m1 = _mm256_and_ps(a1, _mm256_castsi256_ps(sign));
__m256 m2 = _mm256_and_ps(a2, _mm256_castsi256_ps(sign));
__m256 m3 = _mm256_and_ps(a3, _mm256_castsi256_ps(sign));
m0 = _mm256_max_ps(_mm256_max_ps(m0, m1), _mm256_max_ps(m2, m3));
m0 = _mm256_max_ps(m0, _mm256_permute2f128_ps(m0, m0, 1));
m0 = _mm256_max_ps(m0, _mm256_shuffle_ps(m0, m0, _MM_SHUFFLE(1, 0, 3, 2)));
m0 = _mm256_max_ps(m0, _mm256_shuffle_ps(m0, m0, _MM_SHUFFLE(2, 3, 0, 1)));
const float absmax = _mm256_cvtss_f32(m0);
if (absmax == 0.0f) {
_mm256_storeu_si256(reinterpret_cast<__m256i*>(values + start), _mm256_setzero_si256());
scales[start / 32] = 0.0f;
continue;
}
const __m256 mult = _mm256_set1_ps(127.0f / absmax);
__m256i q0 = _mm256_cvtps_epi32(_mm256_round_ps(_mm256_mul_ps(a0, mult),
_MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC));
__m256i q1 = _mm256_cvtps_epi32(_mm256_round_ps(_mm256_mul_ps(a1, mult),
_MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC));
__m256i q2 = _mm256_cvtps_epi32(_mm256_round_ps(_mm256_mul_ps(a2, mult),
_MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC));
__m256i q3 = _mm256_cvtps_epi32(_mm256_round_ps(_mm256_mul_ps(a3, mult),
_MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC));
const __m256i p01 = _mm256_packs_epi32(q0, q1);
const __m256i p23 = _mm256_packs_epi32(q2, q3);
const __m256i p = _mm256_packs_epi16(p01, p23);
// packs work within 128-bit lanes: 32-bit units arrive as
// [A0,B0,C0,D0,A1,B1,C1,D1]; permute to linear [A0,A1,B0,B1,...].
const __m256i idx = _mm256_setr_epi32(0, 4, 1, 5, 2, 6, 3, 7);
const __m256i out = _mm256_permutevar8x32_epi32(p, idx);
_mm256_storeu_si256(reinterpret_cast<__m256i*>(values + start), out);
scales[start / 32] = absmax / 127.0f;
}
for (; start < n; ++start) {
// Partial tail: same scalar math as quantize_row_i8.
float absmax = 0.0f;
const std::size_t bend = std::min(start + 32, n);
for (std::size_t i = start; i < bend; ++i) absmax = std::max(absmax, std::abs(input[i]));
if (absmax == 0.0f) {
for (std::size_t i = start; i < bend; ++i) values[i] = 0;
scales[start / 32] = 0.0f;
} else {
const float norm = 127.0f / absmax;
for (std::size_t i = start; i < bend; ++i) {
long q = std::lrint(static_cast<double>(input[i]) * norm);
values[i] = static_cast<std::int8_t>(std::clamp<long>(q, -127, 127));
}
scales[start / 32] = absmax / 127.0f;
}
start = bend - 1;
}
}
// first = w[0..15], second = w[16..31], linear order, dotted against linear
// int16 acts. Sub-8 replaces the LUT (codes 0..15 map to code-8 exactly).
static void unpack_int4_w16(const std::uint8_t* weights, const __m128i& mask, const __m128i& eight,
__m256i& first, __m256i& second) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i low = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i high = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
first = _mm256_cvtepi8_epi16(low);
second = _mm256_cvtepi8_epi16(high);
}
// One 32-weight int4 block, int16-quantized activation -> block dot in int32.
static __m256i int4_block_dot(const std::uint8_t* weights, const __m128i& mask, const __m128i& eight,
const std::int16_t* act) {
__m256i first, second;
unpack_int4_w16(weights, mask, eight, first, second);
return _mm256_add_epi32(
_mm256_madd_epi16(first, _mm256_loadu_si256(reinterpret_cast<const __m256i*>(act))),
_mm256_madd_epi16(second, _mm256_loadu_si256(reinterpret_cast<const __m256i*>(act + 16))));
}
float dot_int4_q8_avx2(const std::uint8_t* weights, const std::int16_t* act,
const float* scales, const float* act_scales, std::size_t n) {
const __m128i mask = _mm_set1_epi8(15);
const __m128i eight = _mm_set1_epi8(8);
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 128 <= n; i += 128, weights += 64, act += 128) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 128), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act + 256), _MM_HINT_T0);
const __m256 wscale0 = _mm256_mul_ps(_mm256_broadcast_ss(scales + i / 32),
_mm256_broadcast_ss(act_scales + i / 32));
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(int4_block_dot(weights, mask, eight, act)),
wscale0, sum0);
const __m256 wscale1 = _mm256_mul_ps(_mm256_broadcast_ss(scales + i / 32 + 1),
_mm256_broadcast_ss(act_scales + i / 32 + 1));
sum1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(int4_block_dot(weights + 16, mask, eight, act + 32)),
wscale1, sum1);
const __m256 wscale2 = _mm256_mul_ps(_mm256_broadcast_ss(scales + i / 32 + 2),
_mm256_broadcast_ss(act_scales + i / 32 + 2));
sum2 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(int4_block_dot(weights + 32, mask, eight, act + 64)),
wscale2, sum2);
const __m256 wscale3 = _mm256_mul_ps(_mm256_broadcast_ss(scales + i / 32 + 3),
_mm256_broadcast_ss(act_scales + i / 32 + 3));
sum3 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(int4_block_dot(weights + 48, mask, eight, act + 96)),
wscale3, sum3);
}
sum0 = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
for (; i + 32 <= n; i += 32, weights += 16, act += 32)
sum0 = _mm256_fmadd_ps(
_mm256_cvtepi32_ps(int4_block_dot(weights, mask, eight, act)),
_mm256_mul_ps(_mm256_broadcast_ss(scales + i / 32),
_mm256_broadcast_ss(act_scales + i / 32)), sum0);
float result = reduce_sum(sum0);
if (i < n) {
// Partial tail block, split-half block-relative indexing.
float block_sum = 0;
const float scale = scales[i / 32] * act_scales[i / 32];
const std::size_t m = n - i;
for (std::size_t t = 0; t < m; ++t) {
const int nibble = t < 16 ? (weights[t] & 15) : ((weights[t - 16] >> 4) & 15);
block_sum += static_cast<float>(nibble - 8) * static_cast<float>(act[t]);
}
result += block_sum * scale;
}
return result;
}
// One 16-weight fp4 block (8 bytes): nibbles -> int16 half values -> int32.
static __m256i fp4_block_dot(const std::uint8_t* weights, const __m128i& mask, const __m128i& lut,
const std::int16_t* act) {
const __m128i packed = _mm_loadl_epi64(reinterpret_cast<const __m128i*>(weights));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m256i w16 = _mm256_cvtepi8_epi16(_mm_shuffle_epi8(lut, _mm_unpacklo_epi8(low, high)));
return _mm256_madd_epi16(w16, _mm256_loadu_si256(reinterpret_cast<const __m256i*>(act)));
}
float dot_fp4_q8_avx2(const std::uint8_t* weights, const std::int16_t* act,
const std::uint8_t* scales, const float* act_scales, std::size_t n) {
const __m128i mask = _mm_set1_epi8(15);
const __m128i lut = _mm_loadu_si128(reinterpret_cast<const __m128i*>(fp4_element_lut()));
const float* scale_lut = fp4_scale_lut();
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
// Two 16-weight fp4 blocks per 32-weight activation block: the activation
// scale broadcast is shared, the weight scale comes from the E4M3 LUT.
for (; i + 64 <= n; i += 64, weights += 32, act += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 64), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act + 128), _MM_HINT_T0);
const float act_scale0 = act_scales[i / 32];
const __m256 ascale0 = _mm256_broadcast_ss(&act_scale0);
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights, mask, lut, act)),
_mm256_mul_ps(ascale0, _mm256_broadcast_ss(scale_lut + scales[i / 16])), sum0);
sum1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights + 8, mask, lut, act + 16)),
_mm256_mul_ps(ascale0, _mm256_broadcast_ss(scale_lut + scales[i / 16 + 1])), sum1);
const float act_scale1 = act_scales[i / 32 + 1];
const __m256 ascale1 = _mm256_broadcast_ss(&act_scale1);
sum2 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights + 16, mask, lut, act + 32)),
_mm256_mul_ps(ascale1, _mm256_broadcast_ss(scale_lut + scales[i / 16 + 2])), sum2);
sum3 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights + 24, mask, lut, act + 48)),
_mm256_mul_ps(ascale1, _mm256_broadcast_ss(scale_lut + scales[i / 16 + 3])), sum3);
}
sum0 = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
for (; i + 32 <= n; i += 32, weights += 16, act += 32) {
const __m256 ascale = _mm256_broadcast_ss(act_scales + i / 32);
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights, mask, lut, act)),
_mm256_mul_ps(ascale, _mm256_broadcast_ss(scale_lut + scales[i / 16])), sum0);
sum0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(fp4_block_dot(weights + 8, mask, lut, act + 16)),
_mm256_mul_ps(ascale, _mm256_broadcast_ss(scale_lut + scales[i / 16 + 1])), sum0);
}
float result = reduce_sum(sum0);
if (i < n) {
// Partial tail block, read only to the row's logical end.
const auto* elements = fp4_element_lut();
float block_sum = 0;
for (std::size_t j = 0; i < n; ++i, ++j)
block_sum += static_cast<float>(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) *
scale_lut[scales[i / 16]] * static_cast<float>(act[j]);
result += block_sum * act_scales[i / 32];
}
return result;
}
// Canonical 32-block deinterleave with AVX2 (generic /arch:AVX2, no VNNI).
// Per full [c,c+32): OUT[c+k]=IN[c+2k], OUT[c+16+k]=IN[c+2k+1]. Tail is left
// untouched (callers never read it). Scalar order, vector throughput: 4
// loads + 4 shuffles + 4 permutes + 4 stores per 32 vs 64 scalar copies.
void permute_act32_avx2(const float* input, std::size_t n, float* out) {
const __m256i idx = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
std::size_t c = 0;
for (; c + 32 <= n; c += 32) {
const __m256 in0 = _mm256_loadu_ps(input + c);
const __m256 in1 = _mm256_loadu_ps(input + c + 8);
const __m256 in2 = _mm256_loadu_ps(input + c + 16);
const __m256 in3 = _mm256_loadu_ps(input + c + 24);
const __m256 ev0 = _mm256_permutevar8x32_ps(
_mm256_shuffle_ps(in0, in1, _MM_SHUFFLE(2, 0, 2, 0)), idx);
const __m256 od0 = _mm256_permutevar8x32_ps(
_mm256_shuffle_ps(in0, in1, _MM_SHUFFLE(3, 1, 3, 1)), idx);
const __m256 ev1 = _mm256_permutevar8x32_ps(
_mm256_shuffle_ps(in2, in3, _MM_SHUFFLE(2, 0, 2, 0)), idx);
const __m256 od1 = _mm256_permutevar8x32_ps(
_mm256_shuffle_ps(in2, in3, _MM_SHUFFLE(3, 1, 3, 1)), idx);
_mm256_storeu_ps(out + c, ev0);
_mm256_storeu_ps(out + c + 8, ev1);
_mm256_storeu_ps(out + c + 16, od0);
_mm256_storeu_ps(out + c + 24, od1);
}
// Tail (<32) is intentionally untouched: callers never read PERM tails.
}
// Test-gated SiLU*up (same stable sigmoid as scalar; AVX2 TU only unrolls
// and prefetches because AVX2 has no vector exp — exp dominates, so this is
// not wired to the decode path; it exists for gated A/B only).
void silu_mul_avx2(const float* gate, const float* up, float* out, std::size_t n) {
std::size_t i = 0;
for (; i + 4 <= n; i += 4) {
_mm_prefetch(reinterpret_cast<const char*>(gate + i + 16), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(up + i + 16), _MM_HINT_T0);
for (int k = 0; k < 4; ++k) {
const float value = gate[i + k];
const float sigmoid = value >= 0 ? 1.0f / (1.0f + std::exp(-value)) :
std::exp(value) / (1.0f + std::exp(value));
out[i + k] = (value * sigmoid) * up[i + k];
}
}
for (; i < n; ++i) {
const float value = gate[i];
const float sigmoid = value >= 0 ? 1.0f / (1.0f + std::exp(-value)) :
std::exp(value) / (1.0f + std::exp(value));
out[i] = (value * sigmoid) * up[i];
}
}
// ---- Vector-exp activation blocks (decode hot path) ----
// Degree-6 minimax exp, ≤1 ULP vs libm over [-104, 88.7] (Remez-fit +
// float32 hill-climb; +inf above 88.7, 0 below -104, NaN passthrough).
// Deterministic: explicit FMA intrinsics, no tables, no data branches, so
// the same input bits always give the same output bits. Scalar tail twin
// below is lane-wise bit-identical to the vector lanes.
namespace {
inline __m256 vexp_poly6(__m256 x) {
__m256i n = _mm256_cvtps_epi32(_mm256_mul_ps(x, _mm256_set1_ps(1.4426950216f)));
__m256 nf = _mm256_cvtepi32_ps(n);
__m256 r = _mm256_fnmadd_ps(nf, _mm256_set1_ps(0.693359375f), x);
r = _mm256_fnmadd_ps(nf, _mm256_set1_ps(-2.1219444e-4f), r);
__m256 p = _mm256_set1_ps(0.0013963687233626842f);
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.00837346725165844f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.04166526347398758f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.16666468977928162f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.5000000596046448f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(1.0f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(1.0f));
__m256i nc = _mm256_min_epi32(n, _mm256_set1_epi32(127));
__m256 s = _mm256_castsi256_ps(
_mm256_slli_epi32(_mm256_add_epi32(nc, _mm256_set1_epi32(127)), 23));
__m256 y = _mm256_mul_ps(p, s);
y = _mm256_blendv_ps(y, _mm256_mul_ps(y, _mm256_set1_ps(2.0f)),
_mm256_castsi256_ps(_mm256_cmpeq_epi32(n, _mm256_set1_epi32(128))));
__m256 s2 = _mm256_castsi256_ps(
_mm256_slli_epi32(_mm256_add_epi32(n, _mm256_set1_epi32(151)), 23));
y = _mm256_blendv_ps(y,
_mm256_mul_ps(_mm256_mul_ps(p, s2), _mm256_set1_ps(5.960464477539063e-8f)),
_mm256_castsi256_ps(_mm256_cmpgt_epi32(_mm256_set1_epi32(-126), n)));
y = _mm256_blendv_ps(y, _mm256_set1_ps(std::numeric_limits<float>::infinity()),
_mm256_cmp_ps(x, _mm256_set1_ps(88.7f), _CMP_GT_OQ));
y = _mm256_blendv_ps(y, _mm256_set1_ps(0.0f),
_mm256_cmp_ps(x, _mm256_set1_ps(-104.0f), _CMP_LT_OQ));
y = _mm256_blendv_ps(y, x, _mm256_cmp_ps(x, x, _CMP_UNORD_Q));
return y;
}
// Scalar twin of vexp_poly6 (tail path; lane-wise bit-identical). fmaf
// inlines to vfmadd213ss under /arch:AVX2 (single rounding, like vfma ps).
inline float sexp_poly6(float x) {
if (x > 88.7f) return std::numeric_limits<float>::infinity();
if (x < -104.0f) return 0.0f;
if (x != x) return x;
int n = _mm_cvtss_si32(_mm_set_ss(x * 1.4426950216f));
float nf = static_cast<float>(n);
float r = std::fmaf(-nf, 0.693359375f, x);
r = std::fmaf(-nf, -2.1219444e-4f, r);
float p = 0.0013963687233626842f;
p = std::fmaf(p, r, 0.00837346725165844f);
p = std::fmaf(p, r, 0.04166526347398758f);
p = std::fmaf(p, r, 0.16666468977928162f);
p = std::fmaf(p, r, 0.5000000596046448f);
p = std::fmaf(p, r, 1.0f);
p = std::fmaf(p, r, 1.0f);
if (n > 127) return (p * 1.7014118346046923e38f) * 2.0f;
if (n < -126) {
float s2;
const std::uint32_t bits = static_cast<std::uint32_t>(n + 151) << 23;
std::memcpy(&s2, &bits, 4);
return (p * s2) * 5.960464477539063e-8f;
}
float s;
const std::uint32_t bits = static_cast<std::uint32_t>(n + 127) << 23;
std::memcpy(&s, &bits, 4);
return p * s;
}
// Unified stable sigmoid 1/(1+exp(-x)): exact for x>=0 (same ops as the
// scalar branch), ≤1 ULP elsewhere; safe at ±inf (no NaN: denom >= 1).
inline __m256 vsigmoid(__m256 x) {
__m256 e = vexp_poly6(_mm256_sub_ps(_mm256_setzero_ps(), x));
return _mm256_div_ps(_mm256_set1_ps(1.0f),
_mm256_add_ps(_mm256_set1_ps(1.0f), e));
}
inline float ssigmoid(float x) {
return 1.0f / (1.0f + sexp_poly6(-x));
}
} // namespace
void act_exp_avx2(float* x, std::size_t n) {
std::size_t i = 0;
for (; i + 8 <= n; i += 8) {
__m256 v = _mm256_loadu_ps(x + i);
_mm256_storeu_ps(x + i, vexp_poly6(v));
}
for (; i < n; ++i) x[i] = sexp_poly6(x[i]);
}
void act_silu_avx2(float* x, std::size_t n) {
std::size_t i = 0;
const __m256 one = _mm256_set1_ps(1.0f);
for (; i + 8 <= n; i += 8) {
__m256 v = _mm256_loadu_ps(x + i);
__m256 e = vexp_poly6(_mm256_sub_ps(_mm256_setzero_ps(), v));
__m256 y = _mm256_div_ps(v, _mm256_add_ps(one, e));
_mm256_storeu_ps(x + i, y);
}
for (; i < n; ++i) x[i] = x[i] * ssigmoid(x[i]);
}
void act_silu_mul_avx2(const float* gate, const float* up, float* out, std::size_t n) {
// Same per-element order as the scalar reference: clamp, silu, mul.
// out may alias gate (same-index read/write only, no cross-lane reuse).
std::size_t i = 0;
const __m256 lo = _mm256_set1_ps(-15.0f), hi = _mm256_set1_ps(15.0f);
const __m256 one = _mm256_set1_ps(1.0f);
for (; i + 8 <= n; i += 8) {
__m256 g = _mm256_loadu_ps(gate + i);
__m256 u = _mm256_loadu_ps(up + i);
g = _mm256_min_ps(_mm256_max_ps(g, lo), hi);
__m256 e = vexp_poly6(_mm256_sub_ps(_mm256_setzero_ps(), g));
__m256 y = _mm256_div_ps(g, _mm256_add_ps(one, e));
_mm256_storeu_ps(out + i, _mm256_mul_ps(y, u));
}
for (; i < n; ++i) {
float gv = gate[i];
if (gv < -15.0f) gv = -15.0f;
else if (gv > 15.0f) gv = 15.0f;
out[i] = gv * ssigmoid(gv) * up[i];
}
}
void act_sigmoid_avx2(float* x, std::size_t n) {
std::size_t i = 0;
for (; i + 8 <= n; i += 8) {
__m256 v = _mm256_loadu_ps(x + i);
_mm256_storeu_ps(x + i, vsigmoid(v));
}
for (; i < n; ++i) x[i] = ssigmoid(x[i]);
}
void act_sigmoid_mul_avx2(float* o, const float* g, std::size_t n) {
// o may not alias g (callers pass distinct buffers).
std::size_t i = 0;
for (; i + 8 <= n; i += 8) {
__m256 ov = _mm256_loadu_ps(o + i);
__m256 gv = _mm256_loadu_ps(g + i);
_mm256_storeu_ps(o + i, _mm256_mul_ps(ov, vsigmoid(gv)));
}
for (; i < n; ++i) o[i] *= ssigmoid(g[i]);
}
void act_silu_mul_plain_avx2(const float* gate, const float* up, float* out, std::size_t n) {
// Dense path: stable SiLU without the Surjo +-15 clamp.
std::size_t i = 0;
const __m256 one = _mm256_set1_ps(1.0f);
for (; i + 8 <= n; i += 8) {
__m256 g = _mm256_loadu_ps(gate + i);
__m256 u = _mm256_loadu_ps(up + i);
__m256 e = vexp_poly6(_mm256_sub_ps(_mm256_setzero_ps(), g));
__m256 y = _mm256_div_ps(g, _mm256_add_ps(one, e));
_mm256_storeu_ps(out + i, _mm256_mul_ps(y, u));
}
for (; i < n; ++i) {
float gv = gate[i];
out[i] = gv * ssigmoid(gv) * up[i];
}
}
// ---- Accurate single-precision erf/GELU (FWKV decode hot path) ----
// Split-region branchless design (small: float64 Horner; mid: exp*R(t);
// |x|>=4: +-1; NaN payload preserved), ≤1.6 ULP vs double truth over
// [-6,6]; scalar tail twin is lane-wise bit-identical to the vector lanes.
// GELU tail note: 0.5*x*(1+erf) cancels for x<<0 (large relative ULP on
// ~1e-9 values, bit-identical) — absolute error stays ~1e-7, PPL-gated.
namespace {
// Coefficients: least-squares fits + float32-ULP coordinate descent.
constexpr double kErfS0 = 1.1283791670946941;
constexpr double kErfS1 = -0.3761263888986921;
constexpr double kErfS2 = 0.11283791313498021;
constexpr double kErfS3 = -0.026866133625394584;
constexpr double kErfS4 = 0.0052237850067526487;
constexpr double kErfS5 = -0.0008542670561199478;
constexpr double kErfS6 = 0.00011956946752664233;
constexpr double kErfS7 = -1.3911700737181953e-05;
constexpr double kErfS8 = 1.0595274867464255e-06;
constexpr float kErfPmid = 0.3275911f;
constexpr float kErfM0 = 0.0020323586650192738f;
constexpr float kErfM1 = 0.156896248f;
constexpr float kErfM2 = 0.3502644f;
constexpr float kErfM3 = -0.373835027217865f;
constexpr float kErfM4 = 1.2575676441192627f;
constexpr float kErfM5 = -1.2146756649017334f;
constexpr float kErfM6 = 0.9994310736656189f;
constexpr float kErfM7 = -0.17759732902050018f;
inline float erf_exp_unit(float x) {
float nf = std::floor(std::fmaf(x, 1.4426950408889634f, 0.5f));
nf = (!(nf >= -30.0f)) ? -30.0f : nf;
nf = (!(nf <= 127.0f)) ? 127.0f : nf;
const auto n = static_cast<std::int32_t>(nf);
const float fn = nf;
float r = std::fmaf(-fn, 0.693145751953125f, x);
r = std::fmaf(-fn, 1.428606765330187045e-06f, r);
float p = 0.001394858118146658f;
p = std::fmaf(p, r, 0.008375128731131554f);
p = std::fmaf(p, r, 0.041666217148303986f);
p = std::fmaf(p, r, 0.16666415333747864f);
p = std::fmaf(p, r, 0.5f);
p = std::fmaf(p, r, 1.0f);
p = std::fmaf(p, r, 1.0f);
const std::uint32_t sbits = static_cast<std::uint32_t>(n + 127) << 23;
float s;
std::memcpy(&s, &sbits, 4);
return p * s;
}
inline float serf_f32(float x) {
std::uint32_t ux;
std::memcpy(&ux, &x, 4);
const std::uint32_t axu = ux & 0x7FFFFFFFu;
float ax;
std::memcpy(&ax, &axu, 4);
const std::uint32_t sxu = (ux & 0x80000000u) | 0x3F800000u;
float sx;
std::memcpy(&sx, &sxu, 4);
const double axd = static_cast<double>(ax);
const double zd = axd * axd;
const float z = ax * ax;
double psd = kErfS8;
psd = std::fma(psd, zd, kErfS7);
psd = std::fma(psd, zd, kErfS6);
psd = std::fma(psd, zd, kErfS5);
psd = std::fma(psd, zd, kErfS4);
psd = std::fma(psd, zd, kErfS3);
psd = std::fma(psd, zd, kErfS2);
psd = std::fma(psd, zd, kErfS1);
psd = std::fma(psd, zd, kErfS0);
const float ps = static_cast<float>(psd);
const float e_small = x * ps;
const float ex = erf_exp_unit(-z);
const float t = 1.0f / std::fmaf(kErfPmid, ax, 1.0f);
float q = kErfM7;
q = std::fmaf(q, t, kErfM6);
q = std::fmaf(q, t, kErfM5);
q = std::fmaf(q, t, kErfM4);
q = std::fmaf(q, t, kErfM3);
q = std::fmaf(q, t, kErfM2);
q = std::fmaf(q, t, kErfM1);
q = std::fmaf(q, t, kErfM0);
const float e_mid = sx * (1.0f - ex * q);
const std::uint32_t m_big = static_cast<std::uint32_t>(-static_cast<std::int32_t>(ax >= 4.0f));
const std::uint32_t m_small = static_cast<std::uint32_t>(-static_cast<std::int32_t>(ax <= 1.0f));
const std::uint32_t m_nan = static_cast<std::uint32_t>(-static_cast<std::int32_t>(ax != ax));
std::uint32_t a, b, r;
std::memcpy(&a, &e_mid, 4);
std::memcpy(&b, &sx, 4);
r = (a & ~m_big) | (b & m_big);
float e_midbig;
std::memcpy(&e_midbig, &r, 4);
std::memcpy(&a, &e_midbig, 4);
std::memcpy(&b, &e_small, 4);
r = (a & ~m_small) | (b & m_small);
float res;
std::memcpy(&res, &r, 4);
std::memcpy(&a, &res, 4);
r = (a & ~m_nan) | (ux & m_nan);
std::memcpy(&res, &r, 4);
return res;
}
inline float sgelu_f32(float x) {
const float y = x * 0.7071067811865475f;
const float e = serf_f32(y);
return (0.5f * x) * (1.0f + e);
}
inline __m256 verf_exp_unit(__m256 x) {
__m256 nf = _mm256_floor_ps(_mm256_fmadd_ps(x, _mm256_set1_ps(1.4426950408889634f),
_mm256_set1_ps(0.5f)));
nf = _mm256_min_ps(_mm256_max_ps(nf, _mm256_set1_ps(-30.0f)), _mm256_set1_ps(127.0f));
__m256i ni = _mm256_cvtps_epi32(nf);
__m256 fn = nf;
__m256 r = _mm256_fnmadd_ps(fn, _mm256_set1_ps(0.693145751953125f), x);
r = _mm256_fnmadd_ps(fn, _mm256_set1_ps(1.428606765330187045e-06f), r);
__m256 p = _mm256_set1_ps(0.001394858118146658f);
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.008375128731131554f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.041666217148303986f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.16666415333747864f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(0.5f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(1.0f));
p = _mm256_fmadd_ps(p, r, _mm256_set1_ps(1.0f));
__m256 s = _mm256_castsi256_ps(
_mm256_slli_epi32(_mm256_add_epi32(ni, _mm256_set1_epi32(127)), 23));
return _mm256_mul_ps(p, s);
}
inline __m256 verf_f32(__m256 x) {
const __m256i absm = _mm256_set1_epi32(0x7FFFFFFF);
const __m256i sgnm = _mm256_set1_epi32(0x80000000);
const __m256 one = _mm256_set1_ps(1.0f);
__m256i xi = _mm256_castps_si256(x);
__m256 ax = _mm256_castsi256_ps(_mm256_and_si256(xi, absm));
__m256 sx = _mm256_castsi256_ps(
_mm256_or_si256(_mm256_and_si256(xi, sgnm), _mm256_castps_si256(one)));
__m256 z = _mm256_mul_ps(ax, ax);
__m128 ax_lo = _mm256_castps256_ps128(ax);
__m128 ax_hi = _mm256_extractf128_ps(ax, 1);
__m256d axd_lo = _mm256_cvtps_pd(ax_lo);
__m256d axd_hi = _mm256_cvtps_pd(ax_hi);
__m256d zd_lo = _mm256_mul_pd(axd_lo, axd_lo);
__m256d zd_hi = _mm256_mul_pd(axd_hi, axd_hi);
__m256d psd_lo = _mm256_set1_pd(kErfS8);
__m256d psd_hi = _mm256_set1_pd(kErfS8);
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS7));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS7));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS6));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS6));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS5));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS5));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS4));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS4));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS3));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS3));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS2));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS2));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS1));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS1));
psd_lo = _mm256_fmadd_pd(psd_lo, zd_lo, _mm256_set1_pd(kErfS0));
psd_hi = _mm256_fmadd_pd(psd_hi, zd_hi, _mm256_set1_pd(kErfS0));
__m128 ps_lo = _mm256_cvtpd_ps(psd_lo);
__m128 ps_hi = _mm256_cvtpd_ps(psd_hi);
__m256 ps = _mm256_insertf128_ps(_mm256_castps128_ps256(ps_lo), ps_hi, 1);
__m256 e_small = _mm256_mul_ps(x, ps);
__m256 ex = verf_exp_unit(_mm256_sub_ps(_mm256_setzero_ps(), z));
__m256 t = _mm256_div_ps(one, _mm256_fmadd_ps(_mm256_set1_ps(kErfPmid), ax, one));
__m256 q = _mm256_set1_ps(kErfM7);
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM6));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM5));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM4));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM3));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM2));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM1));
q = _mm256_fmadd_ps(q, t, _mm256_set1_ps(kErfM0));
__m256 e_mid = _mm256_mul_ps(sx, _mm256_sub_ps(one, _mm256_mul_ps(ex, q)));
__m256 big_m = _mm256_cmp_ps(ax, _mm256_set1_ps(4.0f), _CMP_GE_OQ);
__m256 sml_m = _mm256_cmp_ps(ax, _mm256_set1_ps(1.0f), _CMP_LE_OQ);
__m256 nan_m = _mm256_cmp_ps(ax, ax, _CMP_NEQ_UQ);
__m256 e_midbig = _mm256_blendv_ps(e_mid, sx, big_m);
__m256 res = _mm256_blendv_ps(e_midbig, e_small, sml_m);
res = _mm256_blendv_ps(res, x, nan_m);
return res;
}
inline __m256 vgelu_f32(__m256 x) {
__m256 y = _mm256_mul_ps(x, _mm256_set1_ps(0.7071067811865475f));
__m256 e = verf_f32(y);
__m256 h = _mm256_mul_ps(_mm256_set1_ps(0.5f), x);
return _mm256_mul_ps(h, _mm256_add_ps(_mm256_set1_ps(1.0f), e));
}
} // namespace
void act_gelu_avx2(float* x, std::size_t n) {
std::size_t i = 0;
for (; i + 8 <= n; i += 8) {
__m256 v = _mm256_loadu_ps(x + i);
_mm256_storeu_ps(x + i, vgelu_f32(v));
}
for (; i < n; ++i) x[i] = sgelu_f32(x[i]);
}
} // namespace cism