// 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). #include "kernels.hpp" #include #include #if defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI) #include #endif namespace cism { #if defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI) 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(weights)); const __m256i a = _mm256_loadu_si256(reinterpret_cast(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(a + 64), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(weights + 256), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(input + 64), _MM_HINT_T0); sum0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights)))), _mm512_loadu_ps(input), sum0); sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 16)))), _mm512_loadu_ps(input + 16), sum1); sum2 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 32)))), _mm512_loadu_ps(input + 32), sum2); sum3 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 48)))), _mm512_loadu_ps(input + 48), sum3); sum4 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 64)))), _mm512_loadu_ps(input + 64), sum4); sum5 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 80)))), _mm512_loadu_ps(input + 80), sum5); sum6 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 96)))), _mm512_loadu_ps(input + 96), sum6); sum7 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(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(weights)))), _mm512_loadu_ps(input), sum0); sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 16)))), _mm512_loadu_ps(input + 16), sum1); sum2 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(weights + 32)))), _mm512_loadu_ps(input + 32), sum2); sum3 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(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(weights)))), _mm512_loadu_ps(input), sum0); sum1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32( _mm_loadu_si128(reinterpret_cast(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(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(0x80)); const __m256i ones = _mm256_set1_epi16(1); 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); sum0 += static_cast(block_dot32_vnni(weights, act, bias, ones)) * act_scales[i / 32]; sum1 += static_cast(block_dot32_vnni(weights + 32, act + 32, bias, ones)) * act_scales[i / 32 + 1]; sum2 += static_cast(block_dot32_vnni(weights + 64, act + 64, bias, ones)) * act_scales[i / 32 + 2]; sum3 += static_cast(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(block_dot32_vnni(weights, act, bias, ones)) * act_scales[i / 32]; result += static_cast(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(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(weights[j]) * static_cast(act[j]); result += static_cast(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. #else // 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(weights[j]) * static_cast(act[j]); result += static_cast(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(weights[j]) * static_cast(act[j]); result += static_cast(block_sum) * scale; } return result; } #endif } // namespace cism