Download src/kernels/cpu/iq_avx2.cpp from WineryLabs/Winery-Strata: direct link, hf CLI and curl.
- Browser
- Download file 26 kB
-
https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/kernels/cpu/iq_avx2.cpp
- Command line
-
hf download hf://WineryLabs/Winery-Strata/src/kernels/cpu/iq_avx2.cpp
-
curl -L -o iq_avx2.cpp https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/kernels/cpu/iq_avx2.cpp
26 kB
| // src/kernels/cpu/iq_avx2.cpp - the i-quant expert rows in 256-bit lanes, several tokens at once. | |
| // | |
| // The AVX-512 kernels' (iq_avx512.cpp) multi-token scheme on the CPUs without AVX-512 (AMD Zen 2/3, Intel | |
| // Core 12th-14th gen and Core Ultra): each 32-value chunk is decoded once - grid lookups, one sign vector, | |
| // one scale vector - and every token of the verify window applies them with a load, a `sign`, a `maddubs`, | |
| // a `madd` and an `add`. ggml-cpu's own AVX2 dot products for these formats are single-token: every token | |
| // re-does the codebook lookups and the sign expansion. The sign vector is ggml's bit_selector pattern | |
| // (ggml-cpu/arch/x86/quants.c): pshufb-broadcast of each sign byte, AND with the bit selector, CMPEQ, OR | |
| // with one, `vpsignb`. The arithmetic is ggml's (ggml-cpu/quants.c, the `_generic` references) - only the | |
| // order of the float additions differs. | |
| // | |
| // Formats: IQ2_XXS (16), IQ2_XS (17), IQ3_XXS (18), IQ3_S (21), IQ2_S (22), IQ4_XS (23). IQ1_M stays on ggml-cpu. | |
| namespace strata::kernels::cpu { | |
| namespace { | |
| inline float h2f(uint16_t h) { return _mm_cvtss_f32(_mm_cvtph_ps(_mm_cvtsi32_si128((int) h))); } | |
| inline uint32_t u32(const uint8_t* p) { uint32_t v; std::memcpy(&v, p, 4); return v; } | |
| inline uint16_t u16(const uint8_t* p) { uint16_t v; std::memcpy(&v, p, 2); return v; } | |
| inline uint64_t u64(const uint8_t* p) { uint64_t v; std::memcpy(&v, p, 8); return v; } | |
| // 32 sign bits -> 32 bytes of -1 (bit set) / +1: ggml's bit_selector pattern, shared by all tokens. | |
| inline __m256i sgn_vec(uint32_t m) { | |
| const __m128i bm = _mm_set1_epi32((int) m); | |
| const __m128i lo = _mm_shuffle_epi8(bm, _mm_setr_epi8(0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1)); | |
| const __m128i hi = _mm_shuffle_epi8(bm, _mm_setr_epi8(2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3)); | |
| const __m256i bits = _mm256_inserti128_si256(_mm256_castsi128_si256(lo), hi, 1); | |
| const __m256i sel = _mm256_setr_epi8(1, 2, 4, 8, 16, 32, 64, (char) 0x80, | |
| 1, 2, 4, 8, 16, 32, 64, (char) 0x80, | |
| 1, 2, 4, 8, 16, 32, 64, (char) 0x80, | |
| 1, 2, 4, 8, 16, 32, 64, (char) 0x80); | |
| const __m256i nz = _mm256_cmpeq_epi8(_mm256_and_si256(bits, sel), sel); | |
| return _mm256_or_si256(nz, _mm256_set1_epi8(1)); | |
| } | |
| // The two 16-value scales of one 32-value half (IQ2_XS, IQ2_S): int16 lanes 0-7 = a (values 0-15), | |
| // 8-15 = b (values 16-31). NOT a broadcast dword [a,b] - that would alternate the scales lane by lane. | |
| inline __m256i sc16(int a, int b) { | |
| return _mm256_inserti128_si256(_mm256_castsi128_si256(_mm_set1_epi16((short) a)), | |
| _mm_set1_epi16((short) b), 1); | |
| } | |
| // One scale for the whole 32-value half (IQ2_XXS, IQ3_XXS, IQ3_S). | |
| inline __m256i sc32(int s) { return _mm256_set1_epi16(s); } | |
| // keven_signs_q2xs (ggml keeps it static in arch/x86/quants.c): one u64 per 7-bit sign index, byte k | |
| // = 0xFF when bit k of ksigns_iq2xs[i] is set, 0x01 otherwise. With this a whole 32-value sign vector | |
| // is four scalar loads and a set_epi64x - no ksigns byte packing chain and no bit_selector expansion. | |
| // Built once at load time from the shared ggml-common table so it can never drift from it. | |
| struct EvenSigns { | |
| uint64_t v[128]; | |
| EvenSigns() { | |
| for (int i = 0; i < 128; ++i) { | |
| uint64_t r = 0; | |
| for (int k = 0; k < 8; ++k) | |
| r |= (uint64_t) (((ksigns_iq2xs[i] >> k) & 1) ? 0xFF : 0x01) << (8 * k); | |
| v[i] = r; | |
| } | |
| } | |
| }; | |
| static const EvenSigns even_signs; | |
| inline float hsum8(__m256 v) { | |
| const __m128 lo = _mm256_castps256_ps128(v), hi = _mm256_extractf128_ps(v, 1); | |
| __m128 s = _mm_add_ps(lo, hi); | |
| s = _mm_hadd_ps(s, s); | |
| s = _mm_hadd_ps(s, s); | |
| return _mm_cvtss_f32(s); | |
| } | |
| // E-2 (iq_avx512.cpp) on the AVX-2 path: the pool streams the expert rows from DRAM at ~25 GB/s (4 KB pages | |
| // when large pages are refused), so ask for the bytes a few blocks before the decode needs them - 2048 B is | |
| // two gate/up rows ahead, a row is ~1 KB. Same switch as the AVX-512 kernels: STRATA_IQ_PREFETCH is the | |
| // distance in bytes, 0 = off, default 2048. Measured on a Zen 3 5700X3D (no AVX-512) on IQ3_S decode: -4% on | |
| // the gate/up phase, +1.0 GB/s over the rows, -1.3% ms/round end to end. The non-temporal hint measured | |
| // worse than T0 at the same distance, so this keeps T0. | |
| const int prefetch_ahead = [] { | |
| const char* v = std::getenv("STRATA_IQ_PREFETCH"); | |
| return v ? std::atoi(v) : 2048; | |
| }(); | |
| inline void rows_ahead(const uint8_t* p) { | |
| if (prefetch_ahead <= 0) return; | |
| _mm_prefetch((const char*) p, _MM_HINT_T0); | |
| _mm_prefetch((const char*) p + 64, _MM_HINT_T0); | |
| } | |
| // ---- per format: one 32-value half (values 64*j + 32*half .. +31) -> grid magnitudes, sign vector, scales | |
| template <int TY> struct Fmt32; | |
| template <> struct Fmt32<16> { // IQ2_XXS: d, qs[32] u16 | |
| static constexpr int bytes = 66; | |
| static constexpr float K = 0.125f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| const uint8_t* q = b + 2 + 16 * j + 8 * half; | |
| const uint32_t w0 = u32(q), w1 = u32(q + 4); | |
| g = _mm256_set_epi64x((long long) iq2xxs_grid[w0 >> 24], (long long) iq2xxs_grid[(w0 >> 16) & 255], | |
| (long long) iq2xxs_grid[(w0 >> 8) & 255], (long long) iq2xxs_grid[w0 & 255]); | |
| sgn = _mm256_set_epi64x((long long) even_signs.v[(w1 >> 21) & 127], (long long) even_signs.v[(w1 >> 14) & 127], | |
| (long long) even_signs.v[(w1 >> 7) & 127], (long long) even_signs.v[w1 & 127]); | |
| sc = sc32(2 * (int) (w1 >> 28) + 1); | |
| } | |
| }; | |
| template <> struct Fmt32<17> { // IQ2_XS: d, qs[32] u16 (9-bit grid index + 7-bit sign index), scales[8] | |
| static constexpr int bytes = 74; | |
| static constexpr float K = 0.125f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| uint16_t v[8]; | |
| std::memcpy(v, b + 2 + 16 * j, 16); | |
| const int o = 4 * half; | |
| g = _mm256_set_epi64x((long long) iq2xs_grid[v[o + 3] & 511], (long long) iq2xs_grid[v[o + 2] & 511], | |
| (long long) iq2xs_grid[v[o + 1] & 511], (long long) iq2xs_grid[v[o] & 511]); | |
| uint32_t s = 0; | |
| for (int l = 0; l < 4; ++l) s |= (uint32_t) ksigns_iq2xs[v[o + l] >> 9] << (8 * l); | |
| sgn = sgn_vec(s); | |
| const uint8_t sb = b[66 + 2 * j + half]; | |
| sc = sc16(2 * (sb & 15) + 1, 2 * (sb >> 4) + 1); | |
| } | |
| }; | |
| template <> struct Fmt32<22> { // IQ2_S: d, qs[64] (32 grid bytes, 32 sign bytes), qh[8], scales[8] | |
| static constexpr int bytes = 82; | |
| static constexpr float K = 0.125f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| const uint8_t* qs = b + 2 + 8 * j; | |
| const uint8_t h = b[66 + 2 * j + half]; | |
| const int o = 4 * half; | |
| g = _mm256_set_epi64x((long long) iq2s_grid[qs[o + 3] | ((h << 2) & 0x300)], | |
| (long long) iq2s_grid[qs[o + 2] | ((h << 4) & 0x300)], | |
| (long long) iq2s_grid[qs[o + 1] | ((h << 6) & 0x300)], | |
| (long long) iq2s_grid[qs[o] | ((h << 8) & 0x300)]); | |
| const uint64_t m = u64(b + 2 + 32 + 8 * j); | |
| sgn = sgn_vec(half ? (uint32_t) (m >> 32) : (uint32_t) m); | |
| const uint8_t sb = b[74 + 2 * j + half]; | |
| sc = sc16(2 * (sb & 15) + 1, 2 * (sb >> 4) + 1); | |
| } | |
| }; | |
| template <> struct Fmt32<18> { // IQ3_XXS: d, qs[64] grid bytes, 8 x u32 (4 x 7-bit sign index + 4-bit scale) | |
| static constexpr int bytes = 98; | |
| static constexpr float K = 0.25f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| const uint8_t* q = b + 2 + 16 * j + 8 * half; | |
| g = _mm256_set_epi32((int) iq3xxs_grid[q[7]], (int) iq3xxs_grid[q[6]], (int) iq3xxs_grid[q[5]], (int) iq3xxs_grid[q[4]], | |
| (int) iq3xxs_grid[q[3]], (int) iq3xxs_grid[q[2]], (int) iq3xxs_grid[q[1]], (int) iq3xxs_grid[q[0]]); | |
| const uint32_t w = u32(b + 2 + 64 + 8 * j + 4 * half); | |
| sgn = _mm256_set_epi64x((long long) even_signs.v[(w >> 21) & 127], (long long) even_signs.v[(w >> 14) & 127], | |
| (long long) even_signs.v[(w >> 7) & 127], (long long) even_signs.v[w & 127]); | |
| sc = sc32(2 * (int) (w >> 28) + 1); | |
| } | |
| }; | |
| template <> struct Fmt32<21> { // IQ3_S: d, qs[64], qh[8], signs[32], scales[4] | |
| static constexpr int bytes = 110; | |
| static constexpr float K = 1.0f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| const uint8_t* q = b + 2 + 16 * j + 8 * half; | |
| const uint32_t h = b[66 + 2 * j + half]; | |
| g = _mm256_set_epi32(G3(7), G3(6), G3(5), G3(4), G3(3), G3(2), G3(1), G3(0)); | |
| const uint64_t m = u64(b + 74 + 8 * j); | |
| sgn = sgn_vec(half ? (uint32_t) (m >> 32) : (uint32_t) m); | |
| const uint8_t s = b[106 + j]; | |
| sc = sc32(half ? 2 * (s >> 4) + 1 : 2 * (s & 15) + 1); | |
| } | |
| }; | |
| template <> struct Fmt32<23> { // IQ4_XS: d, scales_h, scales_l[4], qs[128] - 136 B, 8 signed sub-scales | |
| // ggml's ggml_vec_dot_iq4_xs_q8_K: the same 16-value codebook as IQ4_NL, but each of the eight 32-value | |
| // sub-blocks carries its own 6-bit scale, read as two nibbles of scales_l[p] plus two bits of scales_h, and | |
| // used SIGNED as (ls - 32). The scale therefore folds into the int16 operand of madd_epi16 instead of | |
| // becoming a float multiply per sub-block, and the codebook sign is carried the way iq4nl_rows does it. | |
| static constexpr int bytes = 136; | |
| static constexpr float K = 1.0f; | |
| static inline void decode(const uint8_t* b, int j, int half, __m256i& g, __m256i& sgn, __m256i& sc) { | |
| const int H = 2 * j + half; // the eight 32-value sub-blocks, in the order the activation bytes come | |
| const __m128i values = _mm_loadu_si128((const __m128i*) kvalues_iq4nl); | |
| const __m128i bits = _mm_loadu_si128((const __m128i*) (b + 8 + 16 * H)); | |
| const __m128i m4 = _mm_set1_epi8(0x0f); | |
| const __m256i q4 = _mm256_inserti128_si256( | |
| _mm256_castsi128_si256(_mm_shuffle_epi8(values, _mm_and_si128(bits, m4))), | |
| _mm_shuffle_epi8(values, _mm_and_si128(_mm_srli_epi16(bits, 4), m4)), 1); | |
| g = _mm256_sign_epi8(q4, q4); // |w|, the unsigned operand of maddubs | |
| sgn = _mm256_sign_epi8(_mm256_set1_epi8(1), q4); // w's sign, applied to the activation | |
| const int ls = ((b[4 + (H >> 1)] >> (4 * (H & 1))) & 0xf) | (((u16(b + 2) >> (2 * H)) & 3) << 4); | |
| sc = _mm256_set1_epi16((short) (ls - 32)); | |
| } | |
| }; | |
| template <int TY, int NT> | |
| inline void row_dot(const uint8_t* row, int nblocks, const block_q8_K* const* y, float* res) { | |
| __m256 accf[NT]; | |
| for (int t = 0; t < NT; ++t) accf[t] = _mm256_setzero_ps(); | |
| for (int i = 0; i < nblocks; ++i) { | |
| const uint8_t* blk = row + (size_t) i * Fmt32<TY>::bytes; | |
| rows_ahead(blk + prefetch_ahead); | |
| __m256i acci[NT]; | |
| for (int t = 0; t < NT; ++t) acci[t] = _mm256_setzero_si256(); | |
| for (int j = 0; j < 4; ++j) { | |
| for (int half = 0; half < 2; ++half) { | |
| __m256i g, sgn, sc; | |
| Fmt32<TY>::decode(blk, j, half, g, sgn, sc); | |
| const int off = 64 * j + 32 * half; | |
| for (int t = 0; t < NT; ++t) { | |
| const __m256i yv = _mm256_loadu_si256((const __m256i*) (y[t][i].qs + off)); | |
| const __m256i ys = _mm256_sign_epi8(yv, sgn); | |
| acci[t] = _mm256_add_epi32(acci[t], _mm256_madd_epi16(_mm256_maddubs_epi16(g, ys), sc)); | |
| } | |
| } | |
| } | |
| const float dx = h2f(u16(blk)) * Fmt32<TY>::K; | |
| for (int t = 0; t < NT; ++t) | |
| accf[t] = _mm256_fmadd_ps(_mm256_set1_ps(dx * y[t][i].d), _mm256_cvtepi32_ps(acci[t]), accf[t]); | |
| } | |
| for (int t = 0; t < NT; ++t) res[t] = hsum8(accf[t]); | |
| } | |
| // ---- IQ2_XS (17) with ggml's vectorized sign decode (arch/x86/quants.c, ggml_vec_dot_iq2_xs_q8_K): | |
| // the 7-bit sign indices never touch ksigns_iq2xs - the 8th sign bit is reconstructed by parity (two | |
| // shifts, a xor and ggml's bit_helper pshufb) and the per-half sign vectors come out of shuffles of one | |
| // 32-byte load. The 8 scale bytes become all 16 half-scales with ggml's unpack trick. The grid lookups | |
| // stay scalar loads into set_epi64x (ggml does the same). Every token still only pays a load, a sign, | |
| // a maddubs, a madd and an add per half, into alternating accumulators. | |
| template <int NT> | |
| inline void row_dot_iq2xs(const uint8_t* row, int nblocks, const block_q8_K* const* y, float* res) { | |
| static const uint8_t bit_sel[32] = { | |
| 1, 2, 4, 8, 16, 32, 64, (uint8_t) 0x80, 1, 2, 4, 8, 16, 32, 64, (uint8_t) 0x80, | |
| 1, 2, 4, 8, 16, 32, 64, (uint8_t) 0x80, 1, 2, 4, 8, 16, 32, 64, (uint8_t) 0x80 }; | |
| static const char shuf_even[32] = { // sign bits of u16 0,1,2,3 (even bytes of the sign register) | |
| 0, 0, 0, 0, 0, 0, 0, 0, 2, 2, 2, 2, 2, 2, 2, 2, | |
| 4, 4, 4, 4, 4, 4, 4, 4, 6, 6, 6, 6, 6, 6, 6, 6 }; | |
| static const char shuf_odd[32] = { // sign bits of u16 4,5,6,7 | |
| 8, 8, 8, 8, 8, 8, 8, 8, 10, 10, 10, 10, 10, 10, 10, 10, | |
| 12, 12, 12, 12, 12, 12, 12, 12, 14, 14, 14, 14, 14, 14, 14, 14 }; | |
| static const uint8_t bit_help[32] = { // 0x80 when the 4-bit index has odd popcount (parity bit 7) | |
| 0x00, (uint8_t) 0x80, (uint8_t) 0x80, 0x00, (uint8_t) 0x80, 0x00, 0x00, (uint8_t) 0x80, | |
| (uint8_t) 0x80, 0x00, 0x00, (uint8_t) 0x80, 0x00, (uint8_t) 0x80, (uint8_t) 0x80, 0x00, | |
| 0x00, (uint8_t) 0x80, (uint8_t) 0x80, 0x00, (uint8_t) 0x80, 0x00, 0x00, (uint8_t) 0x80, | |
| (uint8_t) 0x80, 0x00, 0x00, (uint8_t) 0x80, 0x00, (uint8_t) 0x80, (uint8_t) 0x80, 0x00 }; | |
| static const uint8_t k_sc_shuffle[128] = { // half h -> scale bytes 2h (lanes 0-7), 2h+1 (lanes 8-15) | |
| 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, | |
| 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, | |
| 4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5, | |
| 6, 6, 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7, | |
| 8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9, | |
| 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11, | |
| 12, 12, 12, 12, 12, 12, 12, 12, 13, 13, 13, 13, 13, 13, 13, 13, | |
| 14, 14, 14, 14, 14, 14, 14, 14, 15, 15, 15, 15, 15, 15, 15, 15 }; | |
| const __m256i bsel = _mm256_loadu_si256((const __m256i*) bit_sel); | |
| const __m256i shev = _mm256_loadu_si256((const __m256i*) shuf_even); | |
| const __m256i shod = _mm256_loadu_si256((const __m256i*) shuf_odd); | |
| const __m256i bhelp = _mm256_loadu_si256((const __m256i*) bit_help); | |
| const __m256i m511 = _mm256_set1_epi16(511); | |
| const __m128i m4 = _mm_set1_epi8(0xf), m1 = _mm_set1_epi8(1); | |
| const __m256i one8 = _mm256_set1_epi8(1); | |
| __m256 accf[NT]; | |
| for (int t = 0; t < NT; ++t) accf[t] = _mm256_setzero_ps(); | |
| for (int i = 0; i < nblocks; ++i) { | |
| const uint8_t* blk = row + (size_t) i * 74; | |
| rows_ahead(blk + prefetch_ahead); | |
| // the 8 scale bytes -> 16 half-scales of 2*s+1, interleaved [a0, b0, a1, b1, ...] (ggml's unpack) | |
| __m128i st = _mm_set1_epi64x((long long) u64(blk + 66)); | |
| st = _mm_unpacklo_epi8(_mm_and_si128(st, m4), _mm_and_si128(_mm_srli_epi16(st, 4), m4)); | |
| const __m128i scales = _mm_add_epi8(_mm_slli_epi16(st, 1), m1); | |
| __m256i acc[NT][2]; | |
| for (int t = 0; t < NT; ++t) acc[t][0] = acc[t][1] = _mm256_setzero_si256(); | |
| for (int jj = 0; jj < 2; ++jj) { // 128 values (16 u16) per step | |
| const __m256i q2 = _mm256_loadu_si256((const __m256i*) (blk + 2 + 32 * jj)); | |
| const __m256i p7 = _mm256_srli_epi16(q2, 9); // 7 sign bits per u16, even bytes | |
| const __m256i p3 = _mm256_srli_epi16(q2, 13); // sign index bits 4-6 | |
| const __m256i fsb = _mm256_or_si256(p7, | |
| _mm256_shuffle_epi8(bhelp, _mm256_xor_si256(p7, p3))); // + the parity sign bit 7 | |
| const __m256i s01 = _mm256_broadcastsi128_si256(_mm256_castsi256_si128(fsb)); // u16 0-7 | |
| const __m256i s23 = _mm256_broadcastsi128_si256(_mm256_extracti128_si256(fsb, 1)); // u16 8-15 | |
| alignas(32) uint16_t gi[16]; | |
| _mm256_store_si256((__m256i*) gi, _mm256_and_si256(q2, m511)); | |
| for (int h = 0; h < 4; ++h) { // the four 32-value halves of this group | |
| const int H = 4 * jj + h; // global half index -> scales byte pair | |
| const __m256i g = _mm256_set_epi64x((long long) iq2xs_grid[gi[4 * h + 3]], | |
| (long long) iq2xs_grid[gi[4 * h + 2]], | |
| (long long) iq2xs_grid[gi[4 * h + 1]], | |
| (long long) iq2xs_grid[gi[4 * h]]); | |
| const __m256i sb = _mm256_shuffle_epi8(h < 2 ? s01 : s23, (h & 1) ? shod : shev); | |
| const __m256i sgn = _mm256_or_si256( | |
| _mm256_cmpeq_epi8(_mm256_and_si256(sb, bsel), bsel), one8); | |
| const __m256i sc = _mm256_cvtepi8_epi16( | |
| _mm_shuffle_epi8(scales, _mm_loadu_si128((const __m128i*) k_sc_shuffle + H))); | |
| const int off = 128 * jj + 32 * h; | |
| for (int t = 0; t < NT; ++t) { | |
| const __m256i yv = _mm256_loadu_si256((const __m256i*) (y[t][i].qs + off)); | |
| const __m256i ys = _mm256_sign_epi8(yv, sgn); | |
| acc[t][h & 1] = _mm256_add_epi32(acc[t][h & 1], | |
| _mm256_madd_epi16(_mm256_maddubs_epi16(g, ys), sc)); | |
| } | |
| } | |
| } | |
| const float dx = h2f(u16(blk)) * 0.125f; | |
| for (int t = 0; t < NT; ++t) | |
| accf[t] = _mm256_fmadd_ps(_mm256_set1_ps(dx * y[t][i].d), | |
| _mm256_cvtepi32_ps(_mm256_add_epi32(acc[t][0], acc[t][1])), accf[t]); | |
| } | |
| for (int t = 0; t < NT; ++t) res[t] = hsum8(accf[t]); | |
| } | |
| template <int TY, int NT> | |
| inline void row_dot_any(const uint8_t* row, int nblocks, const block_q8_K* const* y, float* res) { | |
| if constexpr (TY == 17) row_dot_iq2xs<NT>(row, nblocks, y, res); | |
| else row_dot<TY, NT>(row, nblocks, y, res); | |
| } | |
| template <int TY, int NT> | |
| void gu_rows(const uint8_t* blob, size_t gu_row, size_t up_off, int n, const void* const* act, float* const* ff, | |
| int r0, int r1) { | |
| const block_q8_K* y[NT]; | |
| for (int t = 0; t < NT; ++t) y[t] = (const block_q8_K*) act[t]; | |
| const int nb = n / QK_K; | |
| float g[NT], u[NT]; | |
| for (int r = r0; r < r1; ++r) { | |
| row_dot_any<TY, NT>(blob + (size_t) r * gu_row, nb, y, g); | |
| row_dot_any<TY, NT>(blob + up_off + (size_t) r * gu_row, nb, y, u); | |
| for (int t = 0; t < NT; ++t) ff[t][r] = (g[t] / (1.f + std::exp(-g[t]))) * u[t]; | |
| } | |
| } | |
| template <int TY, int NT> | |
| void dot_rows(const uint8_t* w, size_t row_bytes, int n, const void* const* act, float* const* out, int r0, int r1) { | |
| const block_q8_K* y[NT]; | |
| for (int t = 0; t < NT; ++t) y[t] = (const block_q8_K*) act[t]; | |
| float res[NT]; | |
| for (int r = r0; r < r1; ++r) { | |
| row_dot_any<TY, NT>(w + (size_t) r * row_bytes, n / QK_K, y, res); | |
| for (int t = 0; t < NT; ++t) out[t][r] = res[t]; | |
| } | |
| } | |
| template <int TY> | |
| void gu_rows_nt(int nt, const uint8_t* blob, size_t gu_row, size_t up_off, int n, const void* const* act, | |
| float* const* ff, int r0, int r1) { | |
| switch (nt) { | |
| case 1: gu_rows<TY, 1>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 2: gu_rows<TY, 2>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 3: gu_rows<TY, 3>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 4: gu_rows<TY, 4>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 5: gu_rows<TY, 5>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 6: gu_rows<TY, 6>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 7: gu_rows<TY, 7>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| default: gu_rows<TY, 8>(blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| } | |
| } | |
| template <int TY> | |
| void dot_rows_nt(int nt, const uint8_t* w, size_t row_bytes, int n, const void* const* act, float* const* out, int r0, | |
| int r1) { | |
| switch (nt) { | |
| case 1: dot_rows<TY, 1>(w, row_bytes, n, act, out, r0, r1); break; | |
| case 2: dot_rows<TY, 2>(w, row_bytes, n, act, out, r0, r1); break; | |
| case 3: dot_rows<TY, 3>(w, row_bytes, n, act, out, r0, r1); break; | |
| case 4: dot_rows<TY, 4>(w, row_bytes, n, act, out, r0, r1); break; | |
| default: for (int t0 = 0; t0 < nt; t0 += 4) { | |
| const int k = nt - t0 < 4 ? nt - t0 : 4; | |
| dot_rows_nt<TY>(k, w, row_bytes, n, act + t0, out + t0, r0, r1); | |
| } | |
| } | |
| } | |
| } // namespace | |
| bool iq256_supported(int type) noexcept { | |
| return type == 16 || type == 17 || type == 18 || type == 21 || type == 22 || type == 23; | |
| } | |
| void iq256_gu_rows(int type, const uint8_t* blob, size_t gu_row, size_t up_off, int n, const void* const* act, int nt, | |
| float* const* ff, int r0, int r1) { | |
| switch (type) { | |
| case 16: gu_rows_nt<16>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 17: gu_rows_nt<17>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 18: gu_rows_nt<18>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 21: gu_rows_nt<21>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 22: gu_rows_nt<22>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| case 23: gu_rows_nt<23>(nt, blob, gu_row, up_off, n, act, ff, r0, r1); break; | |
| default: break; | |
| } | |
| } | |
| void iq256_rows(int type, const uint8_t* w, size_t row_bytes, int n, const void* const* act, int nt, float* const* out, | |
| int r0, int r1) { | |
| switch (type) { | |
| case 16: dot_rows_nt<16>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| case 17: dot_rows_nt<17>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| case 18: dot_rows_nt<18>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| case 21: dot_rows_nt<21>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| case 22: dot_rows_nt<22>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| case 23: dot_rows_nt<23>(nt, w, row_bytes, n, act, out, r0, r1); break; | |
| default: break; | |
| } | |
| } | |
| // ---- IQ4_NL (type 20): 32-value blocks of 18 bytes (f16 d + 16 nibble bytes) against Q8_0 activations. | |
| // ggml-cpu's own AVX-2 dot (ggml_vec_dot_iq4_nl_q8_0) is single-token: every token redoes the nibble -> | |
| // kvalues_iq4nl pshufb decode and the |w| half of the signed x signed product. Here both are computed once | |
| // per block; every token then costs a sign, a maddubs, a madd and an fmadd. The arithmetic is ggml's - only | |
| // the order of the float additions differs. | |
| template <int NT> | |
| void iq4nl_rows(const uint8_t* w, size_t row_bytes, int n, const block_q8_0* const* y, float* const* out, | |
| int r0, int r1) { | |
| const __m128i values = _mm_loadu_si128((const __m128i*) kvalues_iq4nl); | |
| const __m128i m4b = _mm_set1_epi8(0x0f); | |
| const __m256i ones = _mm256_set1_epi16(1); | |
| const int nb = n / QK4_NL; | |
| for (int r = r0; r < r1; ++r) { | |
| const uint8_t* row = w + (size_t) r * row_bytes; | |
| __m256 accf[NT]; | |
| for (int t = 0; t < NT; ++t) accf[t] = _mm256_setzero_ps(); | |
| for (int ib = 0; ib < nb; ++ib) { | |
| const uint8_t* blk = row + (size_t) ib * sizeof(block_iq4_nl); | |
| rows_ahead(blk + prefetch_ahead); | |
| const __m128i bits = _mm_loadu_si128((const __m128i*) (blk + 2)); | |
| const __m128i lo = _mm_and_si128(bits, m4b); // values 0..15 | |
| const __m128i hi = _mm_and_si128(_mm_srli_epi16(bits, 4), m4b); // values 16..31 | |
| const __m256i q4 = _mm256_inserti128_si256(_mm256_castsi128_si256(_mm_shuffle_epi8(values, lo)), | |
| _mm_shuffle_epi8(values, hi), 1); | |
| const __m256i aq = _mm256_sign_epi8(q4, q4); // |w|: the unsigned operand of maddubs | |
| const float dx = h2f(u16(blk)); | |
| for (int t = 0; t < NT; ++t) { | |
| const block_q8_0& b = y[t][ib]; | |
| const __m256i q8 = _mm256_loadu_si256((const __m256i*) b.qs); | |
| const __m256i p = _mm256_madd_epi16(_mm256_maddubs_epi16(aq, _mm256_sign_epi8(q8, q4)), ones); | |
| accf[t] = _mm256_fmadd_ps(_mm256_set1_ps(dx * h2f(b.d)), _mm256_cvtepi32_ps(p), accf[t]); | |
| } | |
| } | |
| for (int t = 0; t < NT; ++t) out[t][r] = hsum8(accf[t]); | |
| } | |
| } | |
| void iq4nl256_down_rows(const uint8_t* w, size_t row_bytes, int n, const void* const* hq, int nt, float* const* out, | |
| int r0, int r1) { | |
| const block_q8_0* y[8]; | |
| for (int t = 0; t < nt; ++t) y[t] = (const block_q8_0*) hq[t]; | |
| switch (nt) { | |
| case 1: iq4nl_rows<1>(w, row_bytes, n, y, out, r0, r1); break; | |
| case 2: iq4nl_rows<2>(w, row_bytes, n, y, out, r0, r1); break; | |
| case 3: iq4nl_rows<3>(w, row_bytes, n, y, out, r0, r1); break; | |
| case 4: iq4nl_rows<4>(w, row_bytes, n, y, out, r0, r1); break; | |
| default: for (int t0 = 0; t0 < nt; t0 += 4) { | |
| const int k = nt - t0 < 4 ? nt - t0 : 4; | |
| iq4nl256_down_rows(w, row_bytes, n, hq + t0, k, out + t0, r0, r1); | |
| } | |
| } | |
| } | |
| } // namespace strata::kernels::cpu | |