Spaces:
Sleeping
Sleeping
Download native/kernels_vnni.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 17.8 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_vnni.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/kernels_vnni.cpp
-
curl -L -o kernels_vnni.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_vnni.cpp
17.8 kB
| // AVX-VNNI / AVX512-VNNI integer-MAC path (Zen4/Zen5, Alder Lake+, SPR). | |
| // | |
| // One kernel: int8 weights x int8 activations (the VNNI-native 8-bit ABI) | |
| // with per-32 fp32 dequant scales, selected only when the CPU has the exact | |
| // encoding this TU was built for (VEX dpbusd on GCC -mavxvnni, EVEX dpbusd | |
| // on MSVC /arch:AVX512). Pre-VNNI machines never execute this file: the | |
| // int16 Q8 AVX2 path (kernels_avx2.cpp) stays their fast path and scalar | |
| // stays the fallback. | |
| // | |
| // Math (per 32-block, integer-exact): dpbusd needs (u8, s8), weights are | |
| // s8, so bias weights by +128 (xor 0x80) and correct: | |
| // sum(w*a) = dpbusd(w^0x80, a) - 128*sum(a). | |
| // int32 is exact here (|sum| <= 32*127*127 = 516k, |lane| <= ~131k), so any | |
| // horizontal-sum order is bit-identical and the fused single-hsum form below | |
| // equals the old two-hsum form lane-for-lane (checked in numpy over 20k | |
| // random blocks; test_dpbusd_bias_correction_math covers the identity). | |
| // Float accumulation uses 4 independent chains like the AVX2 Q8 kernels; | |
| // order differs from AVX2 (scalar mul+add vs FMA), so numeric changes stay | |
| // PPL-gated (int8 PPL gate holds at +0.3% on the AVX2 reference; VNNI | |
| // hardware CI must re-confirm, see VALIDATION.md). Weight row scales stay on | |
| // the caller side, mirroring dot_int8_q8_avx2. | |
| // | |
| // AVX2-standard audit (against dot_int8_q8_avx2 in kernels_avx2.cpp): | |
| // 128-wide outer loop (4x32, 4 float chains in fixed order) + 64-wide | |
| // (2x32, sequential adds, bitwise-identical to two 32-iterations) + 32-wide | |
| // + scalar tail matches the Q8-family blocking; the 64-wide step is the Q8 | |
| // analogue of the fp32 kernels' 64-wide inner blocking (halves tail-loop | |
| // overhead, no float-order change). T0 prefetches (+256 elements on both | |
| // int8 streams, 2 iterations ahead) mirror dot_int8_q8_avx2 exactly (its | |
| // int16 act prefetch covers twice the bytes for the same 2-iteration lead); | |
| // x86 prefetches never fault, skipped when n < 128. The partial tail reads | |
| // only to the row's logical end with act_scales[i/32]. | |
| // | |
| // Provenance / performance note: the integer math is proven in numpy on any | |
| // CPU; the intrinsic mapping itself is VNNI-hardware-CI-pending (this box is | |
| // Zen3, no VNNI). Single-token decode stays bandwidth-bound (weights stream | |
| // from DRAM either way), so the dpbusd win concentrates in prefill and | |
| // spec-verify GEMM work (tokens resident, weights reused from cache). | |
| namespace cism { | |
| namespace { | |
| // Register-only horizontal sum of 8 int32 lanes (deterministic, no | |
| // store-forwarding stall). Reordering vs a stored-lanes sum is bit-identical: | |
| // int32 never overflows here, so addition is associative. | |
| inline int hsum_epi32(__m256i v) { | |
| __m128i s = _mm_add_epi32(_mm256_castsi256_si128(v), _mm256_extracti128_si256(v, 1)); | |
| s = _mm_add_epi32(s, _mm_shuffle_epi32(s, _MM_SHUFFLE(1, 0, 3, 2))); | |
| s = _mm_add_epi32(s, _mm_shuffle_epi32(s, _MM_SHUFFLE(0, 0, 0, 1))); | |
| return _mm_cvtsi128_si32(s); | |
| } | |
| // Exact integer block dot for 32 int8 pairs via one dpbusd + correction. | |
| // bias/ones are hoisted by the caller (loop-invariant broadcasts). sum(a) is | |
| // derived from the already-loaded act register (extract halves, no act | |
| // reloads) and the -128*sum(a) correction is fused into the dp lanes | |
| // (slli x128 + sub) so a single hsum finishes the block. A psadbw-based | |
| // sum(a) was considered and rejected: it needs an extra act xor for the same | |
| // single hsum, saving nothing after this fusion. | |
| inline int block_dot32_vnni(const std::int8_t* weights, const std::int8_t* act, | |
| __m256i bias, __m256i ones) { | |
| const __m256i w = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(weights)); | |
| const __m256i a = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(act)); | |
| const __m256i dp = _mm256_dpbusd_epi32( | |
| _mm256_setzero_si256(), _mm256_xor_si256(w, bias), a); | |
| const __m256i s = _mm256_add_epi32( | |
| _mm256_madd_epi16(_mm256_cvtepi8_epi16(_mm256_castsi256_si128(a)), ones), | |
| _mm256_madd_epi16(_mm256_cvtepi8_epi16(_mm256_extracti128_si256(a, 1)), ones)); | |
| return hsum_epi32(_mm256_sub_epi32(dp, _mm256_slli_epi32(s, 7))); | |
| } | |
| } // namespace | |
| // AVX512 zmm fp32 dot (needs only AVX512F; gated on VNNI presence like the | |
| // rest of this TU to keep one dispatch story). 8 accumulators over 128-wide | |
| // blocks mirror the AVX2 kernel's structure; chunking (128/64/32/16) and the | |
| // final horizontal+scalar reduce differ, so results match AVX2 within normal | |
| // fp32 rounding (PPL-gated, NOT bitwise-identical — validate on VNNI tin). | |
| // VL-free reduce (split to 256-bit halves) so GCC needs no extra -mavx512vl. | |
| static float reduce_sum512(__m512 value) { | |
| const __m256 added = _mm256_add_ps(_mm512_castps512_ps256(value), | |
| _mm256_castpd_ps(_mm512_extractf64x4_pd(_mm512_castps_pd(value), 1))); | |
| __m128 sum = _mm_add_ps(_mm256_castps256_ps128(added), _mm256_extractf128_ps(added, 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); | |
| } | |
| float dot_fp32_avx512(const float* a, const float* b, std::size_t n) { | |
| __m512 sum0 = _mm512_setzero_ps(), sum1 = _mm512_setzero_ps(); | |
| __m512 sum2 = _mm512_setzero_ps(), sum3 = _mm512_setzero_ps(); | |
| __m512 sum4 = _mm512_setzero_ps(), sum5 = _mm512_setzero_ps(); | |
| __m512 sum6 = _mm512_setzero_ps(), sum7 = _mm512_setzero_ps(); | |
| std::size_t i = 0; | |
| for (; i + 128 <= n; i += 128, a += 128, b += 128) { | |
| _mm_prefetch(reinterpret_cast<const char*>(a + 64), _MM_HINT_T0); | |
| _mm_prefetch(reinterpret_cast<const char*>(b + 64), _MM_HINT_T0); | |
| sum0 = _mm512_fmadd_ps(_mm512_loadu_ps(a), _mm512_loadu_ps(b), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 16), _mm512_loadu_ps(b + 16), sum1); | |
| sum2 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 32), _mm512_loadu_ps(b + 32), sum2); | |
| sum3 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 48), _mm512_loadu_ps(b + 48), sum3); | |
| sum4 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 64), _mm512_loadu_ps(b + 64), sum4); | |
| sum5 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 80), _mm512_loadu_ps(b + 80), sum5); | |
| sum6 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 96), _mm512_loadu_ps(b + 96), sum6); | |
| sum7 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 112), _mm512_loadu_ps(b + 112), sum7); | |
| } | |
| for (; i + 64 <= n; i += 64, a += 64, b += 64) { | |
| sum0 = _mm512_fmadd_ps(_mm512_loadu_ps(a), _mm512_loadu_ps(b), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 16), _mm512_loadu_ps(b + 16), sum1); | |
| sum2 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 32), _mm512_loadu_ps(b + 32), sum2); | |
| sum3 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 48), _mm512_loadu_ps(b + 48), sum3); | |
| } | |
| for (; i + 32 <= n; i += 32, a += 32, b += 32) { | |
| sum0 = _mm512_fmadd_ps(_mm512_loadu_ps(a), _mm512_loadu_ps(b), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_loadu_ps(a + 16), _mm512_loadu_ps(b + 16), sum1); | |
| } | |
| for (; i + 16 <= n; i += 16, a += 16, b += 16) | |
| sum0 = _mm512_fmadd_ps(_mm512_loadu_ps(a), _mm512_loadu_ps(b), sum0); | |
| __m512 sum = _mm512_add_ps(_mm512_add_ps(_mm512_add_ps(sum0, sum1), _mm512_add_ps(sum2, sum3)), | |
| _mm512_add_ps(_mm512_add_ps(sum4, sum5), _mm512_add_ps(sum6, sum7))); | |
| float result = reduce_sum512(sum); | |
| for (; i < n; ++i, ++a, ++b) result += *a * *b; | |
| return result; | |
| } | |
| // AVX512 zmm int8 dot (fp32 activations): 64-byte loads + widen + FMA at 2x | |
| // the AVX2 width, 8 accumulators, masked <16 tail (maskz loads cannot fault). | |
| // Same per-element math as dot_int8_avx2; chunking/reduce differ (PPL-gated). | |
| float dot_int8_avx512(const std::int8_t* weights, const float* input, std::size_t n) { | |
| __m512 sum0 = _mm512_setzero_ps(), sum1 = _mm512_setzero_ps(); | |
| __m512 sum2 = _mm512_setzero_ps(), sum3 = _mm512_setzero_ps(); | |
| __m512 sum4 = _mm512_setzero_ps(), sum5 = _mm512_setzero_ps(); | |
| __m512 sum6 = _mm512_setzero_ps(), sum7 = _mm512_setzero_ps(); | |
| std::size_t i = 0; | |
| for (; i + 128 <= n; i += 128, weights += 128, input += 128) { | |
| _mm_prefetch(reinterpret_cast<const char*>(weights + 256), _MM_HINT_T0); | |
| _mm_prefetch(reinterpret_cast<const char*>(input + 64), _MM_HINT_T0); | |
| sum0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights)))), | |
| _mm512_loadu_ps(input), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16)))), | |
| _mm512_loadu_ps(input + 16), sum1); | |
| sum2 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 32)))), | |
| _mm512_loadu_ps(input + 32), sum2); | |
| sum3 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 48)))), | |
| _mm512_loadu_ps(input + 48), sum3); | |
| sum4 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 64)))), | |
| _mm512_loadu_ps(input + 64), sum4); | |
| sum5 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 80)))), | |
| _mm512_loadu_ps(input + 80), sum5); | |
| sum6 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 96)))), | |
| _mm512_loadu_ps(input + 96), sum6); | |
| sum7 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 112)))), | |
| _mm512_loadu_ps(input + 112), sum7); | |
| } | |
| for (; i + 64 <= n; i += 64, weights += 64, input += 64) { | |
| sum0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights)))), | |
| _mm512_loadu_ps(input), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16)))), | |
| _mm512_loadu_ps(input + 16), sum1); | |
| sum2 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 32)))), | |
| _mm512_loadu_ps(input + 32), sum2); | |
| sum3 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 48)))), | |
| _mm512_loadu_ps(input + 48), sum3); | |
| } | |
| for (; i + 32 <= n; i += 32, weights += 32, input += 32) { | |
| sum0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights)))), | |
| _mm512_loadu_ps(input), sum0); | |
| sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + 16)))), | |
| _mm512_loadu_ps(input + 16), sum1); | |
| } | |
| for (; i + 16 <= n; i += 16, weights += 16, input += 16) | |
| sum0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( | |
| _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights)))), | |
| _mm512_loadu_ps(input), sum0); | |
| __m512 sum = _mm512_add_ps(_mm512_add_ps(_mm512_add_ps(sum0, sum1), _mm512_add_ps(sum2, sum3)), | |
| _mm512_add_ps(_mm512_add_ps(sum4, sum5), _mm512_add_ps(sum6, sum7))); | |
| float result = reduce_sum512(sum); | |
| if (i < n) { | |
| // Masked <16 tail: maskz loads cannot fault, same math as scalar. | |
| // Heap rows may end at a page boundary, so no unmasked over-read. | |
| const std::size_t k = n - i; | |
| const __mmask16 m = static_cast<__mmask16>((1u << k) - 1u); | |
| __m128i wb = _mm_maskz_loadu_epi8(m, weights); | |
| __m512 wf = _mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(wb)); | |
| __m512 a = _mm512_maskz_loadu_ps(m, input); | |
| result += reduce_sum512(_mm512_mul_ps(wf, a)); | |
| } | |
| return result; | |
| } | |
| float dot_int8_q8_vnni(const std::int8_t* weights, const std::int8_t* act, | |
| const float* act_scales, std::size_t n) { | |
| float sum0 = 0, sum1 = 0, sum2 = 0, sum3 = 0; | |
| std::size_t i = 0; | |
| const __m256i bias = _mm256_set1_epi8(static_cast<char>(0x80)); | |
| const __m256i ones = _mm256_set1_epi16(1); | |
| 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); | |
| sum0 += static_cast<float>(block_dot32_vnni(weights, act, bias, ones)) * act_scales[i / 32]; | |
| sum1 += static_cast<float>(block_dot32_vnni(weights + 32, act + 32, bias, ones)) * act_scales[i / 32 + 1]; | |
| sum2 += static_cast<float>(block_dot32_vnni(weights + 64, act + 64, bias, ones)) * act_scales[i / 32 + 2]; | |
| sum3 += static_cast<float>(block_dot32_vnni(weights + 96, act + 96, bias, ones)) * act_scales[i / 32 + 3]; | |
| } | |
| float result = (sum0 + sum1) + (sum2 + sum3); | |
| for (; i + 64 <= n; i += 64, weights += 64, act += 64) { | |
| // 64-wide fallback: two 32-blocks in order, sequential scalar adds. | |
| // Bitwise-identical to two 32-wide iterations; halves tail-loop | |
| // overhead (Q8 analogue of the fp32 kernels' 64-wide blocking). No | |
| // prefetch: tail-resident, mirrors AVX2 Q8 (prefetch only in 128-loop). | |
| result += static_cast<float>(block_dot32_vnni(weights, act, bias, ones)) * act_scales[i / 32]; | |
| result += static_cast<float>(block_dot32_vnni(weights + 32, act + 32, bias, ones)) * act_scales[i / 32 + 1]; | |
| } | |
| for (; i + 32 <= n; i += 32, weights += 32, act += 32) | |
| result += static_cast<float>(block_dot32_vnni(weights, act, bias, ones)) * act_scales[i / 32]; | |
| if (i < n) { | |
| // Partial tail block, read only to the row's logical end. The scale | |
| // is hoisted before the loop (start is 32-aligned with <32 left, so | |
| // this is the same block index the post-loop form read). | |
| const float scale = act_scales[i / 32]; | |
| int block_sum = 0; | |
| for (std::size_t j = 0; i < n; ++i, ++j) | |
| block_sum += static_cast<int>(weights[j]) * static_cast<int>(act[j]); | |
| result += static_cast<float>(block_sum) * scale; | |
| } | |
| return result; | |
| } | |
| // int4-VNNI / fp4-VNNI: deliberately deferred (no kernels.hpp signature | |
| // change, no dead code) — int4/fp4 stay on the AVX2 Q8 path. Both would fit | |
| // this dpbusd ABI via pshufb LUT-decode to int8 then block_dot32_vnni above | |
| // (int4 nibbles -> {-8..7}; fp4 E2M1 -> half-values 0,+-1,+-2,+-3,+-4,+-6, | |
| // +-8,+-12 with the pre-halved E4M3 scale, exact like the AVX2 fp4 path), | |
| // but the LUT cost dominates: the pshufb decode (mask+srli+shuffle+unpack) | |
| // is identical on both paths and dpbusd only replaces the final pmaddwd, | |
| // while a VNNI 4-bit kernel would additionally need an int16->int8 act | |
| // down-convert (Int4Q8/Fp4Q8 ABI takes int16 acts; no int8-act 4-bit | |
| // signature exists and quantize_row_i8/dispatch wiring lives in | |
| // runtime.cpp, outside this TU's ownership) plus a two-scale (per-16 weight | |
| // + per-32 act) epilogue per 32-block for fp4. Net: no win without a Matrix | |
| // ABI change + new dispatch, and no VNNI silicon on this box to validate an | |
| // otherwise unwired kernel. Revisit with Zen4/ADL silicon + PPL re-gate. | |
| // Compiler without VNNI support: scalar-correct implementation so the symbol | |
| // always links. Dispatch (kernels.cpp) never selects it without VNNI macros, | |
| // so this is dead code on such builds — correctness over speed here. | |
| float dot_int8_q8_vnni(const std::int8_t* weights, const std::int8_t* act, | |
| const float* act_scales, std::size_t n) { | |
| float result = 0; | |
| std::size_t i = 0; | |
| for (; i + 32 <= n; i += 32, weights += 32, act += 32) { | |
| int block_sum = 0; | |
| for (std::size_t j = 0; j < 32; ++j) | |
| block_sum += static_cast<int>(weights[j]) * static_cast<int>(act[j]); | |
| result += static_cast<float>(block_sum) * act_scales[i / 32]; | |
| } | |
| if (i < n) { | |
| // Partial tail, read only to the row's logical end. Scale hoisted | |
| // before the loop (start is 32-aligned with <32 left, so this equals | |
| // the post-loop act_scales[i/32]); mirrors the VNNI fast path. | |
| const float scale = act_scales[i / 32]; | |
| int block_sum = 0; | |
| for (std::size_t j = 0; i < n; ++i, ++j) | |
| block_sum += static_cast<int>(weights[j]) * static_cast<int>(act[j]); | |
| result += static_cast<float>(block_sum) * scale; | |
| } | |
| return result; | |
| } | |
| } // namespace cism | |