#include "kernels.hpp" #include #include #include #include #include #include #include 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(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(weights + 1024), _MM_HINT_T0); sum0 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights))), _mm256_loadu_ps(input), sum0); sum1 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 8))), _mm256_loadu_ps(input + 8), sum1); sum2 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 16))), _mm256_loadu_ps(input + 16), sum2); sum3 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 24))), _mm256_loadu_ps(input + 24), sum3); sum4 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 32))), _mm256_loadu_ps(input + 32), sum4); sum5 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 40))), _mm256_loadu_ps(input + 40), sum5); sum6 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(weights + 48))), _mm256_loadu_ps(input + 48), sum6); sum7 = _mm256_fmadd_ps(_mm256_cvtph_ps(_mm_loadu_si128(reinterpret_cast(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(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(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(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(*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(weights + 32), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(act_orig + 64), _MM_HINT_T0); { const __m128i packed = _mm_loadu_si128(reinterpret_cast(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(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(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(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(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(weights + 32), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(act_perm + 64), _MM_HINT_T0); { const __m128i packed = _mm_loadu_si128(reinterpret_cast(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(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(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(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(a + 128), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(weights + 256), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(weights + 32 * k))), _mm256_loadu_si256(reinterpret_cast(act + 32 * k)))); *acc = _mm256_add_epi32(*acc, _mm256_madd_epi16( _mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast(weights + 32 * k + 16))), _mm256_loadu_si256(reinterpret_cast(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(weights))), _mm256_loadu_si256(reinterpret_cast(act)))); acc0 = _mm256_add_epi32(acc0, _mm256_madd_epi16( _mm256_cvtepi8_epi16(_mm_loadu_si128(reinterpret_cast(weights + 16))), _mm256_loadu_si256(reinterpret_cast(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(weights[j]) * static_cast(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(input[i]) * norm); values[i] = static_cast(std::clamp(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(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(act))), _mm256_madd_epi16(second, _mm256_loadu_si256(reinterpret_cast(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(weights + 128), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(nibble - 8) * static_cast(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(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(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(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(weights + 64), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) * scale_lut[scales[i / 16]] * static_cast(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(gate + i + 16), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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::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::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(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(n + 151) << 23; std::memcpy(&s2, &bits, 4); return (p * s2) * 5.960464477539063e-8f; } float s; const std::uint32_t bits = static_cast(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(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(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(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(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(-static_cast(ax >= 4.0f)); const std::uint32_t m_small = static_cast(-static_cast(ax <= 1.0f)); const std::uint32_t m_nan = static_cast(-static_cast(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