Spaces:
Sleeping
Sleeping
Download native/kernels.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 24.9 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/kernels.cpp
-
curl -L -o kernels.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels.cpp
24.9 kB
| 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() { | |
| const char* flag = std::getenv("CISM_FUSE_SILU"); | |
| 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 = | |
| has_avx2_cpu(); | |
| false; | |
| 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) { | |
| if (use_act_avx2()) return act_exp_avx2(x, n); | |
| for (std::size_t i = 0; i < n; ++i) x[i] = std::exp(x[i]); | |
| } | |
| void act_silu(float* x, std::size_t n) { | |
| if (use_act_avx2()) return act_silu_avx2(x, n); | |
| 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) { | |
| if (use_act_avx2()) return act_silu_mul_avx2(gate, up, out, n); | |
| silu_mul_scalar(gate, up, out, n); | |
| } | |
| void act_sigmoid(float* x, std::size_t n) { | |
| if (use_act_avx2()) return act_sigmoid_avx2(x, n); | |
| 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) { | |
| if (use_act_avx2()) return act_sigmoid_mul_avx2(o, g, n); | |
| 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) { | |
| if (use_act_avx2()) return act_silu_mul_plain_avx2(gate, up, out, n); | |
| 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) { | |
| if (use_act_avx2()) return act_gelu_avx2(x, n); | |
| 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() { | |
| 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; | |
| 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; | |
| return false; | |
| } | |
| // 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 (!has_avx2()) return false; | |
| int regs[4]; | |
| __cpuidex(regs, 7, 1); | |
| return (regs[0] & (1 << 4)) != 0; | |
| if (!has_avx2()) return false; | |
| unsigned a, b, c, d; | |
| __cpuid_count(7, 1, a, b, c, d); | |
| return (a & (1u << 4)) != 0; | |
| return false; | |
| } | |
| // 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 (!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; | |
| 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; | |
| return false; | |
| } | |
| 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() { | |
| 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; | |
| 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; | |
| return false; | |
| } | |
| 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() { | |
| int regs[4]; | |
| __cpuid(regs, 1); | |
| constexpr int want = (1 << 19) | (1 << 9); | |
| return (regs[2] & want) == want; | |
| unsigned a, b, c, d; | |
| __cpuid(1, a, b, c, d); | |
| constexpr unsigned want = (1u << 19) | (1u << 9); | |
| return (c & want) == want; | |
| return false; | |
| } | |
| // 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 (!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; | |
| 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; | |
| return false; | |
| } | |
| 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() { | |
| (void)cpu_has_avx_vnni(); | |
| return cpu_has_avx512_vnni(); | |
| return cpu_has_avx_vnni(); | |
| } | |
| bool has_avx512_vnni_cpu() { return cpu_has_avx512_vnni(); } | |
| bool has_neon_cpu() { | |
| return true; | |
| return false; | |
| } | |
| 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). | |
| if (has_avx2()) { | |
| quantize_row_i8_avx2(input, n, values, scales); | |
| return; | |
| } | |
| 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 (has_neon_cpu()) return dot_neon; | |
| if (has_avx512_vnni_cpu()) return dot_fp32_avx512; | |
| if (has_avx2()) return dot_avx2; | |
| if (has_avx()) return dot_avx_fp32; | |
| if (has_sse41()) return dot_sse41; | |
| (void)has_sse41(); | |
| (void)has_avx(); | |
| 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; | |
| if (selected == dot_fp32_avx512) return true; | |
| return false; | |
| } | |
| const char* kernel_name() { | |
| if (fp32_kernel() == dot_neon) return "neon"; | |
| if (fp32_kernel() == dot_avx_fp32) return "avx"; | |
| if (fp32_kernel() == dot_sse41) return "sse41"; | |
| if (fp32_kernel() == dot_fp32_avx512) return "avx512"; | |
| return fp32_kernel() == dot_scalar ? "scalar" : "avx2"; | |
| } | |
| Int8DotKernel int8_kernel() { | |
| static const Int8DotKernel kernel = []() -> Int8DotKernel { | |
| if (fp32_kernel() == dot_neon) return dot_int8_neon; | |
| // 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; | |
| } | |
| if (use_avx2_kernels()) return dot_int8_avx2; | |
| if (has_sse41()) return dot_int8_sse41; | |
| return dot_int8_scalar; | |
| }(); | |
| return kernel; | |
| } | |
| Fp16DotKernel fp16_kernel() { | |
| static const Fp16DotKernel kernel = []() -> Fp16DotKernel { | |
| if (fp32_kernel() == dot_neon) return dot_fp16_scalar; | |
| if (has_f16c_cpu()) return dot_fp16_avx2; | |
| return dot_fp16_scalar; | |
| }(); | |
| return kernel; | |
| } | |
| Int4DotKernel int4_kernel() { | |
| static const Int4DotKernel kernel = []() -> Int4DotKernel { | |
| if (fp32_kernel() == dot_neon) return dot_int4_neon; | |
| if (use_avx2_kernels()) return dot_int4_avx2; | |
| if (has_sse41()) return dot_int4_sse41; | |
| return dot_int4_scalar5; | |
| }(); | |
| return kernel; | |
| } | |
| Int8Q8DotKernel int8_q8_kernel() { | |
| static const Int8Q8DotKernel kernel = []() -> Int8Q8DotKernel { | |
| if (use_avx2_kernels()) return dot_int8_q8_avx2; | |
| return nullptr; | |
| }(); | |
| return kernel; | |
| } | |
| Int4Q8DotKernel int4_q8_kernel() { | |
| static const Int4Q8DotKernel kernel = []() -> Int4Q8DotKernel { | |
| if (use_avx2_kernels()) return dot_int4_q8_avx2; | |
| return nullptr; | |
| }(); | |
| return kernel; | |
| } | |
| Fp4Q8DotKernel fp4_q8_kernel() { | |
| static const Fp4Q8DotKernel kernel = []() -> Fp4Q8DotKernel { | |
| if (use_avx2_kernels()) return dot_fp4_q8_avx2; | |
| return nullptr; | |
| }(); | |
| return kernel; | |
| } | |
| Fp4DotKernel fp4_kernel() { | |
| static const Fp4DotKernel kernel = []() -> Fp4DotKernel { | |
| if (fp32_kernel() == dot_neon) return dot_fp4_neon; | |
| if (use_avx2_kernels()) return dot_fp4_avx2; | |
| if (has_sse41()) return dot_fp4_sse41; | |
| return dot_fp4_scalar5; | |
| }(); | |
| return kernel; | |
| } | |
| } // namespace cism | |