test1111111 / native /kernels.cpp
spitfire4794's picture
Serve SurjoLabs/Surjo-50m-SFT-Only (int8, 2T) on 7860
92dcc4e verified
Raw History Blame Contribute Delete
24.9 kB
#include "kernels.hpp"
#include <algorithm>
#include <array>
#include <cmath>
#include <cstdlib>
#include <cstring>
#if (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && defined(_MSC_VER)
#include <intrin.h>
#elif (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && (defined(__GNUC__) || defined(__clang__))
#include <cpuid.h>
#endif
namespace cism {
// IEEE-754 binary16 conversion (round-to-nearest-even), portable: the AVX2
// TU uses vcvtph where available, but packing and scalar fallback need this
// everywhere (including ARM/SSE builds without F16C).
std::uint16_t fp32_to_fp16(float value) {
std::uint32_t x;
std::memcpy(&x, &value, 4);
const std::uint16_t sign = static_cast<std::uint16_t>((x >> 16) & 0x8000u);
const std::uint32_t ax = x & 0x7FFFFFFFu;
const int e32 = static_cast<int>(ax >> 23);
if (e32 == 255) return static_cast<std::uint16_t>(sign | 0x7BFFu); // inf/NaN: saturate (weights are finite)
const int e16 = e32 - 112;
if (e16 >= 31) return static_cast<std::uint16_t>(sign | 0x7BFFu); // overflow: max finite
const std::uint32_t mant = ax & 0x7FFFFFu;
if (e16 <= 0) {
if (e16 < -10) return sign; // underflow to zero
// Subnormal: hidden 1 + RNE shift. shift = 126-E in [14,24].
const int shift = 126 - e32;
const std::uint32_t m = mant | 0x800000u;
const std::uint32_t base = m >> shift;
const std::uint32_t rbit = (m >> (shift - 1)) & 1u;
const std::uint32_t sticky = m & ((shift == 1) ? 0u : ((1u << (shift - 1)) - 1u));
const std::uint32_t half = base + ((rbit && (sticky || (base & 1u))) ? 1u : 0u);
return static_cast<std::uint16_t>(sign | half); // half <= 0x400: valid code
}
const std::uint32_t dropped = mant & 0x1FFFu;
std::uint32_t half_mant = mant >> 13;
if (dropped > 0x1000u || (dropped == 0x1000u && (half_mant & 1u))) {
if (++half_mant == 0x400u) {
// Mantissa carry: exact power of two (e16+1), except at e16==30
// where it overflows to max finite (value was below 65504).
if (e16 == 30) return static_cast<std::uint16_t>(sign | 0x7BFFu);
return static_cast<std::uint16_t>(
sign | (static_cast<std::uint32_t>(e16 + 1) << 10));
}
}
return static_cast<std::uint16_t>(sign | (static_cast<std::uint32_t>(e16) << 10) | half_mant);
}
float fp16_to_fp32(std::uint16_t bits) {
const std::uint32_t sign = (static_cast<std::uint32_t>(bits) & 0x8000u) << 16;
const std::uint32_t exp = (bits >> 10) & 0x1Fu;
const std::uint32_t mant = bits & 0x3FFu;
std::uint32_t f;
if (exp == 0) {
if (mant == 0) {
f = sign;
} else {
int e = -14;
std::uint32_t m = mant;
while (!(m & 0x400u)) {
m <<= 1;
--e;
}
m &= 0x3FFu;
f = sign | (static_cast<std::uint32_t>(e + 127) << 23) | (m << 13);
}
} else if (exp == 31) {
f = sign | 0x7F800000u | (mant << 13);
} else {
f = sign | ((exp + 112) << 23) | (mant << 13);
}
float out;
std::memcpy(&out, &f, 4);
return out;
}
const std::int8_t* fp4_element_lut() {
// Raw E2M1 nibble (sign in bit 3) -> E2M1 value times two, exact in int8.
static const std::int8_t lut[16] = {0, 1, 2, 3, 4, 6, 8, 12,
0, -1, -2, -3, -4, -6, -8, -12};
return lut;
}
float fp4_decode_scale(std::uint8_t bits) {
// E4M3 (fn variant): subnormal mantissa/8 * 2^-6, normal (1+m/8) * 2^(e-7).
const int exponent = (bits >> 3) & 15;
const int mantissa = bits & 7;
float value;
if (exponent == 0) {
value = static_cast<float>(mantissa) * (1.0f / 8.0f) * 0.015625f;
} else {
value = (1.0f + static_cast<float>(mantissa) / 8.0f) *
std::ldexp(1.0f, exponent - 7);
}
return (bits & 128) ? -value : value;
}
const float* fp4_scale_lut() {
// Decoded E4M3 scale, pre-halved: a half-value element times this entry
// is exactly E2M1 * E4M3, so the vector FMA path needs no extra multiply.
static const auto lut = []() {
static std::array<float, 256> table;
for (int bits = 0; bits < 256; ++bits) table[bits] = fp4_decode_scale(static_cast<std::uint8_t>(bits)) * 0.5f;
return table;
}();
return lut.data();
}
std::uint8_t fp4_encode_scale(float value) {
// Round-to-nearest-even into E4M3, saturating at the max normal (448),
// matching torch.float8_e4m3fn for positive finite inputs.
if (!(value > 0)) return 0;
if (value >= 448.0f) return 0x7E;
if (value < 0.015625f) // smallest normal is 2^-6; below that, subnormals
return static_cast<std::uint8_t>(std::nearbyint(value * 512.0f));
int exponent = 0;
std::frexp(value, &exponent);
const int e = exponent - 1; // floor(log2(value))
const float mantissa = std::ldexp(value, -e); // [1, 2)
int m = static_cast<int>(std::nearbyint((mantissa - 1.0f) * 8.0f));
int eb = e + 7;
if (m >= 8) { m = 0; ++eb; }
if (eb > 15) return 0x7E;
if (eb < 1) // rounding landed below the normal range
return static_cast<std::uint8_t>(std::nearbyint(value * 512.0f));
return static_cast<std::uint8_t>((eb << 3) | m);
}
float dot_fp4_scalar(const std::uint8_t* weights, const float* input, const std::uint8_t* scales, std::size_t n) {
const auto* elements = fp4_element_lut();
const float* scale_lut = fp4_scale_lut();
float sum = 0;
for (std::size_t start = 0; start < n; start += 16) {
const float scale = scale_lut[scales[start / 16]];
float block_sum = 0;
for (auto j = start; j < std::min(n, start + 16); ++j) {
const int nibble = (weights[j / 2] >> (4 * (j % 2))) & 15;
block_sum += static_cast<float>(elements[nibble]) * input[j];
}
sum += block_sum * scale;
}
return sum;
}
float dot_scalar(const float* a, const float* b, std::size_t n) {
float result = 0;
for (std::size_t i = 0; i < n; ++i) result += a[i] * b[i];
return result;
}
float dot_int8_scalar(const std::int8_t* weights, const float* input, std::size_t n) {
float sum = 0;
for (std::size_t j = 0; j < n; ++j) sum += static_cast<float>(weights[j]) * input[j];
return sum;
}
float dot_fp16_scalar(const std::uint16_t* weights, const float* input, std::size_t n) {
float sum = 0;
for (std::size_t j = 0; j < n; ++j) sum += fp16_to_fp32(weights[j]) * input[j];
return sum;
}
float dot_int4_scalar(const std::uint8_t* weights, const float* input, const float* scales, std::size_t n) {
// Split-half layout: byte j holds w[j] (low) and w[j+16] (high);
// codes 0..15 map to (code-8), matching the packer exactly.
float sum = 0;
for (std::size_t start = 0; start < n; start += 32) {
float block_sum = 0;
const std::size_t m = std::min(n - start, static_cast<std::size_t>(32));
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) * input[start + t];
}
weights += 16;
sum += block_sum * scales[start / 32];
}
return sum;
}
// 5-arg adapters so the permuted kernel type keeps a working scalar
// fallback on non-AVX2 builds (the permuted buffer is simply unused).
float dot_int4_scalar5(const std::uint8_t* weights, const float* /*act_perm*/,
const float* input, const float* scales, std::size_t n) {
return dot_int4_scalar(weights, input, scales, n);
}
float dot_fp4_scalar5(const std::uint8_t* weights, const float* /*act_perm*/,
const float* input, const std::uint8_t* scales, std::size_t n) {
return dot_fp4_scalar(weights, input, scales, n);
}
// Canonical permuted-activation layout (scalar reference, bitwise stable):
// per full 32-chunk [c,c+32): OUT[c+k]=IN[c+2k], OUT[c+16+k]=IN[c+2k+1].
// Tail elements beyond (n/32)*32 are left undefined, matching the runtime
// permute_act32 contract. 64-wide inner blocking halves loop overhead with
// identical element order (two 32-blocks per iteration, in order).
void permute_act32_blocked(const float* input, std::size_t n, float* out) {
std::size_t c = 0;
for (; c + 64 <= n; c += 64) {
for (std::size_t k = 0; k < 16; ++k) {
out[c + k] = input[c + 2 * k];
out[c + 16 + k] = input[c + 2 * k + 1];
}
for (std::size_t k = 0; k < 16; ++k) {
out[c + 32 + k] = input[c + 32 + 2 * k];
out[c + 48 + k] = input[c + 32 + 2 * k + 1];
}
}
for (; c + 32 <= n; c += 32) {
for (std::size_t k = 0; k < 16; ++k) {
out[c + k] = input[c + 2 * k];
out[c + 16 + k] = input[c + 2 * k + 1];
}
}
}
// Stable SiLU*up epilogue, bit-identical to Session::forward_tokens.
void silu_mul_scalar(const float* gate, const float* up, float* out, std::size_t n) {
for (std::size_t i = 0; 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];
}
}
// Test-gated only: false unless CISM_FUSE_SILU=1. The runtime decode path
// never consults this (fusion stays off after the measured ~7% slowdown);
// it exists so experiments stay explicitly gated and PPL-checked.
bool silu_fusion_enabled() {
#ifdef _MSC_VER
#pragma warning(push)
#pragma warning(disable : 4996)
#endif
const char* flag = std::getenv("CISM_FUSE_SILU");
#ifdef _MSC_VER
#pragma warning(pop)
#endif
return flag != nullptr && flag[0] == '1' && flag[1] == '\0';
}
// Vector-exp activation blocks. Scalar loops use libm with the exact runtime
// formulas (bitwise reference); the AVX2 TU uses a ≤1-ULP polynomial exp.
// Dispatch resolved once per process (the CPU does not change under us).
namespace {
bool use_act_avx2() {
static const bool cached =
#ifdef CISM_HAVE_AVX2
has_avx2_cpu();
#else
false;
#endif
return cached;
}
inline float scalar_sigmoid(float x) {
return x >= 0 ? 1.0f / (1.0f + std::exp(-x)) :
std::exp(x) / (1.0f + std::exp(x));
}
} // namespace
void act_exp(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_exp_avx2(x, n);
#endif
for (std::size_t i = 0; i < n; ++i) x[i] = std::exp(x[i]);
}
void act_silu(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_silu_avx2(x, n);
#endif
for (std::size_t i = 0; i < n; ++i) x[i] = x[i] * scalar_sigmoid(x[i]);
}
void act_silu_mul(const float* gate, const float* up, float* out, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_silu_mul_avx2(gate, up, out, n);
#endif
silu_mul_scalar(gate, up, out, n);
}
void act_sigmoid(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_sigmoid_avx2(x, n);
#endif
for (std::size_t i = 0; i < n; ++i) x[i] = scalar_sigmoid(x[i]);
}
void act_sigmoid_mul(float* o, const float* g, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_sigmoid_mul_avx2(o, g, n);
#endif
for (std::size_t i = 0; i < n; ++i) o[i] *= scalar_sigmoid(g[i]);
}
void act_silu_mul_plain(const float* gate, const float* up, float* out, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_silu_mul_plain_avx2(gate, up, out, n);
#endif
for (std::size_t i = 0; i < n; ++i) {
const float value = gate[i];
out[i] = (value * scalar_sigmoid(value)) * up[i];
}
}
void act_gelu(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
if (use_act_avx2()) return act_gelu_avx2(x, n);
#endif
for (std::size_t i = 0; i < n; ++i) {
const float value = x[i];
x[i] = 0.5f * value * (1.0f + std::erf(value * 0.7071067811865475f));
}
}
static bool has_avx2() {
#if defined(CISM_HAVE_AVX2) && defined(_MSC_VER)
int regs[4];
__cpuid(regs, 0);
if (regs[0] < 7) return false;
__cpuidex(regs, 1, 0);
// Kernels use explicit FMA intrinsics, so require AVX, OSXSAVE, and FMA.
constexpr int osxsave_avx_fma = (1 << 27) | (1 << 28) | (1 << 12);
if ((regs[2] & osxsave_avx_fma) != osxsave_avx_fma) return false;
if ((_xgetbv(0) & 6) != 6) return false;
__cpuidex(regs, 7, 0);
return (regs[1] & (1 << 5)) != 0;
#elif defined(CISM_HAVE_AVX2) && (defined(__GNUC__) || defined(__clang__))
unsigned a, b, c, d;
if (__get_cpuid_max(0, nullptr) < 7) return false;
__cpuid_count(1, 0, a, b, c, d);
constexpr unsigned osxsave_avx_fma = (1u << 27) | (1u << 28) | (1u << 12);
if ((c & osxsave_avx_fma) != osxsave_avx_fma) return false;
unsigned xcr0_low, xcr0_high;
__asm__ volatile("xgetbv" : "=a"(xcr0_low), "=d"(xcr0_high) : "c"(0));
if ((xcr0_low & 6) != 6) return false;
__cpuid_count(7, 0, a, b, c, d);
return (b & (1u << 5)) != 0;
#else
return false;
#endif
}
// AVX-VNNI: CPUID leaf 7 subleaf 1 EAX[4] (Zen4/Zen5, Alder Lake+). VEX.256
// encoding needs only AVX OS state, so AVX2+FMA presence implies the state;
// still require has_avx2() so ancient/weird CPUs never take this branch.
static bool cpu_has_avx_vnni() {
#if (defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI)) && defined(_MSC_VER)
if (!has_avx2()) return false;
int regs[4];
__cpuidex(regs, 7, 1);
return (regs[0] & (1 << 4)) != 0;
#elif (defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI)) && (defined(__GNUC__) || defined(__clang__))
if (!has_avx2()) return false;
unsigned a, b, c, d;
__cpuid_count(7, 1, a, b, c, d);
return (a & (1u << 4)) != 0;
#else
return false;
#endif
}
// AVX512-VNNI (EVEX encoding): leaf 7 subleaf 0 EBX[11] plus the AVX512
// foundation (F/BW/VL/DQ) and ZMM OS state (XCR0 opmask/ZMM_Hi256/Hi16_ZMM).
static bool cpu_has_avx512_vnni() {
#if defined(CISM_HAVE_AVX512_VNNI) && defined(_MSC_VER)
if (!has_avx2()) return false;
int regs[4];
__cpuidex(regs, 7, 0);
constexpr int want = (1 << 16) | (1 << 17) | (1 << 30) | (1 << 31) | (1 << 11);
if ((regs[1] & want) != want) return false;
return (_xgetbv(0) & 0xE0) == 0xE0;
#elif defined(CISM_HAVE_AVX512_VNNI) && (defined(__GNUC__) || defined(__clang__))
if (!has_avx2()) return false;
unsigned a, b, c, d;
__cpuid_count(7, 0, a, b, c, d);
constexpr unsigned want = (1u << 16) | (1u << 17) | (1u << 30) | (1u << 31) | (1u << 11);
if ((b & want) != want) return false;
unsigned lo, hi;
__asm__ volatile("xgetbv" : "=a"(lo), "=d"(hi) : "c"(0));
return (lo & 0xE0u) == 0xE0u;
#else
return false;
#endif
}
bool has_avx2_cpu() { return has_avx2(); }
// F16C: leaf 1 ECX[29] plus the same AVX OS state as above (vcvtph needs
// YMM state). Gated on CISM_HAVE_F16C (compiler flag present).
static bool has_f16c() {
#if defined(CISM_HAVE_F16C) && defined(_MSC_VER)
int regs[4];
__cpuid(regs, 1);
constexpr int osxsave_avx_fma_f16c = (1 << 27) | (1 << 28) | (1 << 12) | (1 << 29);
if ((regs[2] & osxsave_avx_fma_f16c) != osxsave_avx_fma_f16c) return false;
return (_xgetbv(0) & 6) == 6;
#elif defined(CISM_HAVE_F16C) && (defined(__GNUC__) || defined(__clang__))
unsigned a, b, c, d;
__cpuid(1, a, b, c, d);
constexpr unsigned osxsave_avx_fma_f16c = (1u << 27) | (1u << 28) | (1u << 12) | (1u << 29);
if ((c & osxsave_avx_fma_f16c) != osxsave_avx_fma_f16c) return false;
unsigned xcr0_low, xcr0_high;
__asm__ volatile("xgetbv" : "=a"(xcr0_low), "=d"(xcr0_high) : "c"(0));
return (xcr0_low & 6) == 6;
#else
return false;
#endif
}
bool has_f16c_cpu() { return has_f16c(); }
// SSE4.1: leaf 1 ECX[19]. SSSE3 (pshufb, needed by the int4/fp4 LUT path)
// is ECX[9]; require both so the nibble kernels never run without pshufb.
static bool has_sse41() {
#if (defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && defined(_MSC_VER)
int regs[4];
__cpuid(regs, 1);
constexpr int want = (1 << 19) | (1 << 9);
return (regs[2] & want) == want;
#elif (defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && (defined(__GNUC__) || defined(__clang__))
unsigned a, b, c, d;
__cpuid(1, a, b, c, d);
constexpr unsigned want = (1u << 19) | (1u << 9);
return (c & want) == want;
#else
return false;
#endif
}
// AVX (Sandy/Ivy/Bulldozer era): leaf 1 ECX[27,28] (OSXSAVE+AVX) plus XCR0
// SSE+AVX state. No FMA requirement (Sandy lacks it, Bulldozer has FMA4
// not FMA3); the AVX TU uses mul+add only.
static bool has_avx() {
#if defined(CISM_HAVE_AVX) && defined(_MSC_VER)
if (!has_sse41()) return false;
int regs[4];
__cpuid(regs, 1);
constexpr int want = (1 << 27) | (1 << 28);
if ((regs[2] & want) != want) return false;
return (_xgetbv(0) & 6) == 6;
#elif defined(CISM_HAVE_AVX) && (defined(__GNUC__) || defined(__clang__))
if (!has_sse41()) return false;
unsigned a, b, c, d;
__cpuid(1, a, b, c, d);
constexpr unsigned want = (1u << 27) | (1u << 28);
if ((c & want) != want) return false;
unsigned lo, hi;
__asm__ volatile("xgetbv" : "=a"(lo), "=d"(hi) : "c"(0));
return (lo & 6u) == 6u;
#else
return false;
#endif
}
bool has_sse41_cpu() { return has_sse41(); }
bool has_avx_cpu() { return has_avx(); }
// The VNNI TU is compiled EVEX (MSVC /arch:AVX512) or VEX (GCC -mavxvnni);
// each encoding runs only on CPUs supporting exactly it.
bool has_avx_vnni_cpu() {
#if defined(CISM_HAVE_AVX512_VNNI) && !defined(CISM_HAVE_AVX_VNNI)
(void)cpu_has_avx_vnni();
return cpu_has_avx512_vnni();
#else
return cpu_has_avx_vnni();
#endif
}
bool has_avx512_vnni_cpu() { return cpu_has_avx512_vnni(); }
bool has_neon_cpu() {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
return true;
#else
return false;
#endif
}
const char* kernel_variant() {
if (has_neon_cpu()) return "neon";
if (has_avx512_vnni_cpu()) return "avx512-vnni";
if (has_avx_vnni_cpu()) return "avx-vnni";
if (has_avx2_cpu()) return "avx2";
if (has_avx_cpu()) return "avx";
if (has_sse41_cpu()) return "sse41";
return "scalar";
}
void quantize_row_i8(const float* input, std::size_t n, std::int8_t* values, float* scales) {
// int8 sibling of runtime.cpp quantize_row: 127/absmax per 32-element
// block, one fp32 dequant scale per block. Scalar reference for the VNNI
// path; runs on every CPU (also the portable-kernels unit test target).
#ifdef CISM_HAVE_AVX2
if (has_avx2()) {
quantize_row_i8_avx2(input, n, values, scales);
return;
}
#endif
for (std::size_t start = 0; start < n; start += 32) {
const std::size_t end = std::min(start + 32, n);
float absmax = 0.0f;
for (std::size_t i = start; i < end; ++i) absmax = std::max(absmax, std::abs(input[i]));
if (absmax == 0.0f) {
for (std::size_t i = start; i < end; ++i) values[i] = 0;
scales[start / 32] = 0.0f;
continue;
}
const float norm = 127.0f / absmax;
for (std::size_t i = start; i < end; ++i) {
long q = std::lrint(static_cast<double>(input[i]) * norm);
q = std::clamp<long>(q, -127, 127);
values[i] = static_cast<std::int8_t>(q);
}
scales[start / 32] = absmax / 127.0f;
}
}
Int8Q8VnniKernel int8_q8_vnni_kernel() {
static const Int8Q8VnniKernel kernel = []() -> Int8Q8VnniKernel {
if (has_avx_vnni_cpu()) return dot_int8_q8_vnni;
return nullptr;
}();
return kernel;
}
DotKernel fp32_kernel() {
static const DotKernel kernel = []() -> DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (has_neon_cpu()) return dot_neon;
#endif
#if defined(CISM_HAVE_AVX512_VNNI)
if (has_avx512_vnni_cpu()) return dot_fp32_avx512;
#endif
#ifdef CISM_HAVE_AVX2
if (has_avx2()) return dot_avx2;
#endif
#ifdef CISM_HAVE_AVX
if (has_avx()) return dot_avx_fp32;
#endif
#ifdef CISM_HAVE_SSE41
if (has_sse41()) return dot_sse41;
#else
(void)has_sse41();
(void)has_avx();
#endif
return dot_scalar;
}();
return kernel;
}
// AVX2-or-better family: everything dispatched on dot_avx2 stays valid when
// the zmm kernel leads (same ISA superset, same per-row results contract).
static bool use_avx2_kernels() {
const auto selected = fp32_kernel();
if (selected == dot_avx2) return true;
#ifdef CISM_HAVE_AVX512_VNNI
if (selected == dot_fp32_avx512) return true;
#endif
return false;
}
const char* kernel_name() {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (fp32_kernel() == dot_neon) return "neon";
#endif
#ifdef CISM_HAVE_AVX
if (fp32_kernel() == dot_avx_fp32) return "avx";
#endif
#ifdef CISM_HAVE_SSE41
if (fp32_kernel() == dot_sse41) return "sse41";
#endif
#ifdef CISM_HAVE_AVX512_VNNI
if (fp32_kernel() == dot_fp32_avx512) return "avx512";
#endif
return fp32_kernel() == dot_scalar ? "scalar" : "avx2";
}
Int8DotKernel int8_kernel() {
static const Int8DotKernel kernel = []() -> Int8DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (fp32_kernel() == dot_neon) return dot_int8_neon;
#endif
#if defined(CISM_HAVE_AVX512_VNNI)
// zmm int8 measured SLOWER than AVX2 on Xeon 8375C 2-vCPU (87 vs 99
// tok/s: 512-bit throttling eats the width gain), so it is opt-in
// only (CISM_ENABLE_AVX512_INT8=1) for experimentation/other tin.
if (has_avx512_vnni_cpu()) {
const char* flag = std::getenv("CISM_ENABLE_AVX512_INT8");
if (flag != nullptr && flag[0] == '1' && flag[1] == '\0')
return dot_int8_avx512;
}
#endif
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_int8_avx2;
#endif
#ifdef CISM_HAVE_SSE41
if (has_sse41()) return dot_int8_sse41;
#endif
return dot_int8_scalar;
}();
return kernel;
}
Fp16DotKernel fp16_kernel() {
static const Fp16DotKernel kernel = []() -> Fp16DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (fp32_kernel() == dot_neon) return dot_fp16_scalar;
#endif
#ifdef CISM_HAVE_F16C
if (has_f16c_cpu()) return dot_fp16_avx2;
#endif
return dot_fp16_scalar;
}();
return kernel;
}
Int4DotKernel int4_kernel() {
static const Int4DotKernel kernel = []() -> Int4DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (fp32_kernel() == dot_neon) return dot_int4_neon;
#endif
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_int4_avx2;
#endif
#ifdef CISM_HAVE_SSE41
if (has_sse41()) return dot_int4_sse41;
#endif
return dot_int4_scalar5;
}();
return kernel;
}
Int8Q8DotKernel int8_q8_kernel() {
static const Int8Q8DotKernel kernel = []() -> Int8Q8DotKernel {
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_int8_q8_avx2;
#endif
return nullptr;
}();
return kernel;
}
Int4Q8DotKernel int4_q8_kernel() {
static const Int4Q8DotKernel kernel = []() -> Int4Q8DotKernel {
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_int4_q8_avx2;
#endif
return nullptr;
}();
return kernel;
}
Fp4Q8DotKernel fp4_q8_kernel() {
static const Fp4Q8DotKernel kernel = []() -> Fp4Q8DotKernel {
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_fp4_q8_avx2;
#endif
return nullptr;
}();
return kernel;
}
Fp4DotKernel fp4_kernel() {
static const Fp4DotKernel kernel = []() -> Fp4DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
if (fp32_kernel() == dot_neon) return dot_fp4_neon;
#endif
#ifdef CISM_HAVE_AVX2
if (use_avx2_kernels()) return dot_fp4_avx2;
#endif
#ifdef CISM_HAVE_SSE41
if (has_sse41()) return dot_fp4_sse41;
#endif
return dot_fp4_scalar5;
}();
return kernel;
}
} // namespace cism