#include "kernels.hpp" #include #include #include #include #include #if (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && defined(_MSC_VER) #include #elif (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && (defined(__GNUC__) || defined(__clang__)) #include #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((x >> 16) & 0x8000u); const std::uint32_t ax = x & 0x7FFFFFFFu; const int e32 = static_cast(ax >> 23); if (e32 == 255) return static_cast(sign | 0x7BFFu); // inf/NaN: saturate (weights are finite) const int e16 = e32 - 112; if (e16 >= 31) return static_cast(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(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(sign | 0x7BFFu); return static_cast( sign | (static_cast(e16 + 1) << 10)); } } return static_cast(sign | (static_cast(e16) << 10) | half_mant); } float fp16_to_fp32(std::uint16_t bits) { const std::uint32_t sign = (static_cast(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(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(mantissa) * (1.0f / 8.0f) * 0.015625f; } else { value = (1.0f + static_cast(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 table; for (int bits = 0; bits < 256; ++bits) table[bits] = fp4_decode_scale(static_cast(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::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(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::nearbyint(value * 512.0f)); return static_cast((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(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(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(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(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(input[i]) * norm); q = std::clamp(q, -127, 127); values[i] = static_cast(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