test1111111 / native /kernels_vnni.cpp
spitfire4794's picture
Space: AVX512 kernels (zmm fp32/int8, parallel VNNI-Q8) + 2T cap + NEON int4 fix
40d0cd2
Raw History Blame Contribute Delete
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).
#include "kernels.hpp"
#include <cstddef>
#include <cstdint>
#if defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI)
#include <immintrin.h>
#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<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.
#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<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;
}
#endif
} // namespace cism