Spaces:
Sleeping
Sleeping
Download native/kernels_avx2.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 55.2 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_avx2.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/kernels_avx2.cpp
-
curl -L -o kernels_avx2.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_avx2.cpp
55.2 kB
| 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)); | |
| } | |
| // 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; | |
| } | |
| 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 | |