← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q4k_q8k_vnni.c File Reference

VNNI Q4_K x Q8_K matvec kernel (inference only) More...

#include <stddef.h>
#include <stdint.h>
#include <math.h>
#include <stdlib.h>
#include <string.h>
#include "ckernel_engine.h"
#include "ck_threadpool.h"
#include "ckernel_quant.h"
#include "ck_speed_profiles.h"

Go to the source code of this file.

Functions

static void accum_q4_k_packed_meta_x16_q8_k_block (float acc[16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *x)
 
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse (float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
 
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4 (float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
 
static void accum_q4_k_packed_meta_x8_q8_k_block (float acc[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
 
static void accum_q4_k_packed_meta_x8_q8_k_block_mreuse (float acc[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
 
static void accum_q4_k_packed_meta_x8_q8_k_gemv_block (float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
 
static void accum_q4_k_packed_meta_x8_q8_k_superblock (float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
 
static void accum_q4_k_packed_meta_x8_q8_k_superblock_rows (float acc[8][8], float acc_min[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x[8], int rows)
 
static void accum_q4_k_packed_u8_x16_q8_k_block (float acc[16], const block_q4_K_packed_u8_x16 *w, int active, const block_q8_K *x)
 
static void accum_q4_k_packed_vnni_x16_q8_k_16m_superblock (float acc[16][16], float acc_min[16][16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x[16], int rows)
 
static void accum_q4_k_packed_vnni_x16_q8_k_gemv_block (float acc[16], float acc_min[16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x)
 
static void accum_q4_k_packed_vnni_x8_q8_k_4m_superblock (float acc[4][8], float acc_min[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)
 
int ck_q4k_packed_vnni_x16_available (void)
 
int ck_q4k_packed_vnni_x8_available (void)
 
int ck_q4k_packed_vnni_x8_compact_order_available (void)
 
static float ck_q4k_silu_f32 (float x)
 
static int ck_q4k_x16_chunk4_enabled (void)
 
static float dot_q4_k_packed_meta_q8_k_block (const block_q4_K_packed_meta *w, const block_q8_K *x)
 
static float dot_q4_k_packed_u8_q8_k_block (const block_q4_K_packed_u8 *w, const block_q8_K *x)
 
static void dot_q4_k_packed_vnni_x8_q8_k_compact_order (float block_sums[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)
 
static int32_t dot_q4_packed_u8_q8_32_ref (const uint8_t *q4_32, const int8_t *q8_32)
 
void gemm_nt_q4_k_packed_meta_q8_k (const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_meta_q8_k_threaded (const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit (const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_meta_q8_k_tile (const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1)
 
void gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni (const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int D, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse (const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile (const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit (const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_u8_q8_k (const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile (const void *A_q8, const void *B_packed_u8_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
 
void gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order (const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m (const void *A_q8, const void *B_packed_vnni_x16, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m (const void *A_q8, const void *B_packed_vnni_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
 
void gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni (const void *A_q8, const void *B_gate_up, const float *bias, float *C, int M, int D, int K, int threads)
 
static void gemm_q4_gateup_swiglu_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_gateup_swiglu_x16_thread_fn (int ith, int nth, void *args)
 
void gemm_q4_k_q8_k_compact_rows4 (float *output, int output_stride, const void *weights, const void *const input_rows[4], int rows, int output_dim, int input_dim)
 
void gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4 (float *output, const void *weights_packed, const void *input_q8, int rows, int output_dim, int input_dim)
 
static void gemm_q4_packed_meta_nsplit_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x16_mreuse_process_job (const gemm_q4_packed_meta_x16_thread_work_t *a, int job, int mt, int tile_m)
 
static void gemm_q4_packed_meta_x16_mreuse_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x16_mtile_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_mreuse_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_mtile_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_nsplit_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_split_min_4m_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_split_min_8m_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_u8_x16_mtile_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_vnni_x16_q8k_16m_thread_fn (int ith, int nth, void *args)
 
static void gemm_q4_packed_vnni_x8_q8k_4m_job (gemm_q4_packed_vnni_x8_thread_work_t *a, int job, int row_tiles)
 
static void gemm_q4_packed_vnni_x8_q8k_4m_range_fn (int begin, int end, void *args)
 
static void gemm_q4_packed_vnni_x8_q8k_4m_thread_fn (int ith, int nth, void *args)
 
void gemv_q4_k_q8_k_avx2 (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q4_k_q8_k_parallel_vnni (float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
 
void gemv_q4_k_q8_k_ref (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q4_k_q8_k_vnni (float *y, const void *W, const void *x_q8, int M, int K)
 
void pack_q4_k_to_packed_meta (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_meta_x16 (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_meta_x8 (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_u8 (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_u8_x16 (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_vnni_x16 (const void *src, void *dst, int N, int K)
 
void pack_q4_k_to_packed_vnni_x8 (const void *src, void *dst, int N, int K)
 
size_t q4_k_packed_meta_block_size (void)
 
size_t q4_k_packed_meta_x16_block_size (void)
 
size_t q4_k_packed_meta_x8_block_size (void)
 
size_t q4_k_packed_u8_block_size (void)
 
size_t q4_k_packed_u8_x16_block_size (void)
 
size_t q4_k_packed_vnni_x16_block_size (void)
 
size_t q4_k_packed_vnni_x8_block_size (void)
 

Detailed Description

VNNI Q4_K x Q8_K matvec kernel (inference only)

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

After changes: make test && make llamacpp-parity-full

The canonical providers require AVX2. The x8 output-interleaved provider uses 256-bit VNNI from AVX-VNNI or AVX-512 VNNI+VL. The x16 provider is a separate AVX-512 VNNI diagnostic path; production promotion is sweep-gated.

Packed-meta status: Packed layouts are internal providers selected by the declared production dispatcher. Weight-identity caches own their lifetime until runtime shutdown; the kernel map records layout and ISA requirements. Canonical GGUF-layout Q4_K remains the parity fallback for unsupported ISAs and uncovered shapes.

Definition in file gemm_kernels_q4k_q8k_vnni.c.

Function Documentation

◆ accum_q4_k_packed_meta_x16_q8_k_block()

static void accum_q4_k_packed_meta_x16_q8_k_block ( float  acc[16],
const block_q4_K_packed_meta_x16 *  w,
int  active,
const block_q8_K x 
)
inlinestatic

Definition at line 1712 of file gemm_kernels_q4k_q8k_vnni.c.

1716{
1717 const float xd = x->d;
1718 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1719 const int8_t *q8_lo_ptr = &x->qs[j];
1720 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1721 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1722 (int32_t)x->bsums[j / 16 + 1];
1723 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1724 (int32_t)x->bsums[(j + 32) / 16 + 1];
1725
1726#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1727 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1728 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1729#elif defined(__AVX2__)
1730 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1731 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1732#endif
1733
1734 for (int lane = 0; lane < active; ++lane) {
1735 const uint8_t *qs = &w->qs[lane][q_offset];
1736#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1737 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
1738 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
1739#elif defined(__AVX2__)
1740 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1741 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1742 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1743 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo, q8_lo, w->sc[lane][is]);
1744 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi, q8_hi, w->sc[lane][is + 1]);
1745 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1746#else
1747 int32_t sum_lo = 0;
1748 int32_t sum_hi = 0;
1749 for (int l = 0; l < 32; ++l) {
1750 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1751 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1752 }
1753#endif
1754 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
1755 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
1756#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1757 acc[lane] += d * (float)sum_scaled;
1758 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
1759 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
1760#else
1761 acc[lane] += d * (float)w->sc[lane][is] * (float)sum_lo;
1762 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
1763 acc[lane] += d * (float)w->sc[lane][is + 1] * (float)sum_hi;
1764 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
1765#endif
1766 }
1767 }
1768}
#define CK_FP16_TO_FP32(x)
#define QK_K
int8_t qs[256]
int16_t bsums[256/16]

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_q4_packed_meta_x16_mtile_thread_fn().

◆ accum_q4_k_packed_meta_x16_q8_k_block_mreuse()

static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse ( float  acc[8][16],
const block_q4_K_packed_meta_x16 *  w,
int  active,
const block_q8_K A,
int  blocks_per_vec,
int  block_index,
int  m0,
int  m_count 
)
inlinestatic

Definition at line 1561 of file gemm_kernels_q4k_q8k_vnni.c.

1569{
1570 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1571 for (int lane = 0; lane < active; ++lane) {
1572 const uint8_t *qs = &w->qs[lane][q_offset];
1573#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1574 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1575 const __m256i q4_lo = q4_k_unpack_32_vnni_bytes(packed, 0);
1576 const __m256i q4_hi = q4_k_unpack_32_vnni_bytes(packed, 1);
1577#elif defined(__AVX2__)
1578 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1579 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1580 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1581#endif
1582 const float wd = CK_FP16_TO_FP32(w->d[lane]);
1583 const float wdmin = CK_FP16_TO_FP32(w->dmin[lane]);
1584#if !defined(__AVX2__) || (defined(__AVX512VNNI__) && defined(__AVX512VL__))
1585 const float sc_lo = (float)w->sc[lane][is];
1586 const float sc_hi = (float)w->sc[lane][is + 1];
1587#endif
1588 const float min_lo = (float)w->m[lane][is];
1589 const float min_hi = (float)w->m[lane][is + 1];
1590
1591 for (int mt = 0; mt < m_count; ++mt) {
1592 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1593 const int8_t *q8_lo_ptr = &x->qs[j];
1594 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1595 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1596 (int32_t)x->bsums[j / 16 + 1];
1597 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1598 (int32_t)x->bsums[(j + 32) / 16 + 1];
1599#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1600 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1601 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1602 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo, q8_lo);
1603 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi, q8_hi);
1604#elif defined(__AVX2__)
1605 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1606 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1607 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo, q8_lo, w->sc[lane][is]);
1608 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi, q8_hi, w->sc[lane][is + 1]);
1609 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1610#else
1611 int32_t sum_lo = 0;
1612 int32_t sum_hi = 0;
1613 for (int l = 0; l < 32; ++l) {
1614 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1615 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1616 }
1617#endif
1618 const float xd = x->d;
1619 const float d = wd * xd;
1620 const float dmin = wdmin * xd;
1621#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1622 acc[mt][lane] += d * (float)sum_scaled;
1623 acc[mt][lane] -= dmin * min_lo * (float)bsum_lo;
1624 acc[mt][lane] -= dmin * min_hi * (float)bsum_hi;
1625#else
1626 acc[mt][lane] += d * sc_lo * (float)sum_lo;
1627 acc[mt][lane] -= dmin * min_lo * (float)bsum_lo;
1628 acc[mt][lane] += d * sc_hi * (float)sum_hi;
1629 acc[mt][lane] -= dmin * min_hi * (float)bsum_hi;
1630#endif
1631 }
1632 }
1633 }
1634}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(), gemm_q4_gateup_swiglu_x16_thread_fn(), and gemm_q4_packed_meta_x16_mreuse_process_job().

◆ accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4()

static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4 ( float  acc[8][16],
const block_q4_K_packed_meta_x16 *  w,
int  active,
const block_q8_K A,
int  blocks_per_vec,
int  block_index,
int  m0,
int  m_count 
)
inlinestatic

Definition at line 1636 of file gemm_kernels_q4k_q8k_vnni.c.

1644{
1645#if (defined(__AVX512VNNI__) && defined(__AVX512VL__)) || defined(__AVX2__)
1646 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1647 for (int lane0 = 0; lane0 < active; lane0 += 4) {
1648 const int lanes = (lane0 + 4 <= active) ? 4 : (active - lane0);
1649 __m256i q4_lo[4];
1650 __m256i q4_hi[4];
1651 float wd[4], wdmin[4], sc_lo[4], sc_hi[4], min_lo[4], min_hi[4];
1652
1653 for (int l = 0; l < lanes; ++l) {
1654 const int lane = lane0 + l;
1655 const uint8_t *qs = &w->qs[lane][q_offset];
1656 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1657#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1658 q4_lo[l] = q4_k_unpack_32_vnni_bytes(packed, 0);
1659 q4_hi[l] = q4_k_unpack_32_vnni_bytes(packed, 1);
1660#else
1661 q4_lo[l] = q4_k_unpack_32_avx2_bytes(packed, 0);
1662 q4_hi[l] = q4_k_unpack_32_avx2_bytes(packed, 1);
1663#endif
1664 wd[l] = CK_FP16_TO_FP32(w->d[lane]);
1665 wdmin[l] = CK_FP16_TO_FP32(w->dmin[lane]);
1666 sc_lo[l] = (float)w->sc[lane][is];
1667 sc_hi[l] = (float)w->sc[lane][is + 1];
1668 min_lo[l] = (float)w->m[lane][is];
1669 min_hi[l] = (float)w->m[lane][is + 1];
1670 }
1671
1672 for (int mt = 0; mt < m_count; ++mt) {
1673 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1674 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1675 (int32_t)x->bsums[j / 16 + 1];
1676 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1677 (int32_t)x->bsums[(j + 32) / 16 + 1];
1678 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)&x->qs[j]);
1679 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)&x->qs[j + 32]);
1680 const float xd = x->d;
1681 for (int l = 0; l < lanes; ++l) {
1682 const int lane = lane0 + l;
1683#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1684 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo[l], q8_lo);
1685 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi[l], q8_hi);
1686#else
1687 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo[l], q8_lo, w->sc[lane][is]);
1688 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi[l], q8_hi, w->sc[lane][is + 1]);
1689 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1690#endif
1691 const float d = wd[l] * xd;
1692 const float dmin = wdmin[l] * xd;
1693#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1694 acc[mt][lane] += d * (float)sum_scaled;
1695 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1696 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1697#else
1698 acc[mt][lane] += d * sc_lo[l] * (float)sum_lo;
1699 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1700 acc[mt][lane] += d * sc_hi[l] * (float)sum_hi;
1701 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1702#endif
1703 }
1704 }
1705 }
1706 }
1707#else
1708 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc, w, active, A, blocks_per_vec, block_index, m0, m_count);
1709#endif
1710}
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse(float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)

References accum_q4_k_packed_meta_x16_q8_k_block_mreuse(), block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_q4_gateup_swiglu_x16_thread_fn(), and gemm_q4_packed_meta_x16_mreuse_process_job().

◆ accum_q4_k_packed_meta_x8_q8_k_block()

static void accum_q4_k_packed_meta_x8_q8_k_block ( float  acc[8],
const block_q4_K_packed_meta_x8 *  w,
int  active,
const block_q8_K x 
)
inlinestatic

Definition at line 703 of file gemm_kernels_q4k_q8k_vnni.c.

707{
708 const float xd = x->d;
709 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
710 const int8_t *q8_lo_ptr = &x->qs[j];
711 const int8_t *q8_hi_ptr = &x->qs[j + 32];
712 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
713 (int32_t)x->bsums[j / 16 + 1];
714 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
715 (int32_t)x->bsums[(j + 32) / 16 + 1];
716
717#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
718 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
719 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
720#elif defined(__AVX2__)
721 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
722 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
723#endif
724
725 for (int lane = 0; lane < active; ++lane) {
726 const uint8_t *qs = &w->qs[lane][q_offset];
727#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
728 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
729 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
730#elif defined(__AVX2__)
731 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
732 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
733 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
734 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
735 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
736#else
737 int32_t sum_lo = 0;
738 int32_t sum_hi = 0;
739 for (int l = 0; l < 32; ++l) {
740 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
741 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
742 }
743#endif
744 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
745 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
746 acc[lane] += d * (float)w->sc[lane][is] * (float)sum_lo;
747 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
748 acc[lane] += d * (float)w->sc[lane][is + 1] * (float)sum_hi;
749 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
750 }
751 }
752}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by accum_q4_k_packed_meta_x8_q8_k_block_mreuse(), gemm_nt_q4_k_packed_meta_x8_q8_k(), gemm_q4_packed_meta_x8_mtile_thread_fn(), and gemm_q4_packed_meta_x8_nsplit_thread_fn().

◆ accum_q4_k_packed_meta_x8_q8_k_block_mreuse()

static void accum_q4_k_packed_meta_x8_q8_k_block_mreuse ( float  acc[8][8],
const block_q4_K_packed_meta_x8 *  w,
int  active,
const block_q8_K A,
int  blocks_per_vec,
int  block_index,
int  m0,
int  m_count 
)
inlinestatic

Definition at line 1469 of file gemm_kernels_q4k_q8k_vnni.c.

1477{
1478#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1479 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1480 for (int lane0 = 0; lane0 < active; lane0 += 4) {
1481 const int lanes = (lane0 + 4 <= active) ? 4 : (active - lane0);
1482 __m256i q4_lo[4];
1483 __m256i q4_hi[4];
1484 float wd[4], wdmin[4], sc_lo[4], sc_hi[4], min_lo[4], min_hi[4];
1485
1486 for (int l = 0; l < lanes; ++l) {
1487 const int lane = lane0 + l;
1488 const uint8_t *qs = &w->qs[lane][q_offset];
1489 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1490 q4_lo[l] = q4_k_unpack_32_vnni_bytes(packed, 0);
1491 q4_hi[l] = q4_k_unpack_32_vnni_bytes(packed, 1);
1492 wd[l] = CK_FP16_TO_FP32(w->d[lane]);
1493 wdmin[l] = CK_FP16_TO_FP32(w->dmin[lane]);
1494 sc_lo[l] = (float)w->sc[lane][is];
1495 sc_hi[l] = (float)w->sc[lane][is + 1];
1496 min_lo[l] = (float)w->m[lane][is];
1497 min_hi[l] = (float)w->m[lane][is + 1];
1498 }
1499
1500 for (int mt = 0; mt < m_count; ++mt) {
1501 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1502 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1503 (int32_t)x->bsums[j / 16 + 1];
1504 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1505 (int32_t)x->bsums[(j + 32) / 16 + 1];
1506 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)&x->qs[j]);
1507 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)&x->qs[j + 32]);
1508 const float xd = x->d;
1509 for (int l = 0; l < lanes; ++l) {
1510 const int lane = lane0 + l;
1511 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo[l], q8_lo);
1512 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi[l], q8_hi);
1513 const float d = wd[l] * xd;
1514 const float dmin = wdmin[l] * xd;
1515 acc[mt][lane] += d * sc_lo[l] * (float)sum_lo;
1516 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1517 acc[mt][lane] += d * sc_hi[l] * (float)sum_hi;
1518 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1519 }
1520 }
1521 }
1522 }
1523#else
1524 for (int mt = 0; mt < m_count; ++mt) {
1525 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1526 accum_q4_k_packed_meta_x8_q8_k_block(acc[mt], w, active, x);
1527 }
1528#endif
1529}
static void accum_q4_k_packed_meta_x8_q8_k_block(float acc[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)

References accum_q4_k_packed_meta_x8_q8_k_block(), block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_q4_packed_meta_x8_mreuse_thread_fn().

◆ accum_q4_k_packed_meta_x8_q8_k_gemv_block()

static void accum_q4_k_packed_meta_x8_q8_k_gemv_block ( float  acc[8],
float  acc_min[8],
const block_q4_K_packed_meta_x8 *  w,
int  active,
const block_q8_K x 
)
inlinestatic

Definition at line 1393 of file gemm_kernels_q4k_q8k_vnni.c.

1397{
1398 int32_t iacc[8] = {0};
1399 int32_t iacc_min[8] = {0};
1400 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1401 const int8_t *q8_lo_ptr = &x->qs[j];
1402 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1403 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1404 (int32_t)x->bsums[j / 16 + 1];
1405 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1406 (int32_t)x->bsums[(j + 32) / 16 + 1];
1407#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1408 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1409 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1410#elif defined(__AVX2__)
1411 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1412 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1413#endif
1414 for (int lane = 0; lane < active; ++lane) {
1415 const uint8_t *qs = &w->qs[lane][q_offset];
1416#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1417 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
1418 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
1419#elif defined(__AVX2__)
1420 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1421 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1422 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1423 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
1424 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
1425#else
1426 int32_t sum_lo = 0;
1427 int32_t sum_hi = 0;
1428 for (int l = 0; l < 32; ++l) {
1429 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1430 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1431 }
1432#endif
1433 iacc[lane] += (int32_t)w->sc[lane][is] * sum_lo +
1434 (int32_t)w->sc[lane][is + 1] * sum_hi;
1435 iacc_min[lane] += (int32_t)w->m[lane][is] * bsum_lo +
1436 (int32_t)w->m[lane][is + 1] * bsum_hi;
1437 }
1438 }
1439
1440#if defined(__AVX2__)
1441 float d[8] = {0};
1442 float dmin[8] = {0};
1443 for (int lane = 0; lane < active; ++lane) {
1444 d[lane] = CK_FP16_TO_FP32(w->d[lane]);
1445 dmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1446 }
1447 const __m256 xd = _mm256_set1_ps(x->d);
1448 const __m256 acc_vec = _mm256_fmadd_ps(
1449 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc)),
1450 _mm256_mul_ps(_mm256_loadu_ps(d), xd), _mm256_loadu_ps(acc));
1451 const __m256 min_vec = _mm256_fmadd_ps(
1452 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc_min)),
1453 _mm256_mul_ps(_mm256_loadu_ps(dmin), xd), _mm256_loadu_ps(acc_min));
1454 _mm256_storeu_ps(acc, acc_vec);
1455 _mm256_storeu_ps(acc_min, min_vec);
1456#else
1457 for (int lane = 0; lane < active; ++lane) {
1458 const float d = CK_FP16_TO_FP32(w->d[lane]) * x->d;
1459 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * x->d;
1460 acc[lane] = fmaf((float)iacc[lane], d, acc[lane]);
1461 acc_min[lane] = fmaf((float)iacc_min[lane], dmin, acc_min[lane]);
1462 }
1463#endif
1464}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order().

◆ accum_q4_k_packed_meta_x8_q8_k_superblock()

static void accum_q4_k_packed_meta_x8_q8_k_superblock ( float  acc[8],
float  acc_min[8],
const block_q4_K_packed_meta_x8 *  w,
int  active,
const block_q8_K x 
)
inlinestatic

Definition at line 757 of file gemm_kernels_q4k_q8k_vnni.c.

761{
762 const float xd = x->d;
763#if defined(__AVX2__)
764 float d[8] = {0};
765 float dmin[8] = {0};
766 for (int lane = 0; lane < active; ++lane) {
767 d[lane] = CK_FP16_TO_FP32(w->d[lane]);
768 dmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
769 }
770 const __m256 scale = _mm256_mul_ps(_mm256_loadu_ps(d), _mm256_set1_ps(xd));
771 const __m256 min_scale = _mm256_mul_ps(_mm256_loadu_ps(dmin), _mm256_set1_ps(xd));
772#endif
773
774 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
775 int32_t iacc[8] = {0};
776 int32_t iacc_min[8] = {0};
777 const int8_t *q8_lo_ptr = &x->qs[j];
778 const int8_t *q8_hi_ptr = &x->qs[j + 32];
779 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
780 (int32_t)x->bsums[j / 16 + 1];
781 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
782 (int32_t)x->bsums[(j + 32) / 16 + 1];
783
784#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
785 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
786 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
787#elif defined(__AVX2__)
788 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
789 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
790#endif
791 for (int lane = 0; lane < active; ++lane) {
792 const uint8_t *qs = &w->qs[lane][q_offset];
793#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
794 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
795 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
796#elif defined(__AVX2__)
797 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
798 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
799 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
800 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
801 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
802#else
803 int32_t sum_lo = 0;
804 int32_t sum_hi = 0;
805 for (int l = 0; l < 32; ++l) {
806 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
807 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
808 }
809#endif
810 iacc[lane] = (int32_t)w->sc[lane][is] * sum_lo +
811 (int32_t)w->sc[lane][is + 1] * sum_hi;
812 iacc_min[lane] = (int32_t)w->m[lane][is] * bsum_lo +
813 (int32_t)w->m[lane][is + 1] * bsum_hi;
814 }
815
816#if defined(__AVX2__)
817 const __m256 acc_vec = _mm256_fmadd_ps(
818 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc)),
819 scale,
820 _mm256_loadu_ps(acc));
821 const __m256 acc_min_vec = _mm256_fmadd_ps(
822 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc_min)),
823 min_scale,
824 _mm256_loadu_ps(acc_min));
825 _mm256_storeu_ps(acc, acc_vec);
826 _mm256_storeu_ps(acc_min, acc_min_vec);
827#else
828 for (int lane = 0; lane < active; ++lane) {
829 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
830 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
831 acc[lane] = fmaf((float)iacc[lane], d, acc[lane]);
832 acc_min[lane] = fmaf((float)iacc_min[lane], dmin, acc_min[lane]);
833 }
834#endif
835 }
836}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by accum_q4_k_packed_meta_x8_q8_k_superblock_rows(), gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order(), and gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn().

◆ accum_q4_k_packed_meta_x8_q8_k_superblock_rows()

static void accum_q4_k_packed_meta_x8_q8_k_superblock_rows ( float  acc[8][8],
float  acc_min[8][8],
const block_q4_K_packed_meta_x8 *  w,
int  active,
const block_q8_K x[8],
int  rows 
)
inlinestatic

Definition at line 843 of file gemm_kernels_q4k_q8k_vnni.c.

847{
848#if defined(__AVX2__)
849 float wd[8] = {0};
850 float wdmin[8] = {0};
851 for (int lane = 0; lane < active; ++lane) {
852 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
853 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
854 }
855 const __m256 weight_scale = _mm256_loadu_ps(wd);
856 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
857
858 for (int j = 0, is = 0, q_offset = 0;
859 j < QK_K; j += 64, is += 2, q_offset += 32) {
860 int32_t iacc[8][8] = {{0}};
861 int32_t iacc_min[8][8] = {{0}};
862 __m256i q8_lo[8];
863 __m256i q8_hi[8];
864 int32_t bsum_lo[8];
865 int32_t bsum_hi[8];
866
867 for (int row = 0; row < rows; ++row) {
868 q8_lo[row] = _mm256_loadu_si256((const __m256i *)&x[row]->qs[j]);
869 q8_hi[row] = _mm256_loadu_si256((const __m256i *)&x[row]->qs[j + 32]);
870 bsum_lo[row] = (int32_t)x[row]->bsums[j / 16] +
871 (int32_t)x[row]->bsums[j / 16 + 1];
872 bsum_hi[row] = (int32_t)x[row]->bsums[(j + 32) / 16] +
873 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
874 }
875
876 for (int lane = 0; lane < active; ++lane) {
877 const __m256i packed = _mm256_loadu_si256(
878 (const __m256i *)&w->qs[lane][q_offset]);
879#if defined(CK_HAS_AVX_VNNI_256)
880 const __m256i q4_lo = q4_k_unpack_32_vnni_bytes(packed, 0);
881 const __m256i q4_hi = q4_k_unpack_32_vnni_bytes(packed, 1);
882#else
883 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
884 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
885#endif
886 const int32_t scale_lo = (int32_t)w->sc[lane][is];
887 const int32_t scale_hi = (int32_t)w->sc[lane][is + 1];
888 const int32_t min_lo = (int32_t)w->m[lane][is];
889 const int32_t min_hi = (int32_t)w->m[lane][is + 1];
890
891 for (int row = 0; row < rows; ++row) {
892#if defined(CK_HAS_AVX_VNNI_256)
893 const int32_t sum_lo =
894 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo, q8_lo[row]);
895 const int32_t sum_hi =
896 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi, q8_hi[row]);
897#else
898 const int32_t sum_lo =
899 dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo[row]);
900 const int32_t sum_hi =
901 dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi[row]);
902#endif
903 iacc[row][lane] = scale_lo * sum_lo + scale_hi * sum_hi;
904 iacc_min[row][lane] =
905 min_lo * bsum_lo[row] + min_hi * bsum_hi[row];
906 }
907 }
908
909 for (int row = 0; row < rows; ++row) {
910 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
911 const __m256 value = _mm256_fmadd_ps(
912 _mm256_cvtepi32_ps(
913 _mm256_loadu_si256((const __m256i *)iacc[row])),
914 _mm256_mul_ps(weight_scale, row_scale),
915 _mm256_loadu_ps(acc[row]));
916 const __m256 minimum = _mm256_fmadd_ps(
917 _mm256_cvtepi32_ps(
918 _mm256_loadu_si256((const __m256i *)iacc_min[row])),
919 _mm256_mul_ps(weight_min_scale, row_scale),
920 _mm256_loadu_ps(acc_min[row]));
921 _mm256_storeu_ps(acc[row], value);
922 _mm256_storeu_ps(acc_min[row], minimum);
923 }
924 }
925#else
926 for (int row = 0; row < rows; ++row) {
928 acc[row], acc_min[row], w, active, x[row]);
929 }
930#endif
931}
static void accum_q4_k_packed_meta_x8_q8_k_superblock(float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)

References accum_q4_k_packed_meta_x8_q8_k_superblock(), block_q8_K::bsums, CK_FP16_TO_FP32, and QK_K.

Referenced by gemm_q4_packed_meta_x8_split_min_4m_thread_fn(), and gemm_q4_packed_meta_x8_split_min_8m_thread_fn().

◆ accum_q4_k_packed_u8_x16_q8_k_block()

static void accum_q4_k_packed_u8_x16_q8_k_block ( float  acc[16],
const block_q4_K_packed_u8_x16 *  w,
int  active,
const block_q8_K x 
)
inlinestatic

Definition at line 1532 of file gemm_kernels_q4k_q8k_vnni.c.

1536{
1537 const float xd = x->d;
1538 for (int j = 0; j < QK_K; j += 32) {
1539 const int is = j / 32;
1540 const int8_t *q8_ptr = &x->qs[j];
1541 const int32_t bsum = (int32_t)x->bsums[j / 16] + (int32_t)x->bsums[j / 16 + 1];
1542#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1543 const __m256i q8 = _mm256_loadu_si256((const __m256i *)q8_ptr);
1544#endif
1545 for (int lane = 0; lane < active; ++lane) {
1546 const uint8_t *q4 = &w->qs[lane][j];
1547#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1548 const int32_t sumi = dot_q4_packed_u8_q8_32_vnni_q8v(q4, q8);
1549#else
1550 const int32_t sumi = dot_q4_packed_u8_q8_32_ref(q4, q8_ptr);
1551#endif
1552 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
1553 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
1554 acc[lane] += d * (float)w->sc[lane][is] * (float)sumi;
1555 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum;
1556 }
1557 }
1558}
static int32_t dot_q4_packed_u8_q8_32_ref(const uint8_t *q4_32, const int8_t *q8_32)

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, dot_q4_packed_u8_q8_32_ref(), QK_K, and block_q8_K::qs.

Referenced by gemm_q4_packed_u8_x16_mtile_thread_fn().

◆ accum_q4_k_packed_vnni_x16_q8_k_16m_superblock()

static void accum_q4_k_packed_vnni_x16_q8_k_16m_superblock ( float  acc[16][16],
float  acc_min[16][16],
const block_q4_K_packed_vnni_x16 *  w,
const block_q8_K x[16],
int  rows 
)
inlinestatic

Definition at line 1194 of file gemm_kernels_q4k_q8k_vnni.c.

1198{
1199#if defined(CK_HAS_AVX512_VNNI_512)
1200 float wd[16];
1201 float wdmin[16];
1202 for (int lane = 0; lane < 16; ++lane) {
1203 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1204 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1205 }
1206 const __m512 weight_scale = _mm512_loadu_ps(wd);
1207 const __m512 weight_min_scale = _mm512_loadu_ps(wdmin);
1208 const __m512i nibble_mask = _mm512_set1_epi8(0x0f);
1209
1210 for (int pair = 0; pair < QK_K / 64; ++pair) {
1211 const int j = pair * 64;
1212 const int is = pair * 2;
1213 const __m512i scale_lo = _mm512_cvtepu8_epi32(
1214 _mm_loadu_si128((const __m128i *)w->sc[is]));
1215 const __m512i scale_hi = _mm512_cvtepu8_epi32(
1216 _mm_loadu_si128((const __m128i *)w->sc[is + 1]));
1217 const __m512i min_lo = _mm512_cvtepu8_epi32(
1218 _mm_loadu_si128((const __m128i *)w->m[is]));
1219 const __m512i min_hi = _mm512_cvtepu8_epi32(
1220 _mm_loadu_si128((const __m128i *)w->m[is + 1]));
1221
1222 for (int row_base = 0; row_base < rows; row_base += 8) {
1223 const int row_count =
1224 row_base + 8 <= rows ? 8 : (rows - row_base);
1225 __m512i sum_lo[8];
1226 __m512i sum_hi[8];
1227 for (int row = 0; row < row_count; ++row) {
1228 sum_lo[row] = _mm512_setzero_si512();
1229 sum_hi[row] = _mm512_setzero_si512();
1230 }
1231
1232 for (int segment = 0; segment < 8; ++segment) {
1233 const __m512i packed = _mm512_loadu_si512(
1234 (const void *)(w->qs + (size_t)pair * 512u +
1235 (size_t)segment * 64u));
1236 const __m512i q4_lo =
1237 _mm512_and_si512(packed, nibble_mask);
1238 const __m512i q4_hi = _mm512_and_si512(
1239 _mm512_srli_epi16(packed, 4), nibble_mask);
1240 for (int row = 0; row < row_count; ++row) {
1241 int32_t q8_lo_word;
1242 int32_t q8_hi_word;
1243 const block_q8_K *activation = x[row_base + row];
1244 memcpy(&q8_lo_word,
1245 activation->qs + j + segment * 4,
1246 sizeof(q8_lo_word));
1247 memcpy(&q8_hi_word,
1248 activation->qs + j + 32 + segment * 4,
1249 sizeof(q8_hi_word));
1250 sum_lo[row] = _mm512_dpbusd_epi32(
1251 sum_lo[row], q4_lo,
1252 _mm512_set1_epi32(q8_lo_word));
1253 sum_hi[row] = _mm512_dpbusd_epi32(
1254 sum_hi[row], q4_hi,
1255 _mm512_set1_epi32(q8_hi_word));
1256 }
1257 }
1258
1259 for (int row = 0; row < row_count; ++row) {
1260 const block_q8_K *activation = x[row_base + row];
1261 const __m512i weighted = _mm512_add_epi32(
1262 _mm512_mullo_epi32(sum_lo[row], scale_lo),
1263 _mm512_mullo_epi32(sum_hi[row], scale_hi));
1264 const int32_t bsum_lo =
1265 (int32_t)activation->bsums[j / 16] +
1266 (int32_t)activation->bsums[j / 16 + 1];
1267 const int32_t bsum_hi =
1268 (int32_t)activation->bsums[(j + 32) / 16] +
1269 (int32_t)activation->bsums[(j + 32) / 16 + 1];
1270 const __m512i weighted_min = _mm512_add_epi32(
1271 _mm512_mullo_epi32(
1272 min_lo, _mm512_set1_epi32(bsum_lo)),
1273 _mm512_mullo_epi32(
1274 min_hi, _mm512_set1_epi32(bsum_hi)));
1275 const __m512 row_scale =
1276 _mm512_set1_ps(activation->d);
1277 const int output_row = row_base + row;
1278 const __m512 value = _mm512_fmadd_ps(
1279 _mm512_cvtepi32_ps(weighted),
1280 _mm512_mul_ps(weight_scale, row_scale),
1281 _mm512_loadu_ps(acc[output_row]));
1282 const __m512 minimum = _mm512_fmadd_ps(
1283 _mm512_cvtepi32_ps(weighted_min),
1284 _mm512_mul_ps(weight_min_scale, row_scale),
1285 _mm512_loadu_ps(acc_min[output_row]));
1286 _mm512_storeu_ps(acc[output_row], value);
1287 _mm512_storeu_ps(acc_min[output_row], minimum);
1288 }
1289 }
1290 }
1291#else
1292 (void)acc;
1293 (void)acc_min;
1294 (void)w;
1295 (void)x;
1296 (void)rows;
1297#endif
1298}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_q4_packed_vnni_x16_q8k_16m_thread_fn().

◆ accum_q4_k_packed_vnni_x16_q8_k_gemv_block()

static void accum_q4_k_packed_vnni_x16_q8_k_gemv_block ( float  acc[16],
float  acc_min[16],
const block_q4_K_packed_vnni_x16 *  w,
const block_q8_K x 
)
inlinestatic

Definition at line 1300 of file gemm_kernels_q4k_q8k_vnni.c.

1304{
1305#if defined(CK_HAS_AVX512_VNNI_512)
1306 const __m512i nibble_mask = _mm512_set1_epi8(0x0f);
1307 __m512i iacc = _mm512_setzero_si512();
1308 __m512i iacc_min = _mm512_setzero_si512();
1309 for (int pair = 0; pair < QK_K / 64; ++pair) {
1310 const int j = pair * 64;
1311 const int is = pair * 2;
1312 __m512i sum_lo = _mm512_setzero_si512();
1313 __m512i sum_hi = _mm512_setzero_si512();
1314 for (int segment = 0; segment < 8; ++segment) {
1315 const __m512i packed = _mm512_loadu_si512(
1316 (const void *)(w->qs + (size_t)pair * 512u +
1317 (size_t)segment * 64u));
1318 const __m512i q4_lo =
1319 _mm512_and_si512(packed, nibble_mask);
1320 const __m512i q4_hi = _mm512_and_si512(
1321 _mm512_srli_epi16(packed, 4), nibble_mask);
1322 int32_t q8_lo_word;
1323 int32_t q8_hi_word;
1324 memcpy(&q8_lo_word, x->qs + j + segment * 4,
1325 sizeof(q8_lo_word));
1326 memcpy(&q8_hi_word, x->qs + j + 32 + segment * 4,
1327 sizeof(q8_hi_word));
1328 sum_lo = _mm512_dpbusd_epi32(
1329 sum_lo, q4_lo, _mm512_set1_epi32(q8_lo_word));
1330 sum_hi = _mm512_dpbusd_epi32(
1331 sum_hi, q4_hi, _mm512_set1_epi32(q8_hi_word));
1332 }
1333
1334 const __m512i scale_lo = _mm512_cvtepu8_epi32(
1335 _mm_loadu_si128((const __m128i *)w->sc[is]));
1336 const __m512i scale_hi = _mm512_cvtepu8_epi32(
1337 _mm_loadu_si128((const __m128i *)w->sc[is + 1]));
1338 iacc = _mm512_add_epi32(
1339 iacc,
1340 _mm512_add_epi32(
1341 _mm512_mullo_epi32(sum_lo, scale_lo),
1342 _mm512_mullo_epi32(sum_hi, scale_hi)));
1343
1344 const int32_t bsum_lo =
1345 (int32_t)x->bsums[j / 16] +
1346 (int32_t)x->bsums[j / 16 + 1];
1347 const int32_t bsum_hi =
1348 (int32_t)x->bsums[(j + 32) / 16] +
1349 (int32_t)x->bsums[(j + 32) / 16 + 1];
1350 const __m512i min_lo = _mm512_cvtepu8_epi32(
1351 _mm_loadu_si128((const __m128i *)w->m[is]));
1352 const __m512i min_hi = _mm512_cvtepu8_epi32(
1353 _mm_loadu_si128((const __m128i *)w->m[is + 1]));
1354 iacc_min = _mm512_add_epi32(
1355 iacc_min,
1356 _mm512_add_epi32(
1357 _mm512_mullo_epi32(
1358 min_lo, _mm512_set1_epi32(bsum_lo)),
1359 _mm512_mullo_epi32(
1360 min_hi, _mm512_set1_epi32(bsum_hi))));
1361 }
1362
1363 float wd[16];
1364 float wdmin[16];
1365 for (int lane = 0; lane < 16; ++lane) {
1366 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1367 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1368 }
1369 const __m512 xd = _mm512_set1_ps(x->d);
1370 _mm512_storeu_ps(
1371 acc,
1372 _mm512_fmadd_ps(
1373 _mm512_cvtepi32_ps(iacc),
1374 _mm512_mul_ps(_mm512_loadu_ps(wd), xd),
1375 _mm512_loadu_ps(acc)));
1376 _mm512_storeu_ps(
1377 acc_min,
1378 _mm512_fmadd_ps(
1379 _mm512_cvtepi32_ps(iacc_min),
1380 _mm512_mul_ps(_mm512_loadu_ps(wdmin), xd),
1381 _mm512_loadu_ps(acc_min)));
1382#else
1383 (void)acc;
1384 (void)acc_min;
1385 (void)w;
1386 (void)x;
1387#endif
1388}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order().

◆ accum_q4_k_packed_vnni_x8_q8_k_4m_superblock()

static void accum_q4_k_packed_vnni_x8_q8_k_4m_superblock ( float  acc[4][8],
float  acc_min[4][8],
const block_q4_K_packed_vnni_x8 *  w,
const block_q8_K x[4],
int  rows 
)
inlinestatic

Definition at line 937 of file gemm_kernels_q4k_q8k_vnni.c.

941{
942#if defined(CK_HAS_AVX_VNNI_256)
943 float wd[8];
944 float wdmin[8];
945 for (int lane = 0; lane < 8; ++lane) {
946 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
947 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
948 }
949 const __m256 weight_scale = _mm256_loadu_ps(wd);
950 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
951 const __m256i nibble_mask = _mm256_set1_epi8(0x0f);
952
953 for (int pair = 0; pair < QK_K / 64; ++pair) {
954 const int j = pair * 64;
955 const int is = pair * 2;
956 __m256i sum_lo[4];
957 __m256i sum_hi[4];
958 for (int row = 0; row < rows; ++row) {
959 sum_lo[row] = _mm256_setzero_si256();
960 sum_hi[row] = _mm256_setzero_si256();
961 }
962
963 for (int segment = 0; segment < 8; ++segment) {
964 const __m256i packed = _mm256_loadu_si256(
965 (const __m256i *)(w->qs + (size_t)pair * 256u +
966 (size_t)segment * 32u));
967 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
968 const __m256i q4_hi = _mm256_and_si256(
969 _mm256_srli_epi16(packed, 4), nibble_mask);
970 for (int row = 0; row < rows; ++row) {
971 int32_t q8_lo_word;
972 int32_t q8_hi_word;
973 memcpy(&q8_lo_word, x[row]->qs + j + segment * 4,
974 sizeof(q8_lo_word));
975 memcpy(&q8_hi_word, x[row]->qs + j + 32 + segment * 4,
976 sizeof(q8_hi_word));
977 sum_lo[row] = ck_dpbusd_i32x8(
978 sum_lo[row], q4_lo, _mm256_set1_epi32(q8_lo_word));
979 sum_hi[row] = ck_dpbusd_i32x8(
980 sum_hi[row], q4_hi, _mm256_set1_epi32(q8_hi_word));
981 }
982 }
983
984 const __m256i scale_lo = _mm256_cvtepu8_epi32(
985 _mm_loadl_epi64((const __m128i *)w->sc[is]));
986 const __m256i scale_hi = _mm256_cvtepu8_epi32(
987 _mm_loadl_epi64((const __m128i *)w->sc[is + 1]));
988 const __m256i min_lo = _mm256_cvtepu8_epi32(
989 _mm_loadl_epi64((const __m128i *)w->m[is]));
990 const __m256i min_hi = _mm256_cvtepu8_epi32(
991 _mm_loadl_epi64((const __m128i *)w->m[is + 1]));
992
993 for (int row = 0; row < rows; ++row) {
994 const __m256i weighted = _mm256_add_epi32(
995 _mm256_mullo_epi32(sum_lo[row], scale_lo),
996 _mm256_mullo_epi32(sum_hi[row], scale_hi));
997 const int32_t bsum_lo =
998 (int32_t)x[row]->bsums[j / 16] +
999 (int32_t)x[row]->bsums[j / 16 + 1];
1000 const int32_t bsum_hi =
1001 (int32_t)x[row]->bsums[(j + 32) / 16] +
1002 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
1003 const __m256i weighted_min = _mm256_add_epi32(
1004 _mm256_mullo_epi32(min_lo, _mm256_set1_epi32(bsum_lo)),
1005 _mm256_mullo_epi32(min_hi, _mm256_set1_epi32(bsum_hi)));
1006 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
1007 const __m256 value = _mm256_fmadd_ps(
1008 _mm256_cvtepi32_ps(weighted),
1009 _mm256_mul_ps(weight_scale, row_scale),
1010 _mm256_loadu_ps(acc[row]));
1011 const __m256 minimum = _mm256_fmadd_ps(
1012 _mm256_cvtepi32_ps(weighted_min),
1013 _mm256_mul_ps(weight_min_scale, row_scale),
1014 _mm256_loadu_ps(acc_min[row]));
1015 _mm256_storeu_ps(acc[row], value);
1016 _mm256_storeu_ps(acc_min[row], minimum);
1017 }
1018 }
1019#else
1020 (void)acc;
1021 (void)acc_min;
1022 (void)w;
1023 (void)x;
1024 (void)rows;
1025#endif
1026}

References block_q8_K::bsums, CK_FP16_TO_FP32, and QK_K.

Referenced by gemm_q4_packed_vnni_x8_q8k_4m_job().

◆ ck_q4k_packed_vnni_x16_available()

int ck_q4k_packed_vnni_x16_available ( void  )

Definition at line 286 of file gemm_kernels_q4k_q8k_vnni.c.

287{
288#if defined(CK_HAS_AVX512_VNNI_512)
289 return 1;
290#else
291 return 0;
292#endif
293}

Referenced by gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order(), and gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m().

◆ ck_q4k_packed_vnni_x8_available()

int ck_q4k_packed_vnni_x8_available ( void  )

Definition at line 272 of file gemm_kernels_q4k_q8k_vnni.c.

273{
274#if defined(CK_HAS_AVX_VNNI_256)
275 return 1;
276#else
277 return 0;
278#endif
279}

◆ ck_q4k_packed_vnni_x8_compact_order_available()

int ck_q4k_packed_vnni_x8_compact_order_available ( void  )

Definition at line 1178 of file gemm_kernels_q4k_q8k_vnni.c.

1179{
1180#if defined(CK_HAS_AVX_VNNI_256) && defined(__AVX512F__) && \
1181 defined(__AVX512VNNI__) && defined(__AVX512VL__)
1182 return 1;
1183#else
1184 return 0;
1185#endif
1186}

◆ ck_q4k_silu_f32()

static float ck_q4k_silu_f32 ( float  x)
inlinestatic

Definition at line 3671 of file gemm_kernels_q4k_q8k_vnni.c.

3672{
3673 return x / (1.0f + expf(-x));
3674}

Referenced by gemm_q4_gateup_swiglu_thread_fn().

◆ ck_q4k_x16_chunk4_enabled()

static int ck_q4k_x16_chunk4_enabled ( void  )
static

Definition at line 36 of file gemm_kernels_q4k_q8k_vnni.c.

37{
38 static int cached = -1;
39 if (cached < 0) {
40 const char *env = getenv("CK_Q4K_X16_CHUNK4");
41 cached = env ? ck_env_value_truthy(env) : 1;
42 }
43 return cached;
44}
static int ck_env_value_truthy(const char *v)

References ck_env_value_truthy().

Referenced by gemm_q4_gateup_swiglu_x16_thread_fn(), and gemm_q4_packed_meta_x16_mreuse_process_job().

◆ dot_q4_k_packed_meta_q8_k_block()

static float dot_q4_k_packed_meta_q8_k_block ( const block_q4_K_packed_meta *  w,
const block_q8_K x 
)
inlinestatic

Definition at line 668 of file gemm_kernels_q4k_q8k_vnni.c.

670{
671 const float d = CK_FP16_TO_FP32(w->d) * x->d;
672 const float dmin = CK_FP16_TO_FP32(w->dmin) * x->d;
673 float sumf = 0.0f;
674 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
675 const uint8_t *qs = &w->qs[q_offset];
676 const int8_t *q8_lo = &x->qs[j];
677 const int8_t *q8_hi = &x->qs[j + 32];
678
679#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
680 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni(qs, q8_lo, 0);
681 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni(qs, q8_hi, 1);
682#else
683 int32_t sum_lo = 0;
684 int32_t sum_hi = 0;
685 for (int l = 0; l < 32; ++l) {
686 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo[l];
687 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi[l];
688 }
689#endif
690 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
691 (int32_t)x->bsums[j / 16 + 1];
692 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
693 (int32_t)x->bsums[(j + 32) / 16 + 1];
694
695 sumf += d * (float)w->sc[is] * (float)sum_lo;
696 sumf -= dmin * (float)w->m[is] * (float)bsum_lo;
697 sumf += d * (float)w->sc[is + 1] * (float)sum_hi;
698 sumf -= dmin * (float)w->m[is + 1] * (float)bsum_hi;
699 }
700 return sumf;
701}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_nt_q4_k_packed_meta_q8_k(), gemm_nt_q4_k_packed_meta_q8_k_tile(), gemm_q4_packed_meta_nsplit_thread_fn(), and gemm_q4_packed_meta_thread_fn().

◆ dot_q4_k_packed_u8_q8_k_block()

static float dot_q4_k_packed_u8_q8_k_block ( const block_q4_K_packed_u8 *  w,
const block_q8_K x 
)
inlinestatic

Definition at line 648 of file gemm_kernels_q4k_q8k_vnni.c.

650{
651 const float d = CK_FP16_TO_FP32(w->d) * x->d;
652 const float dmin = CK_FP16_TO_FP32(w->dmin) * x->d;
653 float sumf = 0.0f;
654 for (int j = 0; j < QK_K; j += 32) {
655 const int is = j / 32;
656#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
657 const int32_t sumi = dot_q4_packed_u8_q8_32_vnni(&w->qs[j], &x->qs[j]);
658#else
659 const int32_t sumi = dot_q4_packed_u8_q8_32_ref(&w->qs[j], &x->qs[j]);
660#endif
661 const int32_t bsum = (int32_t)x->bsums[j / 16] + (int32_t)x->bsums[j / 16 + 1];
662 sumf += d * (float)w->sc[is] * (float)sumi;
663 sumf -= dmin * (float)w->m[is] * (float)bsum;
664 }
665 return sumf;
666}

References block_q8_K::bsums, CK_FP16_TO_FP32, block_q8_K::d, dot_q4_packed_u8_q8_32_ref(), QK_K, and block_q8_K::qs.

Referenced by gemm_nt_q4_k_packed_u8_q8_k().

◆ dot_q4_k_packed_vnni_x8_q8_k_compact_order()

static void dot_q4_k_packed_vnni_x8_q8_k_compact_order ( float  block_sums[4][8],
const block_q4_K_packed_vnni_x8 *  w,
const block_q8_K x[4],
int  rows 
)
inlinestatic

Definition at line 1032 of file gemm_kernels_q4k_q8k_vnni.c.

1037{
1038#if defined(CK_HAS_AVX_VNNI_256)
1039 float wd[8];
1040 float wdmin[8];
1041 for (int lane = 0; lane < 8; ++lane) {
1042 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1043 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1044 }
1045 const __m256 weight_scale = _mm256_loadu_ps(wd);
1046 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
1047 const __m256i nibble_mask = _mm256_set1_epi8(0x0f);
1048
1049 for (int pair = 0; pair < QK_K / 64; ++pair) {
1050 const int j = pair * 64;
1051 const int is = pair * 2;
1052 __m256i sum_lo[4];
1053 __m256i sum_hi[4];
1054 for (int row = 0; row < rows; ++row) {
1055 sum_lo[row] = _mm256_setzero_si256();
1056 sum_hi[row] = _mm256_setzero_si256();
1057 }
1058
1059 for (int segment = 0; segment < 8; ++segment) {
1060 const __m256i packed = _mm256_loadu_si256(
1061 (const __m256i *)(w->qs + (size_t)pair * 256u +
1062 (size_t)segment * 32u));
1063 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
1064 const __m256i q4_hi = _mm256_and_si256(
1065 _mm256_srli_epi16(packed, 4), nibble_mask);
1066 for (int row = 0; row < rows; ++row) {
1067 int32_t q8_lo_word;
1068 int32_t q8_hi_word;
1069 memcpy(&q8_lo_word, x[row]->qs + j + segment * 4,
1070 sizeof(q8_lo_word));
1071 memcpy(&q8_hi_word, x[row]->qs + j + 32 + segment * 4,
1072 sizeof(q8_hi_word));
1073 sum_lo[row] = ck_dpbusd_i32x8(
1074 sum_lo[row], q4_lo, _mm256_set1_epi32(q8_lo_word));
1075 sum_hi[row] = ck_dpbusd_i32x8(
1076 sum_hi[row], q4_hi, _mm256_set1_epi32(q8_hi_word));
1077 }
1078 }
1079
1080 const __m256 scale_lo = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1081 _mm_loadl_epi64((const __m128i *)w->sc[is])));
1082 const __m256 scale_hi = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1083 _mm_loadl_epi64((const __m128i *)w->sc[is + 1])));
1084 const __m256 min_lo = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1085 _mm_loadl_epi64((const __m128i *)w->m[is])));
1086 const __m256 min_hi = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1087 _mm_loadl_epi64((const __m128i *)w->m[is + 1])));
1088
1089 for (int row = 0; row < rows; ++row) {
1090 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
1091 const __m256 d = _mm256_mul_ps(weight_scale, row_scale);
1092 const __m256 dmin = _mm256_mul_ps(weight_min_scale, row_scale);
1093 __m256 value = _mm256_loadu_ps(block_sums[row]);
1094 value = _mm256_fmadd_ps(
1095 _mm256_mul_ps(d, scale_lo),
1096 _mm256_cvtepi32_ps(sum_lo[row]), value);
1097 const int32_t bsum_lo =
1098 (int32_t)x[row]->bsums[j / 16] +
1099 (int32_t)x[row]->bsums[j / 16 + 1];
1100 value = _mm256_fnmadd_ps(
1101 _mm256_mul_ps(dmin, min_lo),
1102 _mm256_set1_ps((float)bsum_lo), value);
1103 value = _mm256_fmadd_ps(
1104 _mm256_mul_ps(d, scale_hi),
1105 _mm256_cvtepi32_ps(sum_hi[row]), value);
1106 const int32_t bsum_hi =
1107 (int32_t)x[row]->bsums[(j + 32) / 16] +
1108 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
1109 value = _mm256_fnmadd_ps(
1110 _mm256_mul_ps(dmin, min_hi),
1111 _mm256_set1_ps((float)bsum_hi), value);
1112 _mm256_storeu_ps(block_sums[row], value);
1113 }
1114 }
1115#else
1116 (void)block_sums;
1117 (void)w;
1118 (void)x;
1119 (void)rows;
1120#endif
1121}

References block_q8_K::bsums, CK_FP16_TO_FP32, and QK_K.

Referenced by gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4().

◆ dot_q4_packed_u8_q8_32_ref()

static int32_t dot_q4_packed_u8_q8_32_ref ( const uint8_t *  q4_32,
const int8_t *  q8_32 
)
inlinestatic

Definition at line 637 of file gemm_kernels_q4k_q8k_vnni.c.

639{
640 int32_t acc = 0;
641 for (int i = 0; i < 32; ++i) {
642 acc += (int32_t)q4_32[i] * (int32_t)q8_32[i];
643 }
644 return acc;
645}

Referenced by accum_q4_k_packed_u8_x16_q8_k_block(), and dot_q4_k_packed_u8_q8_k_block().

◆ gemm_nt_q4_k_packed_meta_q8_k()

void gemm_nt_q4_k_packed_meta_q8_k ( const void *  A_q8,
const void *  B_packed,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1798 of file gemm_kernels_q4k_q8k_vnni.c.

1803{
1804 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1805 return;
1806 }
1807 const block_q8_K *A = (const block_q8_K *)A_q8;
1808 const block_q4_K_packed_meta *W = (const block_q4_K_packed_meta *)B_packed;
1809 const int blocks_per_vec = K / QK_K;
1810 const int blocks_per_row = K / QK_K;
1811 for (int m = 0; m < M; ++m) {
1812 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1813 float *c_row = C + (size_t)m * (size_t)N;
1814 for (int n = 0; n < N; ++n) {
1815 const block_q4_K_packed_meta *w_row = W + (size_t)n * (size_t)blocks_per_row;
1816 float sum = bias ? bias[n] : 0.0f;
1817 for (int b = 0; b < blocks_per_row; ++b) {
1818 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
1819 }
1820 c_row[n] = sum;
1821 }
1822 }
1823}
static float dot_q4_k_packed_meta_q8_k_block(const block_q4_K_packed_meta *w, const block_q8_K *x)
#define C(color)
Definition show_config.c:39

References C, dot_q4_k_packed_meta_q8_k_block(), and QK_K.

Referenced by gemm_nt_q4_k_packed_meta_q8_k_threaded(), and gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit().

◆ gemm_nt_q4_k_packed_meta_q8_k_threaded()

void gemm_nt_q4_k_packed_meta_q8_k_threaded ( const void *  A_q8,
const void *  B_packed,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 2247 of file gemm_kernels_q4k_q8k_vnni.c.

2253{
2254 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2255 return;
2256 }
2257 ck_threadpool_t *pool = ck_threadpool_global();
2258 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2259 if (active_threads <= 0 || active_threads > pool_threads) {
2260 active_threads = pool_threads;
2261 }
2262 if (active_threads > M) {
2263 active_threads = M;
2264 }
2265 if (active_threads <= 1) {
2266 gemm_nt_q4_k_packed_meta_q8_k(A_q8, B_packed, bias, C, M, N, K);
2267 return;
2268 }
2269 gemm_q4_packed_meta_thread_work_t work = {
2270 .A = (const block_q8_K *)A_q8,
2271 .W = (const block_q4_K_packed_meta *)B_packed,
2272 .bias = bias,
2273 .C = C,
2274 .M = M,
2275 .N = N,
2276 .K = K,
2277 .blocks_per_vec = K / QK_K,
2278 .blocks_per_row = K / QK_K,
2279 };
2280 ck_threadpool_dispatch_n(pool, active_threads, gemm_q4_packed_meta_thread_fn, &work);
2281}
void ck_threadpool_dispatch_n(ck_threadpool_t *pool, int active_threads, ck_work_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
void gemm_nt_q4_k_packed_meta_q8_k(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K)
static void gemm_q4_packed_meta_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_nt_q4_k_packed_meta_q8_k(), gemm_q4_packed_meta_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit()

void gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit ( const void *  A_q8,
const void *  B_packed,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 2310 of file gemm_kernels_q4k_q8k_vnni.c.

2316{
2317 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2318 return;
2319 }
2320 ck_threadpool_t *pool = ck_threadpool_global();
2321 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2322 if (active_threads <= 0 || active_threads > pool_threads) {
2323 active_threads = pool_threads;
2324 }
2325 if (active_threads > N) {
2326 active_threads = N;
2327 }
2328 if (active_threads <= 1) {
2329 gemm_nt_q4_k_packed_meta_q8_k(A_q8, B_packed, bias, C, M, N, K);
2330 return;
2331 }
2332 gemm_q4_packed_meta_thread_work_t work = {
2333 .A = (const block_q8_K *)A_q8,
2334 .W = (const block_q4_K_packed_meta *)B_packed,
2335 .bias = bias,
2336 .C = C,
2337 .M = M,
2338 .N = N,
2339 .K = K,
2340 .blocks_per_vec = K / QK_K,
2341 .blocks_per_row = K / QK_K,
2342 };
2344}
static void gemm_q4_packed_meta_nsplit_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_nt_q4_k_packed_meta_q8_k(), gemm_q4_packed_meta_nsplit_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_q8_k_tile()

void gemm_nt_q4_k_packed_meta_q8_k_tile ( const void *  A_q8,
const void *  B_packed,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  m0,
int  m1,
int  n0,
int  n1 
)

Definition at line 1825 of file gemm_kernels_q4k_q8k_vnni.c.

1831{
1832 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1833 return;
1834 }
1835 if (m0 < 0) m0 = 0;
1836 if (n0 < 0) n0 = 0;
1837 if (m1 > M) m1 = M;
1838 if (n1 > N) n1 = N;
1839 if (m0 >= m1 || n0 >= n1) {
1840 return;
1841 }
1842
1843 const block_q8_K *A = (const block_q8_K *)A_q8;
1844 const block_q4_K_packed_meta *W = (const block_q4_K_packed_meta *)B_packed;
1845 const int blocks_per_vec = K / QK_K;
1846 const int blocks_per_row = K / QK_K;
1847
1848 for (int m = m0; m < m1; ++m) {
1849 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1850 float *c_row = C + (size_t)m * (size_t)N;
1851 for (int n = n0; n < n1; ++n) {
1852 const block_q4_K_packed_meta *w_row = W + (size_t)n * (size_t)blocks_per_row;
1853 float sum = bias ? bias[n] : 0.0f;
1854 for (int b = 0; b < blocks_per_row; ++b) {
1855 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
1856 }
1857 c_row[n] = sum;
1858 }
1859 }
1860}

References C, dot_q4_k_packed_meta_q8_k_block(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni()

void gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni ( const void *  A_q8,
const void *  B_packed_x16,
const float *  bias,
float *  C,
int  M,
int  D,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 3401 of file gemm_kernels_q4k_q8k_vnni.c.

3408{
3409 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || D <= 0 || K <= 0 || (K % QK_K) != 0) {
3410 return;
3411 }
3412 ck_threadpool_t *pool = ck_threadpool_global();
3413 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3414 int tm = tile_m > 0 ? tile_m : 4;
3415 if (tm > 8) tm = 8;
3416 const int groups_d = (D + 15) / 16;
3417 const int mt = (M + tm - 1) / tm;
3418 const int jobs = mt * groups_d;
3419 if (active_threads <= 0 || active_threads > pool_threads) {
3420 active_threads = pool_threads;
3421 }
3422 if (active_threads > jobs) {
3423 active_threads = jobs;
3424 }
3425 if (active_threads < 1) {
3426 active_threads = 1;
3427 }
3428
3429 gemm_q4_gateup_swiglu_x16_work_t work = {
3430 .A = (const block_q8_K *)A_q8,
3431 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3432 .bias = bias,
3433 .C = C,
3434 .M = M,
3435 .D = D,
3436 .K = K,
3437 .tile_m = tm,
3438 .blocks_per_vec = K / QK_K,
3439 .blocks_per_row = K / QK_K,
3440 .groups_d = groups_d,
3441 .jobs = jobs,
3442 };
3443 if (!pool || active_threads <= 1) {
3445 return;
3446 }
3448}
static void gemm_q4_gateup_swiglu_x16_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_gateup_swiglu_x16_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order()

void gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1944 of file gemm_kernels_q4k_q8k_vnni.c.

1947{
1948#if defined(__AVX512F__) && defined(__AVX512BW__) && defined(__AVX512DQ__)
1949 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
1950 (K % QK_K) != 0 || (N % 16) != 0) {
1951 return;
1952 }
1953 const block_q8_K *A = (const block_q8_K *)A_q8;
1954 const block_q4_K_packed_meta_x8 *W =
1955 (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1956 const int blocks_per_row = K / QK_K;
1957
1958 for (int row = 0; row < M; ++row) {
1959 const block_q8_K *a_row = A + (size_t)row * (size_t)blocks_per_row;
1960 float *c_row = C + (size_t)row * (size_t)N;
1961 for (int n0 = 0; n0 < N; n0 += 16) {
1962 const int group0 = n0 / 8;
1963 __m512 acc = _mm512_setzero_ps();
1964 __m512 acc_min = _mm512_setzero_ps();
1965
1966 for (int b = 0; b < blocks_per_row; ++b) {
1967 const block_q4_K_packed_meta_x8 *w0 =
1968 W + (size_t)group0 * (size_t)blocks_per_row + (size_t)b;
1969 const block_q4_K_packed_meta_x8 *w1 =
1970 W + (size_t)(group0 + 1) * (size_t)blocks_per_row + (size_t)b;
1971 const block_q8_K *x = &a_row[b];
1972 float d[16];
1973 float dmin[16];
1974 for (int lane = 0; lane < 16; ++lane) {
1975 const block_q4_K_packed_meta_x8 *w = lane < 8 ? w0 : w1;
1976 const int wl = lane & 7;
1977 d[lane] = CK_FP16_TO_FP32(w->d[wl]);
1978 dmin[lane] = CK_FP16_TO_FP32(w->dmin[wl]);
1979 }
1980 const __m512 scale =
1981 _mm512_mul_ps(_mm512_loadu_ps(d), _mm512_set1_ps(x->d));
1982 const __m512 min_scale =
1983 _mm512_mul_ps(_mm512_loadu_ps(dmin), _mm512_set1_ps(x->d));
1984
1985 for (int j = 0, is = 0, q_offset = 0;
1986 j < QK_K;
1987 j += 64, is += 2, q_offset += 32) {
1988 int32_t iacc[16];
1989 int32_t iacc_min[16];
1990 const int8_t *q8_lo_ptr = &x->qs[j];
1991 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1992 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1993 (int32_t)x->bsums[j / 16 + 1];
1994 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1995 (int32_t)x->bsums[(j + 32) / 16 + 1];
1996#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1997 const __m256i q8_lo =
1998 _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1999 const __m256i q8_hi =
2000 _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
2001#endif
2002 for (int lane = 0; lane < 16; ++lane) {
2003 const block_q4_K_packed_meta_x8 *w = lane < 8 ? w0 : w1;
2004 const int wl = lane & 7;
2005 const uint8_t *qs = &w->qs[wl][q_offset];
2006#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
2007 const int32_t sum_lo =
2008 dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
2009 const int32_t sum_hi =
2010 dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
2011#else
2012 int32_t sum_lo = 0;
2013 int32_t sum_hi = 0;
2014 for (int i = 0; i < 32; ++i) {
2015 sum_lo += (int32_t)(qs[i] & 0x0F) *
2016 (int32_t)q8_lo_ptr[i];
2017 sum_hi += (int32_t)(qs[i] >> 4) *
2018 (int32_t)q8_hi_ptr[i];
2019 }
2020#endif
2021 /* llama.cpp's AVX-512 q4_K_8x8 provider keeps each
2022 * 32-value dot in an int16 lane through PMADDUBSW and
2023 * wrapping VPADDW operations, then widens while
2024 * applying the two sub-block scales. */
2025 const int16_t packed_sum_lo = (int16_t)sum_lo;
2026 const int16_t packed_sum_hi = (int16_t)sum_hi;
2027 iacc[lane] = (int32_t)w->sc[wl][is] * (int32_t)packed_sum_lo +
2028 (int32_t)w->sc[wl][is + 1] * (int32_t)packed_sum_hi;
2029 iacc_min[lane] = (int32_t)w->m[wl][is] * bsum_lo +
2030 (int32_t)w->m[wl][is + 1] * bsum_hi;
2031 }
2032 acc = _mm512_fmadd_ps(
2033 _mm512_cvtepi32_ps(_mm512_loadu_si512(iacc)),
2034 scale,
2035 acc);
2036 acc_min = _mm512_fmadd_ps(
2037 _mm512_cvtepi32_ps(_mm512_loadu_si512(iacc_min)),
2038 min_scale,
2039 acc_min);
2040 }
2041 }
2042 __m512 value = _mm512_sub_ps(acc, acc_min);
2043 if (bias) {
2044 value = _mm512_add_ps(value, _mm512_loadu_ps(bias + n0));
2045 }
2046 _mm512_storeu_ps(c_row + n0, value);
2047 }
2048 }
2049#else
2051 A_q8, B_packed_x8, bias, C, M, N, K);
2052#endif
2053}
void gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)

References block_q8_K::bsums, C, CK_FP16_TO_FP32, block_q8_K::d, gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order(), QK_K, and block_q8_K::qs.

◆ gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse()

void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse ( const void *  A_q8,
const void *  B_packed_x16,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 3228 of file gemm_kernels_q4k_q8k_vnni.c.

3235{
3236 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
3237 return;
3238 }
3239 ck_threadpool_t *pool = ck_threadpool_global();
3240 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3241 const int groups = (N + 15) / 16;
3242 int tm = tile_m > 0 ? tile_m : 4;
3243 if (tm > 8) tm = 8;
3244 const int mt = (M + tm - 1) / tm;
3245 const int jobs = mt * groups;
3246 if (active_threads <= 0 || active_threads > pool_threads) {
3247 active_threads = pool_threads;
3248 }
3249 if (active_threads > jobs) {
3250 active_threads = jobs;
3251 }
3252 gemm_q4_packed_meta_x16_thread_work_t work = {
3253 .A = (const block_q8_K *)A_q8,
3254 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3255 .bias = bias,
3256 .C = C,
3257 .M = M,
3258 .N = N,
3259 .K = K,
3260 .blocks_per_vec = K / QK_K,
3261 .blocks_per_row = K / QK_K,
3262 .groups = groups,
3263 .tile_m = tm,
3264 .jobs = jobs,
3265 };
3266 if (active_threads <= 1) {
3268 return;
3269 }
3271}
static void gemm_q4_packed_meta_x16_mreuse_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x16_mreuse_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile()

void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile ( const void *  A_q8,
const void *  B_packed_x16,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 3273 of file gemm_kernels_q4k_q8k_vnni.c.

3280{
3281 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
3282 return;
3283 }
3284 ck_threadpool_t *pool = ck_threadpool_global();
3285 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3286 const int groups = (N + 15) / 16;
3287 int tm = tile_m > 0 ? tile_m : 4;
3288 if (tm > 8) tm = 8;
3289 const int mt = (M + tm - 1) / tm;
3290 const int jobs = mt * groups;
3291 if (active_threads <= 0 || active_threads > pool_threads) {
3292 active_threads = pool_threads;
3293 }
3294 if (active_threads > jobs) {
3295 active_threads = jobs;
3296 }
3297 gemm_q4_packed_meta_x16_thread_work_t work = {
3298 .A = (const block_q8_K *)A_q8,
3299 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3300 .bias = bias,
3301 .C = C,
3302 .M = M,
3303 .N = N,
3304 .K = K,
3305 .blocks_per_vec = K / QK_K,
3306 .blocks_per_row = K / QK_K,
3307 .groups = groups,
3308 .tile_m = tm,
3309 .jobs = jobs,
3310 };
3311 if (active_threads <= 1) {
3313 return;
3314 }
3316}
static void gemm_q4_packed_meta_x16_mtile_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x16_mtile_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k()

void gemm_nt_q4_k_packed_meta_x8_q8_k ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1862 of file gemm_kernels_q4k_q8k_vnni.c.

1867{
1868 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1869 return;
1870 }
1871 const block_q8_K *A = (const block_q8_K *)A_q8;
1872 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1873 const int blocks_per_vec = K / QK_K;
1874 const int blocks_per_row = K / QK_K;
1875 const int groups = (N + 7) / 8;
1876
1877 for (int m = 0; m < M; ++m) {
1878 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1879 float *c_row = C + (size_t)m * (size_t)N;
1880 for (int g = 0; g < groups; ++g) {
1881 const int n0 = g * 8;
1882 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
1883 float acc[8];
1884 for (int lane = 0; lane < active; ++lane) {
1885 acc[lane] = bias ? bias[n0 + lane] : 0.0f;
1886 }
1887 for (int b = 0; b < blocks_per_row; ++b) {
1888 const block_q4_K_packed_meta_x8 *w_group =
1889 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
1890 accum_q4_k_packed_meta_x8_q8_k_block(acc, w_group, active, &a_row[b]);
1891 }
1892 for (int lane = 0; lane < active; ++lane) {
1893 c_row[n0 + lane] = acc[lane];
1894 }
1895 }
1896 }
1897}

References accum_q4_k_packed_meta_x8_q8_k_block(), C, and QK_K.

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile(), and gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit().

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order()

void gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 2055 of file gemm_kernels_q4k_q8k_vnni.c.

2058{
2059 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2060 return;
2061 }
2062 const block_q8_K *A = (const block_q8_K *)A_q8;
2063 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
2064 const int blocks_per_row = K / QK_K;
2065 const int groups = (N + 7) / 8;
2066 for (int m = 0; m < M; ++m) {
2067 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_row;
2068 float *c_row = C + (size_t)m * (size_t)N;
2069 for (int g = 0; g < groups; ++g) {
2070 const int n0 = g * 8;
2071 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
2072 float acc[8] = {0};
2073 float acc_min[8] = {0};
2074 for (int b = 0; b < blocks_per_row; ++b) {
2075 const block_q4_K_packed_meta_x8 *w_group =
2076 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
2078 acc, acc_min, w_group, active, &a_row[b]);
2079 }
2080 float values[8];
2081#if defined(__AVX2__)
2082 _mm256_storeu_ps(values, _mm256_sub_ps(
2083 _mm256_loadu_ps(acc), _mm256_loadu_ps(acc_min)));
2084#else
2085 for (int lane = 0; lane < active; ++lane) values[lane] = acc[lane] - acc_min[lane];
2086#endif
2087 for (int lane = 0; lane < active; ++lane) {
2088 c_row[n0 + lane] = values[lane] + (bias ? bias[n0 + lane] : 0.0f);
2089 }
2090 }
2091 }
2092}
static void accum_q4_k_packed_meta_x8_q8_k_gemv_block(float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)

References accum_q4_k_packed_meta_x8_q8_k_gemv_block(), C, and QK_K.

Referenced by ck_moe_q4k_llama_projection().

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m()

void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 2947 of file gemm_kernels_q4k_q8k_vnni.c.

2950{
2951 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2952 (K % QK_K) != 0) {
2953 return;
2954 }
2955 ck_threadpool_t *pool = ck_threadpool_global();
2956 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2957 const int groups = (N + 7) / 8;
2958 const int jobs = ((M + 3) / 4) * groups;
2959 if (active_threads <= 0 || active_threads > pool_threads) {
2960 active_threads = pool_threads;
2961 }
2962 if (active_threads > jobs) active_threads = jobs;
2963
2964 gemm_q4_packed_meta_x8_thread_work_t work = {
2965 .A = (const block_q8_K *)A_q8,
2966 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2967 .bias = bias,
2968 .C = C,
2969 .M = M,
2970 .N = N,
2971 .K = K,
2972 .blocks_per_vec = K / QK_K,
2973 .blocks_per_row = K / QK_K,
2974 .groups = groups,
2975 .tile_m = 4,
2976 .jobs = jobs,
2977 };
2978 if (active_threads <= 1 || !pool) {
2980 return;
2981 }
2983 pool, active_threads,
2985}
static void gemm_q4_packed_meta_x8_split_min_4m_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x8_split_min_4m_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m()

void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 2987 of file gemm_kernels_q4k_q8k_vnni.c.

2990{
2991 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2992 (K % QK_K) != 0) {
2993 return;
2994 }
2995 ck_threadpool_t *pool = ck_threadpool_global();
2996 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2997 const int groups = (N + 7) / 8;
2998 const int jobs = ((M + 7) / 8) * groups;
2999 if (active_threads <= 0 || active_threads > pool_threads) {
3000 active_threads = pool_threads;
3001 }
3002 if (active_threads > jobs) active_threads = jobs;
3003
3004 gemm_q4_packed_meta_x8_thread_work_t work = {
3005 .A = (const block_q8_K *)A_q8,
3006 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
3007 .bias = bias,
3008 .C = C,
3009 .M = M,
3010 .N = N,
3011 .K = K,
3012 .blocks_per_vec = K / QK_K,
3013 .blocks_per_row = K / QK_K,
3014 .groups = groups,
3015 .tile_m = 8,
3016 .jobs = jobs,
3017 };
3018 if (active_threads <= 1 || !pool) {
3020 return;
3021 }
3023 pool, active_threads,
3025}
static void gemm_q4_packed_meta_x8_split_min_8m_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x8_split_min_8m_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse()

void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 2905 of file gemm_kernels_q4k_q8k_vnni.c.

2908{
2909 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2910 (K % QK_K) != 0) {
2911 return;
2912 }
2913 ck_threadpool_t *pool = ck_threadpool_global();
2914 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2915 const int groups = (N + 7) / 8;
2916 int tm = tile_m > 0 ? tile_m : 4;
2917 if (tm > 8) tm = 8;
2918 const int jobs = ((M + tm - 1) / tm) * groups;
2919 if (active_threads <= 0 || active_threads > pool_threads) {
2920 active_threads = pool_threads;
2921 }
2922 if (active_threads > jobs) active_threads = jobs;
2923
2924 gemm_q4_packed_meta_x8_thread_work_t work = {
2925 .A = (const block_q8_K *)A_q8,
2926 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2927 .bias = bias,
2928 .C = C,
2929 .M = M,
2930 .N = N,
2931 .K = K,
2932 .blocks_per_vec = K / QK_K,
2933 .blocks_per_row = K / QK_K,
2934 .groups = groups,
2935 .tile_m = tm,
2936 .jobs = jobs,
2937 };
2938 if (active_threads <= 1 || !pool) {
2940 return;
2941 }
2943 pool, active_threads,
2945}
static void gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order()

void gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1899 of file gemm_kernels_q4k_q8k_vnni.c.

1902{
1903 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1904 return;
1905 }
1906 const block_q8_K *A = (const block_q8_K *)A_q8;
1907 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1908 const int blocks_per_row = K / QK_K;
1909 const int groups = (N + 7) / 8;
1910
1911 for (int m = 0; m < M; ++m) {
1912 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_row;
1913 float *c_row = C + (size_t)m * (size_t)N;
1914 for (int g = 0; g < groups; ++g) {
1915 const int n0 = g * 8;
1916 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
1917 float acc[8] = {0};
1918 float acc_min[8] = {0};
1919 for (int b = 0; b < blocks_per_row; ++b) {
1920 const block_q4_K_packed_meta_x8 *w_group =
1921 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
1923 acc, acc_min, w_group, active, &a_row[b]);
1924 }
1925 float values[8];
1926#if defined(__AVX2__)
1927 _mm256_storeu_ps(values, _mm256_sub_ps(_mm256_loadu_ps(acc), _mm256_loadu_ps(acc_min)));
1928#else
1929 for (int lane = 0; lane < active; ++lane) {
1930 values[lane] = acc[lane] - acc_min[lane];
1931 }
1932#endif
1933 for (int lane = 0; lane < active; ++lane) {
1934 float value = values[lane];
1935 if (bias) {
1936 value += bias[n0 + lane];
1937 }
1938 c_row[n0 + lane] = value;
1939 }
1940 }
1941 }
1942}

References accum_q4_k_packed_meta_x8_q8_k_superblock(), C, and QK_K.

Referenced by gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order().

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse()

void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 2860 of file gemm_kernels_q4k_q8k_vnni.c.

2867{
2868 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2869 return;
2870 }
2871 ck_threadpool_t *pool = ck_threadpool_global();
2872 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2873 const int groups = (N + 7) / 8;
2874 int tm = tile_m > 0 ? tile_m : 4;
2875 if (tm > 8) tm = 8;
2876 const int mt = (M + tm - 1) / tm;
2877 const int jobs = mt * groups;
2878 if (active_threads <= 0 || active_threads > pool_threads) {
2879 active_threads = pool_threads;
2880 }
2881 if (active_threads > jobs) {
2882 active_threads = jobs;
2883 }
2884 gemm_q4_packed_meta_x8_thread_work_t work = {
2885 .A = (const block_q8_K *)A_q8,
2886 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2887 .bias = bias,
2888 .C = C,
2889 .M = M,
2890 .N = N,
2891 .K = K,
2892 .blocks_per_vec = K / QK_K,
2893 .blocks_per_row = K / QK_K,
2894 .groups = groups,
2895 .tile_m = tm,
2896 .jobs = jobs,
2897 };
2898 if (active_threads <= 1) {
2900 return;
2901 }
2903}
static void gemm_q4_packed_meta_x8_mreuse_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_meta_x8_mreuse_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile()

void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 2814 of file gemm_kernels_q4k_q8k_vnni.c.

2821{
2822 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2823 return;
2824 }
2825 ck_threadpool_t *pool = ck_threadpool_global();
2826 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2827 const int groups = (N + 7) / 8;
2828 int tm = tile_m > 0 ? tile_m : 4;
2829 if (tm > 8) tm = 8;
2830 const int mt = (M + tm - 1) / tm;
2831 const int jobs = mt * groups;
2832 if (active_threads <= 0 || active_threads > pool_threads) {
2833 active_threads = pool_threads;
2834 }
2835 if (active_threads > jobs) {
2836 active_threads = jobs;
2837 }
2838 if (active_threads <= 1) {
2839 gemm_nt_q4_k_packed_meta_x8_q8_k(A_q8, B_packed_x8, bias, C, M, N, K);
2840 return;
2841 }
2842 gemm_q4_packed_meta_x8_thread_work_t work = {
2843 .A = (const block_q8_K *)A_q8,
2844 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2845 .bias = bias,
2846 .C = C,
2847 .M = M,
2848 .N = N,
2849 .K = K,
2850 .blocks_per_vec = K / QK_K,
2851 .blocks_per_row = K / QK_K,
2852 .groups = groups,
2853 .tile_m = tm,
2854 .jobs = jobs,
2855 };
2857}
static void gemm_q4_packed_meta_x8_mtile_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_x8_q8_k(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_nt_q4_k_packed_meta_x8_q8_k(), gemm_q4_packed_meta_x8_mtile_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit()

void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit ( const void *  A_q8,
const void *  B_packed_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 2381 of file gemm_kernels_q4k_q8k_vnni.c.

2387{
2388 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2389 return;
2390 }
2391 ck_threadpool_t *pool = ck_threadpool_global();
2392 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2393 const int groups = (N + 7) / 8;
2394 if (active_threads <= 0 || active_threads > pool_threads) {
2395 active_threads = pool_threads;
2396 }
2397 if (active_threads > groups) {
2398 active_threads = groups;
2399 }
2400 if (active_threads <= 1) {
2401 gemm_nt_q4_k_packed_meta_x8_q8_k(A_q8, B_packed_x8, bias, C, M, N, K);
2402 return;
2403 }
2404 gemm_q4_packed_meta_x8_thread_work_t work = {
2405 .A = (const block_q8_K *)A_q8,
2406 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2407 .bias = bias,
2408 .C = C,
2409 .M = M,
2410 .N = N,
2411 .K = K,
2412 .blocks_per_vec = K / QK_K,
2413 .blocks_per_row = K / QK_K,
2414 .groups = groups,
2415 };
2417}
static void gemm_q4_packed_meta_x8_nsplit_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_nt_q4_k_packed_meta_x8_q8_k(), gemm_q4_packed_meta_x8_nsplit_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_u8_q8_k()

void gemm_nt_q4_k_packed_u8_q8_k ( const void *  A_q8,
const void *  B_packed,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1770 of file gemm_kernels_q4k_q8k_vnni.c.

1775{
1776 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1777 return;
1778 }
1779 const block_q8_K *A = (const block_q8_K *)A_q8;
1780 const block_q4_K_packed_u8 *W = (const block_q4_K_packed_u8 *)B_packed;
1781 const int blocks_per_vec = K / QK_K;
1782 const int blocks_per_row = K / QK_K;
1783 for (int m = 0; m < M; ++m) {
1784 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1785 float *c_row = C + (size_t)m * (size_t)N;
1786 for (int n = 0; n < N; ++n) {
1787 const block_q4_K_packed_u8 *w_row = W + (size_t)n * (size_t)blocks_per_row;
1788 float sum = bias ? bias[n] : 0.0f;
1789 for (int b = 0; b < blocks_per_row; ++b) {
1790 sum += dot_q4_k_packed_u8_q8_k_block(&w_row[b], &a_row[b]);
1791 }
1792 c_row[n] = sum;
1793 }
1794 }
1795}
static float dot_q4_k_packed_u8_q8_k_block(const block_q4_K_packed_u8 *w, const block_q8_K *x)

References C, dot_q4_k_packed_u8_q8_k_block(), and QK_K.

◆ gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile()

void gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile ( const void *  A_q8,
const void *  B_packed_u8_x16,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  tile_m,
int  active_threads 
)

Definition at line 3494 of file gemm_kernels_q4k_q8k_vnni.c.

3501{
3502 if (!A_q8 || !B_packed_u8_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) return;
3503 ck_threadpool_t *pool = ck_threadpool_global();
3504 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3505 const int groups = (N + 15) / 16;
3506 int tm = tile_m > 0 ? tile_m : 4;
3507 if (tm > 8) tm = 8;
3508 const int mt = (M + tm - 1) / tm;
3509 const int jobs = mt * groups;
3510 if (active_threads <= 0 || active_threads > pool_threads) active_threads = pool_threads;
3511 if (active_threads > jobs) active_threads = jobs;
3512 gemm_q4_packed_u8_x16_thread_work_t work = {
3513 .A = (const block_q8_K *)A_q8,
3514 .W = (const block_q4_K_packed_u8_x16 *)B_packed_u8_x16,
3515 .bias = bias,
3516 .C = C,
3517 .M = M,
3518 .N = N,
3519 .K = K,
3520 .blocks_per_vec = K / QK_K,
3521 .blocks_per_row = K / QK_K,
3522 .groups = groups,
3523 .tile_m = tm,
3524 .jobs = jobs,
3525 };
3526 if (active_threads <= 1) {
3528 return;
3529 }
3531}
static void gemm_q4_packed_u8_x16_mtile_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_packed_u8_x16_mtile_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order()

void gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order ( const void *  A_q8,
const void *  B_packed_x16,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 2094 of file gemm_kernels_q4k_q8k_vnni.c.

2098{
2099 if (!A_q8 || !B_packed_x16 || !C ||
2100 M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0 ||
2102 return;
2103 }
2104 const block_q8_K *A = (const block_q8_K *)A_q8;
2105 const block_q4_K_packed_vnni_x16 *W =
2106 (const block_q4_K_packed_vnni_x16 *)B_packed_x16;
2107 const int blocks_per_row = K / QK_K;
2108 const int groups = (N + 15) / 16;
2109 for (int m = 0; m < M; ++m) {
2110 const block_q8_K *a_row =
2111 A + (size_t)m * (size_t)blocks_per_row;
2112 float *c_row = C + (size_t)m * (size_t)N;
2113 for (int g = 0; g < groups; ++g) {
2114 const int n0 = g * 16;
2115 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
2116 float acc[16] = {0};
2117 float acc_min[16] = {0};
2118 for (int b = 0; b < blocks_per_row; ++b) {
2119 const block_q4_K_packed_vnni_x16 *w_group =
2120 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
2122 acc, acc_min, w_group, &a_row[b]);
2123 }
2124 float values[16];
2125#if defined(CK_HAS_AVX512_VNNI_512)
2126 _mm512_storeu_ps(values, _mm512_sub_ps(
2127 _mm512_loadu_ps(acc), _mm512_loadu_ps(acc_min)));
2128#else
2129 for (int lane = 0; lane < active; ++lane) {
2130 values[lane] = acc[lane] - acc_min[lane];
2131 }
2132#endif
2133 for (int lane = 0; lane < active; ++lane) {
2134 c_row[n0 + lane] =
2135 values[lane] + (bias ? bias[n0 + lane] : 0.0f);
2136 }
2137 }
2138 }
2139}
static void accum_q4_k_packed_vnni_x16_q8_k_gemv_block(float acc[16], float acc_min[16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x)
int ck_q4k_packed_vnni_x16_available(void)

References accum_q4_k_packed_vnni_x16_q8_k_gemv_block(), C, ck_q4k_packed_vnni_x16_available(), and QK_K.

◆ gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m()

void gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m ( const void *  A_q8,
const void *  B_packed_vnni_x16,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 3070 of file gemm_kernels_q4k_q8k_vnni.c.

3073{
3074 if (!A_q8 || !B_packed_vnni_x16 || !C ||
3075 M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0 ||
3077 return;
3078 }
3079 ck_threadpool_t *pool = ck_threadpool_global();
3080 const int pool_threads = pool ? ck_threadpool_capacity(pool) : 1;
3081 const int groups = (N + 15) / 16;
3082 const int jobs = ((M + 15) / 16) * groups;
3083 if (active_threads <= 0 || active_threads > pool_threads) {
3084 active_threads = pool_threads;
3085 }
3086 if (active_threads > jobs) active_threads = jobs;
3087
3088 gemm_q4_packed_vnni_x16_thread_work_t work = {
3089 .A = (const block_q8_K *)A_q8,
3090 .W = (const block_q4_K_packed_vnni_x16 *)B_packed_vnni_x16,
3091 .bias = bias,
3092 .C = C,
3093 .M = M,
3094 .N = N,
3095 .blocks_per_row = K / QK_K,
3096 .groups = groups,
3097 };
3098 if (active_threads <= 1 || !pool) {
3100 return;
3101 }
3103 pool, active_threads,
3105}
int ck_threadpool_capacity(const ck_threadpool_t *pool)
static void gemm_q4_packed_vnni_x16_q8k_16m_thread_fn(int ith, int nth, void *args)

References C, ck_q4k_packed_vnni_x16_available(), ck_threadpool_capacity(), ck_threadpool_dispatch_n(), ck_threadpool_global(), gemm_q4_packed_vnni_x16_q8k_16m_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m()

void gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m ( const void *  A_q8,
const void *  B_packed_vnni_x8,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  active_threads 
)

Definition at line 3027 of file gemm_kernels_q4k_q8k_vnni.c.

3030{
3031 if (!A_q8 || !B_packed_vnni_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
3032 (K % QK_K) != 0) {
3033 return;
3034 }
3035 ck_threadpool_t *pool = ck_threadpool_global();
3036 const int pool_threads = pool ? ck_threadpool_capacity(pool) : 1;
3037 const int groups = (N + 7) / 8;
3038 const int jobs = ((M + 3) / 4) * groups;
3039 if (active_threads <= 0 || active_threads > pool_threads) {
3040 active_threads = pool_threads;
3041 }
3042 if (active_threads > jobs) active_threads = jobs;
3043
3044 gemm_q4_packed_vnni_x8_thread_work_t work = {
3045 .A = (const block_q8_K *)A_q8,
3046 .W = (const block_q4_K_packed_vnni_x8 *)B_packed_vnni_x8,
3047 .bias = bias,
3048 .C = C,
3049 .M = M,
3050 .N = N,
3051 .blocks_per_row = K / QK_K,
3052 .groups = groups,
3053 };
3054 if (active_threads <= 1 || !pool) {
3056 return;
3057 }
3059 const int grain = 4;
3061 pool, active_threads, 0, jobs, grain,
3063 } else {
3065 pool, active_threads,
3067 }
3068}
void ck_threadpool_parallel_for_n(ck_threadpool_t *pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void *args)
int ck_gemm_dynamic_schedule_enabled(void)
static void gemm_q4_packed_vnni_x8_q8k_4m_range_fn(int begin, int end, void *args)
static void gemm_q4_packed_vnni_x8_q8k_4m_thread_fn(int ith, int nth, void *args)

References C, ck_gemm_dynamic_schedule_enabled(), ck_threadpool_capacity(), ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_parallel_for_n(), gemm_q4_packed_vnni_x8_q8k_4m_range_fn(), gemm_q4_packed_vnni_x8_q8k_4m_thread_fn(), and QK_K.

◆ gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni()

void gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni ( const void *  A_q8,
const void *  B_gate_up,
const float *  bias,
float *  C,
int  M,
int  D,
int  K,
int  threads 
)

Definition at line 3722 of file gemm_kernels_q4k_q8k_vnni.c.

3730{
3731#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3732 if (!A_q8 || !B_gate_up || !C || M <= 0 || D <= 0 || K <= 0 || (K % QK_K) != 0) {
3733 return;
3734 }
3735
3736 ck_threadpool_t *pool = ck_threadpool_global();
3737 int active = threads > 0 ? threads : (pool ? ck_threadpool_n_threads(pool) : 1);
3738 if (active < 1) active = 1;
3739 if (active > D) active = D;
3740
3741 gemm_q4_gateup_swiglu_work_t work = {
3742 .A = (const block_q8_K *)A_q8,
3743 .W = (const block_q4_K *)B_gate_up,
3744 .bias = bias,
3745 .C = C,
3746 .M = M,
3747 .D = D,
3748 .K = K,
3749 .blocks_per_vec = K / QK_K,
3750 .blocks_per_row = K / QK_K,
3751 };
3752
3753 if (active <= 1 || !pool) {
3755 return;
3756 }
3758#else
3759 (void)A_q8;
3760 (void)B_gate_up;
3761 (void)bias;
3762 (void)C;
3763 (void)M;
3764 (void)D;
3765 (void)K;
3766 (void)threads;
3767#endif
3768}
static void gemm_q4_gateup_swiglu_thread_fn(int ith, int nth, void *args)

References C, ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_gateup_swiglu_thread_fn(), and QK_K.

◆ gemm_q4_gateup_swiglu_thread_fn()

static void gemm_q4_gateup_swiglu_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 3688 of file gemm_kernels_q4k_q8k_vnni.c.

3689{
3690#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3691 gemm_q4_gateup_swiglu_work_t *a = (gemm_q4_gateup_swiglu_work_t *)args;
3692 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
3693
3694 const int dd = (a->D + nth - 1) / nth;
3695 const int d0 = dd * ith;
3696 const int d1 = (d0 + dd < a->D) ? (d0 + dd) : a->D;
3697 if (d0 >= a->D) return;
3698
3699 for (int d = d0; d < d1; ++d) {
3700 const block_q4_K *w_gate = a->W + (size_t)d * (size_t)a->blocks_per_row;
3701 const block_q4_K *w_up = a->W + (size_t)(a->D + d) * (size_t)a->blocks_per_row;
3702 const float b_gate = a->bias ? a->bias[d] : 0.0f;
3703 const float b_up = a->bias ? a->bias[a->D + d] : 0.0f;
3704 for (int m = 0; m < a->M; ++m) {
3705 const block_q8_K *x = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3706 float gate = b_gate;
3707 float up = b_up;
3708 for (int b = 0; b < a->blocks_per_row; ++b) {
3709 gate += dot_q4_k_q8_k_vnni_block(&w_gate[b], &x[b]);
3710 up += dot_q4_k_q8_k_vnni_block(&w_up[b], &x[b]);
3711 }
3712 a->C[(size_t)m * (size_t)a->D + (size_t)d] = ck_q4k_silu_f32(gate) * up;
3713 }
3714 }
3715#else
3716 (void)ith;
3717 (void)nth;
3718 (void)args;
3719#endif
3720}
static float ck_q4k_silu_f32(float x)

References ck_q4k_silu_f32().

Referenced by gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni().

◆ gemm_q4_gateup_swiglu_x16_thread_fn()

static void gemm_q4_gateup_swiglu_x16_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 3335 of file gemm_kernels_q4k_q8k_vnni.c.

3336{
3337 gemm_q4_gateup_swiglu_x16_work_t *a = (gemm_q4_gateup_swiglu_x16_work_t *)args;
3338 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3339 return;
3340 }
3341
3342 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3343 if (tile_m > 8) tile_m = 8;
3344 const int mt = (a->M + tile_m - 1) / tile_m;
3345
3346 for (int job = ith; job < a->jobs; job += nth) {
3347 const int g = job / mt;
3348 const int tm = job - g * mt;
3349 const int m0 = tm * tile_m;
3350 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3351 const int m_count = m1 - m0;
3352 const int d0 = g * 16;
3353 const int active = (d0 + 16 <= a->D) ? 16 : (a->D - d0);
3354 if (m_count <= 0 || active <= 0 || g >= a->groups_d) {
3355 continue;
3356 }
3357
3358 float acc_gate[8][16];
3359 float acc_up[8][16];
3360 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
3361 for (int lane = 0; lane < active; ++lane) {
3362 acc_gate[mt_lane][lane] = a->bias ? a->bias[d0 + lane] : 0.0f;
3363 acc_up[mt_lane][lane] = a->bias ? a->bias[a->D + d0 + lane] : 0.0f;
3364 }
3365 }
3366
3367 for (int b = 0; b < a->blocks_per_row; ++b) {
3368 const block_q4_K_packed_meta_x16 *w_gate =
3369 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3370 const block_q4_K_packed_meta_x16 *w_up =
3371 a->W + (size_t)(a->groups_d + g) * (size_t)a->blocks_per_row + (size_t)b;
3373 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc_gate, w_gate, active, a->A,
3374 a->blocks_per_vec, b, m0, m_count);
3375 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc_up, w_up, active, a->A,
3376 a->blocks_per_vec, b, m0, m_count);
3377 } else {
3378 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc_gate, w_gate, active, a->A,
3379 a->blocks_per_vec, b, m0, m_count);
3380 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc_up, w_up, active, a->A,
3381 a->blocks_per_vec, b, m0, m_count);
3382 }
3383 }
3384
3385 for (int m = m0; m < m1; ++m) {
3386 float *c_row = a->C + (size_t)m * (size_t)a->D;
3387#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
3388#pragma clang loop vectorize(enable) interleave(enable)
3389#elif defined(__GNUC__)
3390#pragma GCC ivdep
3391#endif
3392 for (int lane = 0; lane < active; ++lane) {
3393 const float gate = acc_gate[m - m0][lane];
3394 const float up = acc_up[m - m0][lane];
3395 c_row[d0 + lane] = (gate / (1.0f + expf(-gate))) * up;
3396 }
3397 }
3398 }
3399}
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
static int ck_q4k_x16_chunk4_enabled(void)

References accum_q4_k_packed_meta_x16_q8_k_block_mreuse(), accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(), and ck_q4k_x16_chunk4_enabled().

Referenced by gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni().

◆ gemm_q4_k_q8_k_compact_rows4()

void gemm_q4_k_q8_k_compact_rows4 ( float *  output,
int  output_stride,
const void *  weights,
const void *const  input_rows[4],
int  rows,
int  output_dim,
int  input_dim 
)

Definition at line 3615 of file gemm_kernels_q4k_q8k_vnni.c.

3622{
3623 if (!output || !weights || !input_rows || rows <= 0 || rows > 4 ||
3624 output_stride < output_dim || output_dim <= 0 || input_dim <= 0 ||
3625 (input_dim % QK_K) != 0) {
3626 return;
3627 }
3628 for (int row = 0; row < rows; ++row) {
3629 if (!input_rows[row]) return;
3630 }
3631
3632#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3633 const block_q4_K *blocks = (const block_q4_K *)weights;
3634 const int blocks_per_row = input_dim / QK_K;
3635 const block_q8_K *inputs[4] = {
3636 (const block_q8_K *)input_rows[0],
3637 (const block_q8_K *)input_rows[rows > 1 ? 1 : 0],
3638 (const block_q8_K *)input_rows[rows > 2 ? 2 : 0],
3639 (const block_q8_K *)input_rows[rows > 3 ? 3 : 0],
3640 };
3641 for (int n = 0; n < output_dim; ++n) {
3642 const block_q4_K *weight_row =
3643 blocks + (size_t)n * (size_t)blocks_per_row;
3644 float sums[4] = {0.0f, 0.0f, 0.0f, 0.0f};
3645 for (int block = 0; block < blocks_per_row; ++block) {
3646 const block_q8_K *block_rows[4] = {
3647 &inputs[0][block], &inputs[1][block],
3648 &inputs[2][block], &inputs[3][block],
3649 };
3650 float block_sums[4];
3651 dot_q4_k_q8_k_vnni_block_rows4(
3652 &weight_row[block], block_rows, rows, block_sums);
3653 for (int row = 0; row < rows; ++row) {
3654 sums[row] += block_sums[row];
3655 }
3656 }
3657 for (int row = 0; row < rows; ++row) {
3658 output[(size_t)row * (size_t)output_stride + (size_t)n] = sums[row];
3659 }
3660 }
3661#else
3662 for (int row = 0; row < rows; ++row) {
3664 output + (size_t)row * (size_t)output_stride,
3665 weights, input_rows[row], output_dim, input_dim);
3666 }
3667#endif
3668}
void gemv_q4_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q4_k_q8_k(), and QK_K.

Referenced by ck_moe_q4k_q5k_bucket_work().

◆ gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4()

void gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4 ( float *  output,
const void *  weights_packed,
const void *  input_q8,
int  rows,
int  output_dim,
int  input_dim 
)

Definition at line 1123 of file gemm_kernels_q4k_q8k_vnni.c.

1130{
1131 if (!output || !weights_packed || !input_q8 || rows <= 0 || rows > 4 ||
1132 output_dim <= 0 || input_dim <= 0 || (input_dim % QK_K) != 0) {
1133 return;
1134 }
1135#if defined(CK_HAS_AVX_VNNI_256)
1136 const block_q8_K *input = (const block_q8_K *)input_q8;
1137 const block_q4_K_packed_vnni_x8 *weights =
1138 (const block_q4_K_packed_vnni_x8 *)weights_packed;
1139 const int blocks_per_row = input_dim / QK_K;
1140 const int groups = (output_dim + 7) / 8;
1141 for (int group = 0; group < groups; ++group) {
1142 const int n0 = group * 8;
1143 const int active = n0 + 8 <= output_dim ? 8 : output_dim - n0;
1144 float acc[4][8] = {{0}};
1145 for (int block = 0; block < blocks_per_row; ++block) {
1146 float block_sums[4][8] = {{0}};
1147 const block_q8_K *input_rows[4] = {NULL, NULL, NULL, NULL};
1148 for (int row = 0; row < rows; ++row) {
1149 input_rows[row] = input +
1150 (size_t)row * (size_t)blocks_per_row + (size_t)block;
1151 }
1153 block_sums,
1154 weights + (size_t)group * (size_t)blocks_per_row +
1155 (size_t)block,
1156 input_rows,
1157 rows);
1158 for (int row = 0; row < rows; ++row) {
1159 const __m256 prior = _mm256_loadu_ps(acc[row]);
1160 const __m256 current = _mm256_loadu_ps(block_sums[row]);
1161 _mm256_storeu_ps(acc[row], _mm256_add_ps(prior, current));
1162 }
1163 }
1164 for (int row = 0; row < rows; ++row) {
1165 for (int lane = 0; lane < active; ++lane) {
1166 output[(size_t)row * (size_t)output_dim +
1167 (size_t)n0 + (size_t)lane] = acc[row][lane];
1168 }
1169 }
1170 }
1171#else
1172 (void)rows;
1173 (void)output_dim;
1174 (void)input_dim;
1175#endif
1176}
static void dot_q4_k_packed_vnni_x8_q8_k_compact_order(float block_sums[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)

References dot_q4_k_packed_vnni_x8_q8_k_compact_order(), and QK_K.

Referenced by ck_moe_q4k_q5k_bucket_work().

◆ gemm_q4_packed_meta_nsplit_thread_fn()

static void gemm_q4_packed_meta_nsplit_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2284 of file gemm_kernels_q4k_q8k_vnni.c.

2285{
2286 gemm_q4_packed_meta_thread_work_t *a = (gemm_q4_packed_meta_thread_work_t *)args;
2287 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2288 return;
2289 }
2290 const int dn = (a->N + nth - 1) / nth;
2291 const int n0 = dn * ith;
2292 const int n1 = (n0 + dn < a->N) ? (n0 + dn) : a->N;
2293 if (n0 >= a->N) {
2294 return;
2295 }
2296 for (int m = 0; m < a->M; ++m) {
2297 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2298 float *c_row = a->C + (size_t)m * (size_t)a->N;
2299 for (int n = n0; n < n1; ++n) {
2300 const block_q4_K_packed_meta *w_row = a->W + (size_t)n * (size_t)a->blocks_per_row;
2301 float sum = a->bias ? a->bias[n] : 0.0f;
2302 for (int b = 0; b < a->blocks_per_row; ++b) {
2303 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
2304 }
2305 c_row[n] = sum;
2306 }
2307 }
2308}

References dot_q4_k_packed_meta_q8_k_block().

Referenced by gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit().

◆ gemm_q4_packed_meta_thread_fn()

static void gemm_q4_packed_meta_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2221 of file gemm_kernels_q4k_q8k_vnni.c.

2222{
2223 gemm_q4_packed_meta_thread_work_t *a = (gemm_q4_packed_meta_thread_work_t *)args;
2224 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2225 return;
2226 }
2227 const int dm = (a->M + nth - 1) / nth;
2228 const int m0 = dm * ith;
2229 const int m1 = (m0 + dm < a->M) ? (m0 + dm) : a->M;
2230 if (m0 >= a->M) {
2231 return;
2232 }
2233 for (int m = m0; m < m1; ++m) {
2234 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2235 float *c_row = a->C + (size_t)m * (size_t)a->N;
2236 for (int n = 0; n < a->N; ++n) {
2237 const block_q4_K_packed_meta *w_row = a->W + (size_t)n * (size_t)a->blocks_per_row;
2238 float sum = a->bias ? a->bias[n] : 0.0f;
2239 for (int b = 0; b < a->blocks_per_row; ++b) {
2240 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
2241 }
2242 c_row[n] = sum;
2243 }
2244 }
2245}

References dot_q4_k_packed_meta_q8_k_block().

Referenced by gemm_nt_q4_k_packed_meta_q8_k_threaded().

◆ gemm_q4_packed_meta_x16_mreuse_process_job()

static void gemm_q4_packed_meta_x16_mreuse_process_job ( const gemm_q4_packed_meta_x16_thread_work_t *  a,
int  job,
int  mt,
int  tile_m 
)
inlinestatic

Definition at line 3156 of file gemm_kernels_q4k_q8k_vnni.c.

3161{
3162 const int g = job / mt;
3163 const int tm = job - g * mt;
3164 const int m0 = tm * tile_m;
3165 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3166 const int m_count = m1 - m0;
3167 const int n0 = g * 16;
3168 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3169 if (m0 >= a->M || g >= a->groups || m_count <= 0) {
3170 return;
3171 }
3172
3173#if defined(__AVX2__)
3174 const block_q4_K_packed_meta_x16 *w_group =
3175 a->W + (size_t)g * (size_t)a->blocks_per_row;
3176 for (int m = m0; m < m1; ++m) {
3177 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3178 float *c_row = a->C + (size_t)m * (size_t)a->N;
3179 for (int lane = 0; lane < active; ++lane) {
3180 float value = dot_q4_k_packed_meta_x16_q8_k_llama_avx2(
3181 w_group, a->blocks_per_row, lane, a_row);
3182 c_row[n0 + lane] = value + (a->bias ? a->bias[n0 + lane] : 0.0f);
3183 }
3184 }
3185#else
3186 float acc[8][16];
3187 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
3188 for (int lane = 0; lane < active; ++lane) {
3189 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3190 }
3191 }
3192 for (int b = 0; b < a->blocks_per_row; ++b) {
3193 const block_q4_K_packed_meta_x16 *w_block =
3194 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3196 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc, w_block, active, a->A,
3197 a->blocks_per_vec, b, m0, m_count);
3198 } else {
3199 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc, w_block, active, a->A,
3200 a->blocks_per_vec, b, m0, m_count);
3201 }
3202 }
3203 for (int m = m0; m < m1; ++m) {
3204 float *c_row = a->C + (size_t)m * (size_t)a->N;
3205 for (int lane = 0; lane < active; ++lane) {
3206 c_row[n0 + lane] = acc[m - m0][lane];
3207 }
3208 }
3209#endif
3210}

References accum_q4_k_packed_meta_x16_q8_k_block_mreuse(), accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(), and ck_q4k_x16_chunk4_enabled().

Referenced by gemm_q4_packed_meta_x16_mreuse_thread_fn().

◆ gemm_q4_packed_meta_x16_mreuse_thread_fn()

static void gemm_q4_packed_meta_x16_mreuse_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 3212 of file gemm_kernels_q4k_q8k_vnni.c.

3213{
3214 gemm_q4_packed_meta_x16_thread_work_t *a = (gemm_q4_packed_meta_x16_thread_work_t *)args;
3215 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3216 return;
3217 }
3218 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3219 if (tile_m > 8) tile_m = 8;
3220 const int mt = (a->M + tile_m - 1) / tile_m;
3221 const int total = mt * a->groups;
3222
3223 for (int job = ith; job < total; job += nth) {
3225 }
3226}
static void gemm_q4_packed_meta_x16_mreuse_process_job(const gemm_q4_packed_meta_x16_thread_work_t *a, int job, int mt, int tile_m)

References gemm_q4_packed_meta_x16_mreuse_process_job().

Referenced by gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse().

◆ gemm_q4_packed_meta_x16_mtile_thread_fn()

static void gemm_q4_packed_meta_x16_mtile_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 3108 of file gemm_kernels_q4k_q8k_vnni.c.

3109{
3110 gemm_q4_packed_meta_x16_thread_work_t *a = (gemm_q4_packed_meta_x16_thread_work_t *)args;
3111 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3112 return;
3113 }
3114 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3115 if (tile_m > 8) tile_m = 8;
3116 const int mt = (a->M + tile_m - 1) / tile_m;
3117 const int total = mt * a->groups;
3118
3119 for (int job = ith; job < total; job += nth) {
3120 const int g = job / mt;
3121 const int tm = job - g * mt;
3122 const int m0 = tm * tile_m;
3123 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3124 const int n0 = g * 16;
3125 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3126 if (m0 >= a->M || g >= a->groups) {
3127 continue;
3128 }
3129
3130 float acc[8][16];
3131 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
3132 for (int lane = 0; lane < active; ++lane) {
3133 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3134 }
3135 }
3136
3137 for (int b = 0; b < a->blocks_per_row; ++b) {
3138 const block_q4_K_packed_meta_x16 *w_group =
3139 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3140 for (int m = m0; m < m1; ++m) {
3141 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3142 accum_q4_k_packed_meta_x16_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
3143 }
3144 }
3145
3146 for (int m = m0; m < m1; ++m) {
3147 float *c_row = a->C + (size_t)m * (size_t)a->N;
3148 for (int lane = 0; lane < active; ++lane) {
3149 c_row[n0 + lane] = acc[m - m0][lane];
3150 }
3151 }
3152 }
3153}
static void accum_q4_k_packed_meta_x16_q8_k_block(float acc[16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *x)

References accum_q4_k_packed_meta_x16_q8_k_block().

Referenced by gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile().

◆ gemm_q4_packed_meta_x8_mreuse_thread_fn()

static void gemm_q4_packed_meta_x8_mreuse_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2467 of file gemm_kernels_q4k_q8k_vnni.c.

2468{
2469 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2470 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2471 return;
2472 }
2473 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2474 if (tile_m > 8) tile_m = 8;
2475 const int mt = (a->M + tile_m - 1) / tile_m;
2476 const int total = mt * a->groups;
2477
2478 for (int job = ith; job < total; job += nth) {
2479 const int g = job / mt;
2480 const int tm = job - g * mt;
2481 const int m0 = tm * tile_m;
2482 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2483 const int m_count = m1 - m0;
2484 const int n0 = g * 8;
2485 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2486 if (m0 >= a->M || g >= a->groups || m_count <= 0) {
2487 continue;
2488 }
2489
2490 float acc[8][8];
2491 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
2492 for (int lane = 0; lane < active; ++lane) {
2493 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2494 }
2495 }
2496
2497 for (int b = 0; b < a->blocks_per_row; ++b) {
2498 const block_q4_K_packed_meta_x8 *w_group =
2499 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2500 accum_q4_k_packed_meta_x8_q8_k_block_mreuse(acc, w_group, active, a->A,
2501 a->blocks_per_vec, b, m0, m_count);
2502 }
2503
2504 for (int m = m0; m < m1; ++m) {
2505 float *c_row = a->C + (size_t)m * (size_t)a->N;
2506 for (int lane = 0; lane < active; ++lane) {
2507 c_row[n0 + lane] = acc[m - m0][lane];
2508 }
2509 }
2510 }
2511}
static void accum_q4_k_packed_meta_x8_q8_k_block_mreuse(float acc[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)

References accum_q4_k_packed_meta_x8_q8_k_block_mreuse().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse().

◆ gemm_q4_packed_meta_x8_mtile_thread_fn()

static void gemm_q4_packed_meta_x8_mtile_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2419 of file gemm_kernels_q4k_q8k_vnni.c.

2420{
2421 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2422 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2423 return;
2424 }
2425 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2426 if (tile_m > 8) tile_m = 8;
2427 const int mt = (a->M + tile_m - 1) / tile_m;
2428 const int total = mt * a->groups;
2429
2430 for (int job = ith; job < total; job += nth) {
2431 const int g = job / mt;
2432 const int tm = job - g * mt;
2433 const int m0 = tm * tile_m;
2434 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2435 const int n0 = g * 8;
2436 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2437 if (m0 >= a->M || g >= a->groups) {
2438 continue;
2439 }
2440
2441 float acc[8][8];
2442 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
2443 for (int lane = 0; lane < active; ++lane) {
2444 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2445 }
2446 }
2447
2448 for (int b = 0; b < a->blocks_per_row; ++b) {
2449 const block_q4_K_packed_meta_x8 *w_group =
2450 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2451 for (int m = m0; m < m1; ++m) {
2452 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2453 accum_q4_k_packed_meta_x8_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
2454 }
2455 }
2456
2457 for (int m = m0; m < m1; ++m) {
2458 float *c_row = a->C + (size_t)m * (size_t)a->N;
2459 for (int lane = 0; lane < active; ++lane) {
2460 c_row[n0 + lane] = acc[m - m0][lane];
2461 }
2462 }
2463 }
2464}

References accum_q4_k_packed_meta_x8_q8_k_block().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile().

◆ gemm_q4_packed_meta_x8_nsplit_thread_fn()

static void gemm_q4_packed_meta_x8_nsplit_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2346 of file gemm_kernels_q4k_q8k_vnni.c.

2347{
2348 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2349 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2350 return;
2351 }
2352 const int dg = (a->groups + nth - 1) / nth;
2353 const int g0 = dg * ith;
2354 const int g1 = (g0 + dg < a->groups) ? (g0 + dg) : a->groups;
2355 if (g0 >= a->groups) {
2356 return;
2357 }
2358
2359 for (int m = 0; m < a->M; ++m) {
2360 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2361 float *c_row = a->C + (size_t)m * (size_t)a->N;
2362 for (int g = g0; g < g1; ++g) {
2363 const int n0 = g * 8;
2364 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2365 float acc[8];
2366 for (int lane = 0; lane < active; ++lane) {
2367 acc[lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2368 }
2369 for (int b = 0; b < a->blocks_per_row; ++b) {
2370 const block_q4_K_packed_meta_x8 *w_group =
2371 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2372 accum_q4_k_packed_meta_x8_q8_k_block(acc, w_group, active, &a_row[b]);
2373 }
2374 for (int lane = 0; lane < active; ++lane) {
2375 c_row[n0 + lane] = acc[lane];
2376 }
2377 }
2378 }
2379}

References accum_q4_k_packed_meta_x8_q8_k_block().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit().

◆ gemm_q4_packed_meta_x8_split_min_4m_thread_fn()

static void gemm_q4_packed_meta_x8_split_min_4m_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2576 of file gemm_kernels_q4k_q8k_vnni.c.

2578{
2579 gemm_q4_packed_meta_x8_thread_work_t *a =
2580 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2581 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2582 return;
2583 }
2584
2585 const int row_tiles = (a->M + 3) / 4;
2586 const int total = row_tiles * a->groups;
2587 for (int job = ith; job < total; job += nth) {
2588 const int g = job / row_tiles;
2589 const int row_tile = job - g * row_tiles;
2590 const int m0 = row_tile * 4;
2591 const int rows = (m0 + 4 <= a->M) ? 4 : (a->M - m0);
2592 const int n0 = g * 8;
2593 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2594 if (rows <= 0 || active <= 0 || g >= a->groups) {
2595 continue;
2596 }
2597
2598 float acc[8][8] = {{0}};
2599 float acc_min[8][8] = {{0}};
2600 for (int b = 0; b < a->blocks_per_row; ++b) {
2601 const block_q4_K_packed_meta_x8 *w_group =
2602 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2603 const block_q8_K *x[8] = {NULL};
2604 for (int row = 0; row < rows; ++row) {
2605 x[row] = a->A + (size_t)(m0 + row) *
2606 (size_t)a->blocks_per_vec + (size_t)b;
2607 }
2609 acc, acc_min, w_group, active, x, rows);
2610 }
2611
2612 for (int row = 0; row < rows; ++row) {
2613 float values[8];
2614#if defined(__AVX2__)
2615 _mm256_storeu_ps(values,
2616 _mm256_sub_ps(_mm256_loadu_ps(acc[row]),
2617 _mm256_loadu_ps(acc_min[row])));
2618#else
2619 for (int lane = 0; lane < active; ++lane) {
2620 values[lane] = acc[row][lane] - acc_min[row][lane];
2621 }
2622#endif
2623 float *c_row = a->C + (size_t)(m0 + row) * (size_t)a->N;
2624 for (int lane = 0; lane < active; ++lane) {
2625 float value = values[lane];
2626 if (a->bias) value += a->bias[n0 + lane];
2627 c_row[n0 + lane] = value;
2628 }
2629 }
2630 }
2631}
static void accum_q4_k_packed_meta_x8_q8_k_superblock_rows(float acc[8][8], float acc_min[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x[8], int rows)

References accum_q4_k_packed_meta_x8_q8_k_superblock_rows().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m().

◆ gemm_q4_packed_meta_x8_split_min_8m_thread_fn()

static void gemm_q4_packed_meta_x8_split_min_8m_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2633 of file gemm_kernels_q4k_q8k_vnni.c.

2635{
2636 gemm_q4_packed_meta_x8_thread_work_t *a =
2637 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2638 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2639 return;
2640 }
2641
2642 const int row_tiles = (a->M + 7) / 8;
2643 const int total = row_tiles * a->groups;
2644 for (int job = ith; job < total; job += nth) {
2645 const int g = job / row_tiles;
2646 const int row_tile = job - g * row_tiles;
2647 const int m0 = row_tile * 8;
2648 const int rows = (m0 + 8 <= a->M) ? 8 : (a->M - m0);
2649 const int n0 = g * 8;
2650 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2651 if (rows <= 0 || active <= 0 || g >= a->groups) {
2652 continue;
2653 }
2654
2655 float acc[8][8] = {{0}};
2656 float acc_min[8][8] = {{0}};
2657 for (int b = 0; b < a->blocks_per_row; ++b) {
2658 const block_q4_K_packed_meta_x8 *w_group =
2659 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2660 const block_q8_K *x[8] = {NULL};
2661 for (int row = 0; row < rows; ++row) {
2662 x[row] = a->A + (size_t)(m0 + row) *
2663 (size_t)a->blocks_per_vec + (size_t)b;
2664 }
2666 acc, acc_min, w_group, active, x, rows);
2667 }
2668
2669 for (int row = 0; row < rows; ++row) {
2670 float values[8];
2671#if defined(__AVX2__)
2672 _mm256_storeu_ps(values,
2673 _mm256_sub_ps(_mm256_loadu_ps(acc[row]),
2674 _mm256_loadu_ps(acc_min[row])));
2675#else
2676 for (int lane = 0; lane < active; ++lane) {
2677 values[lane] = acc[row][lane] - acc_min[row][lane];
2678 }
2679#endif
2680 float *c_row = a->C + (size_t)(m0 + row) * (size_t)a->N;
2681 for (int lane = 0; lane < active; ++lane) {
2682 float value = values[lane];
2683 if (a->bias) value += a->bias[n0 + lane];
2684 c_row[n0 + lane] = value;
2685 }
2686 }
2687 }
2688}

References accum_q4_k_packed_meta_x8_q8_k_superblock_rows().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m().

◆ gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn()

static void gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2516 of file gemm_kernels_q4k_q8k_vnni.c.

2518{
2519 gemm_q4_packed_meta_x8_thread_work_t *a =
2520 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2521 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2522 return;
2523 }
2524
2525 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2526 if (tile_m > 8) tile_m = 8;
2527 const int mt = (a->M + tile_m - 1) / tile_m;
2528 const int total = mt * a->groups;
2529
2530 for (int job = ith; job < total; job += nth) {
2531 const int g = job / mt;
2532 const int tm = job - g * mt;
2533 const int m0 = tm * tile_m;
2534 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2535 const int m_count = m1 - m0;
2536 const int n0 = g * 8;
2537 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2538 if (m_count <= 0 || g >= a->groups) {
2539 continue;
2540 }
2541
2542 float acc[8][8] = {{0}};
2543 float acc_min[8][8] = {{0}};
2544 for (int b = 0; b < a->blocks_per_row; ++b) {
2545 const block_q4_K_packed_meta_x8 *w_group =
2546 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2547 for (int m = m0; m < m1; ++m) {
2548 const block_q8_K *a_row =
2549 a->A + (size_t)m * (size_t)a->blocks_per_vec;
2551 acc[m - m0], acc_min[m - m0], w_group, active, &a_row[b]);
2552 }
2553 }
2554
2555 for (int m = m0; m < m1; ++m) {
2556 float values[8];
2557#if defined(__AVX2__)
2558 _mm256_storeu_ps(values,
2559 _mm256_sub_ps(_mm256_loadu_ps(acc[m - m0]),
2560 _mm256_loadu_ps(acc_min[m - m0])));
2561#else
2562 for (int lane = 0; lane < active; ++lane) {
2563 values[lane] = acc[m - m0][lane] - acc_min[m - m0][lane];
2564 }
2565#endif
2566 float *c_row = a->C + (size_t)m * (size_t)a->N;
2567 for (int lane = 0; lane < active; ++lane) {
2568 float value = values[lane];
2569 if (a->bias) value += a->bias[n0 + lane];
2570 c_row[n0 + lane] = value;
2571 }
2572 }
2573 }
2574}

References accum_q4_k_packed_meta_x8_q8_k_superblock().

Referenced by gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse().

◆ gemm_q4_packed_u8_x16_mtile_thread_fn()

static void gemm_q4_packed_u8_x16_mtile_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 3451 of file gemm_kernels_q4k_q8k_vnni.c.

3452{
3453 gemm_q4_packed_u8_x16_thread_work_t *a = (gemm_q4_packed_u8_x16_thread_work_t *)args;
3454 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
3455 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3456 if (tile_m > 8) tile_m = 8;
3457 const int mt = (a->M + tile_m - 1) / tile_m;
3458 const int total = mt * a->groups;
3459
3460 for (int job = ith; job < total; job += nth) {
3461 const int g = job / mt;
3462 const int tm = job - g * mt;
3463 const int m0 = tm * tile_m;
3464 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3465 const int n0 = g * 16;
3466 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3467 if (m0 >= a->M || g >= a->groups) continue;
3468
3469 float acc[8][16];
3470 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
3471 for (int lane = 0; lane < active; ++lane) {
3472 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3473 }
3474 }
3475
3476 for (int b = 0; b < a->blocks_per_row; ++b) {
3477 const block_q4_K_packed_u8_x16 *w_group =
3478 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3479 for (int m = m0; m < m1; ++m) {
3480 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3481 accum_q4_k_packed_u8_x16_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
3482 }
3483 }
3484
3485 for (int m = m0; m < m1; ++m) {
3486 float *c_row = a->C + (size_t)m * (size_t)a->N;
3487 for (int lane = 0; lane < active; ++lane) {
3488 c_row[n0 + lane] = acc[m - m0][lane];
3489 }
3490 }
3491 }
3492}
static void accum_q4_k_packed_u8_x16_q8_k_block(float acc[16], const block_q4_K_packed_u8_x16 *w, int active, const block_q8_K *x)

References accum_q4_k_packed_u8_x16_q8_k_block().

Referenced by gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile().

◆ gemm_q4_packed_vnni_x16_q8k_16m_thread_fn()

static void gemm_q4_packed_vnni_x16_q8k_16m_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2759 of file gemm_kernels_q4k_q8k_vnni.c.

2761{
2762 gemm_q4_packed_vnni_x16_thread_work_t *a =
2763 (gemm_q4_packed_vnni_x16_thread_work_t *)args;
2764 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2765 return;
2766 }
2767
2768 const int row_tiles = (a->M + 15) / 16;
2769 const int total = row_tiles * a->groups;
2770 for (int job = ith; job < total; job += nth) {
2771 const int group = job / row_tiles;
2772 const int row_tile = job - group * row_tiles;
2773 const int m0 = row_tile * 16;
2774 const int rows = (m0 + 16 <= a->M) ? 16 : (a->M - m0);
2775 const int n0 = group * 16;
2776 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
2777 if (rows <= 0 || active <= 0 || group >= a->groups) {
2778 continue;
2779 }
2780
2781 float acc[16][16] = {{0}};
2782 float acc_min[16][16] = {{0}};
2783 for (int block = 0; block < a->blocks_per_row; ++block) {
2784 const block_q4_K_packed_vnni_x16 *weights =
2785 a->W + (size_t)group * (size_t)a->blocks_per_row +
2786 (size_t)block;
2787 const block_q8_K *x[16] = {NULL};
2788 for (int row = 0; row < rows; ++row) {
2789 x[row] = a->A + (size_t)(m0 + row) *
2790 (size_t)a->blocks_per_row + (size_t)block;
2791 }
2793 acc, acc_min, weights, x, rows);
2794 }
2795
2796 for (int row = 0; row < rows; ++row) {
2797#if defined(CK_HAS_AVX512_VNNI_512)
2798 float values[16];
2799 _mm512_storeu_ps(values, _mm512_sub_ps(
2800 _mm512_loadu_ps(acc[row]),
2801 _mm512_loadu_ps(acc_min[row])));
2802 float *output = a->C + (size_t)(m0 + row) * (size_t)a->N;
2803 for (int lane = 0; lane < active; ++lane) {
2804 output[n0 + lane] = values[lane] +
2805 (a->bias ? a->bias[n0 + lane] : 0.0f);
2806 }
2807#else
2808 (void)active;
2809#endif
2810 }
2811 }
2812}
static void accum_q4_k_packed_vnni_x16_q8_k_16m_superblock(float acc[16][16], float acc_min[16][16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x[16], int rows)

References accum_q4_k_packed_vnni_x16_q8_k_16m_superblock().

Referenced by gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m().

◆ gemm_q4_packed_vnni_x8_q8k_4m_job()

static void gemm_q4_packed_vnni_x8_q8k_4m_job ( gemm_q4_packed_vnni_x8_thread_work_t *  a,
int  job,
int  row_tiles 
)
inlinestatic

Definition at line 2690 of file gemm_kernels_q4k_q8k_vnni.c.

2694{
2695 const int group = job / row_tiles;
2696 const int row_tile = job - group * row_tiles;
2697 const int m0 = row_tile * 4;
2698 const int rows = (m0 + 4 <= a->M) ? 4 : (a->M - m0);
2699 const int n0 = group * 8;
2700 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2701 if (rows <= 0 || active <= 0 || group >= a->groups) {
2702 return;
2703 }
2704
2705 float acc[4][8] = {{0}};
2706 float acc_min[4][8] = {{0}};
2707 for (int block = 0; block < a->blocks_per_row; ++block) {
2708 const block_q4_K_packed_vnni_x8 *weights =
2709 a->W + (size_t)group * (size_t)a->blocks_per_row +
2710 (size_t)block;
2711 const block_q8_K *x[4] = {NULL};
2712 for (int row = 0; row < rows; ++row) {
2713 x[row] = a->A + (size_t)(m0 + row) *
2714 (size_t)a->blocks_per_row + (size_t)block;
2715 }
2717 acc, acc_min, weights, x, rows);
2718 }
2719
2720 for (int row = 0; row < rows; ++row) {
2721 float values[8];
2722 _mm256_storeu_ps(values, _mm256_sub_ps(
2723 _mm256_loadu_ps(acc[row]),
2724 _mm256_loadu_ps(acc_min[row])));
2725 float *output = a->C + (size_t)(m0 + row) * (size_t)a->N;
2726 for (int lane = 0; lane < active; ++lane) {
2727 output[n0 + lane] = values[lane] +
2728 (a->bias ? a->bias[n0 + lane] : 0.0f);
2729 }
2730 }
2731}
static void accum_q4_k_packed_vnni_x8_q8_k_4m_superblock(float acc[4][8], float acc_min[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)

References accum_q4_k_packed_vnni_x8_q8_k_4m_superblock().

Referenced by gemm_q4_packed_vnni_x8_q8k_4m_range_fn(), and gemm_q4_packed_vnni_x8_q8k_4m_thread_fn().

◆ gemm_q4_packed_vnni_x8_q8k_4m_range_fn()

static void gemm_q4_packed_vnni_x8_q8k_4m_range_fn ( int  begin,
int  end,
void *  args 
)
static

Definition at line 2747 of file gemm_kernels_q4k_q8k_vnni.c.

2749{
2750 gemm_q4_packed_vnni_x8_thread_work_t *a =
2751 (gemm_q4_packed_vnni_x8_thread_work_t *)args;
2752 if (!a || begin < 0 || begin >= end) return;
2753 const int row_tiles = (a->M + 3) / 4;
2754 for (int job = begin; job < end; ++job) {
2755 gemm_q4_packed_vnni_x8_q8k_4m_job(a, job, row_tiles);
2756 }
2757}
static void gemm_q4_packed_vnni_x8_q8k_4m_job(gemm_q4_packed_vnni_x8_thread_work_t *a, int job, int row_tiles)
uint32_t end
Definition utf8.c:215

References end, and gemm_q4_packed_vnni_x8_q8k_4m_job().

Referenced by gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m().

◆ gemm_q4_packed_vnni_x8_q8k_4m_thread_fn()

static void gemm_q4_packed_vnni_x8_q8k_4m_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 2733 of file gemm_kernels_q4k_q8k_vnni.c.

2735{
2736 gemm_q4_packed_vnni_x8_thread_work_t *a =
2737 (gemm_q4_packed_vnni_x8_thread_work_t *)args;
2738 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
2739
2740 const int row_tiles = (a->M + 3) / 4;
2741 const int total = row_tiles * a->groups;
2742 for (int job = ith; job < total; job += nth) {
2743 gemm_q4_packed_vnni_x8_q8k_4m_job(a, job, row_tiles);
2744 }
2745}

References gemm_q4_packed_vnni_x8_q8k_4m_job().

Referenced by gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m().

◆ gemv_q4_k_q8_k_avx2()

void gemv_q4_k_q8_k_avx2 ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 118 of file gemm_kernels_q4k_q8k_avx2.c.

122{
123#if defined(__AVX2__)
124 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
125 return;
126 }
127
128 const block_q4_K *blocks = (const block_q4_K *)W;
129 const block_q8_K *x = (const block_q8_K *)x_q8;
130 const int blocks_per_row = K / QK_K;
131
132 for (int row = 0; row < M; ++row) {
133 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
134 y[row] = dot_q4_k_q8_k_avx2(w_row, x, K);
135 }
136#else
137 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
138#endif
139}
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

◆ gemv_q4_k_q8_k_parallel_vnni()

void gemv_q4_k_q8_k_parallel_vnni ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K,
int  ith,
int  nth 
)

Definition at line 3771 of file gemm_kernels_q4k_q8k_vnni.c.

3776{
3777#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3778 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
3779 return;
3780 }
3781 if (ith < 0 || nth <= 0 || ith >= nth) {
3782 return;
3783 }
3784
3785 const int dr = (M + nth - 1) / nth;
3786 const int r0 = dr * ith;
3787 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
3788 if (r0 >= M) {
3789 return;
3790 }
3791
3792 const block_q4_K *blocks = (const block_q4_K *)W;
3793 const block_q8_K *x = (const block_q8_K *)x_q8;
3794 const int blocks_per_row = K / QK_K;
3795
3796 for (int row = r0; row < r1; ++row) {
3797 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
3798 float sum = 0.0f;
3799 for (int b = 0; b < blocks_per_row; ++b) {
3800 sum += dot_q4_k_q8_k_vnni_block(&w_row[b], &x[b]);
3801 }
3802 y[row] = sum;
3803 }
3804#else
3805 (void)y;
3806 (void)W;
3807 (void)x_q8;
3808 (void)M;
3809 (void)K;
3810 (void)ith;
3811 (void)nth;
3812#endif
3813}

References QK_K.

◆ gemv_q4_k_q8_k_ref()

void gemv_q4_k_q8_k_ref ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 201 of file gemm_kernels_q4k_q8k.c.

205{
206 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
207 return;
208 }
209
210 const block_q4_K *blocks = (const block_q4_K *)W;
211 const block_q8_K *x = (const block_q8_K *)x_q8;
212 const int blocks_per_row = K / QK_K;
213
214 for (int row = 0; row < M; ++row) {
215 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
216 y[row] = dot_q4_k_q8_k_ref(w_row, x, K);
217 }
218}
static float dot_q4_k_q8_k_ref(const block_q4_K *w, const block_q8_K *x, int k)

Referenced by gemm_q4_k_q8_k_ref(), gemv_q4_k_q8_k(), and gemv_q4_k_q8_k_vnni().

◆ gemv_q4_k_q8_k_vnni()

void gemv_q4_k_q8_k_vnni ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 3534 of file gemm_kernels_q4k_q8k_vnni.c.

3538{
3539#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3540 const char *fast_env = getenv("CK_ENABLE_Q4K_Q8K_VNNI_FAST");
3541 const int fast_disabled = fast_env && fast_env[0] && fast_env[0] == '0';
3542 if (!fast_disabled && !ck_strict_parity_enabled()) {
3543 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
3544 return;
3545 }
3546
3547 const block_q4_K *blocks = (const block_q4_K *)W;
3548 const block_q8_K *x = (const block_q8_K *)x_q8;
3549 const int blocks_per_row = K / QK_K;
3550
3551 for (int row = 0; row < M; ++row) {
3552 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
3553 float sum = 0.0f;
3554 for (int b = 0; b < blocks_per_row; ++b) {
3555 sum += dot_q4_k_q8_k_vnni_block(&w_row[b], &x[b]);
3556 }
3557 y[row] = sum;
3558 }
3559 return;
3560 }
3561#endif
3562
3563 /* Strict/debug parity keeps the llama-style scalar accumulation path.
3564 * Production AVX-512 hosts use VNNI by default; set
3565 * CK_ENABLE_Q4K_Q8K_VNNI_FAST=0 or CK_DEBUG_Q4K_Q8_REF=1 when attributing
3566 * borderline logit movement against scalar/reference behavior.
3567 */
3568 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
3569}
int ck_strict_parity_enabled(void)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

References ck_strict_parity_enabled(), gemv_q4_k_q8_k_ref(), and QK_K.

Referenced by gemv_q4_k_q8_k(), and gemv_q4_k_q8_k_amx().

◆ pack_q4_k_to_packed_meta()

void pack_q4_k_to_packed_meta ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 505 of file gemm_kernels_q4k_q8k_vnni.c.

506{
507 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
508 return;
509 }
510 const block_q4_K *in = (const block_q4_K *)src;
511 block_q4_K_packed_meta *out = (block_q4_K_packed_meta *)dst;
512 const int blocks_per_row = K / QK_K;
513 for (int n = 0; n < N; ++n) {
514 for (int b = 0; b < blocks_per_row; ++b) {
515 const block_q4_K *sb = in + (size_t)n * (size_t)blocks_per_row + (size_t)b;
516 block_q4_K_packed_meta *pb = out + (size_t)n * (size_t)blocks_per_row + (size_t)b;
517 pb->d = sb->d;
518 pb->dmin = sb->dmin;
519 unpack_q4_k_scales(sb->scales, pb->sc, pb->m);
520 memcpy(pb->qs, sb->qs, sizeof(pb->qs));
521 }
522 }
523}
static void unpack_q4_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
Unpack Q4_K sub-block scales and mins.
uint8_t scales[12]
uint8_t qs[256/2]

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ pack_q4_k_to_packed_meta_x16()

void pack_q4_k_to_packed_meta_x16 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 553 of file gemm_kernels_q4k_q8k_vnni.c.

554{
555 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
556 return;
557 }
558 const block_q4_K *in = (const block_q4_K *)src;
559 block_q4_K_packed_meta_x16 *out = (block_q4_K_packed_meta_x16 *)dst;
560 const int blocks_per_row = K / QK_K;
561 const int groups = (N + 15) / 16;
562 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
563
564 for (int g = 0; g < groups; ++g) {
565 const int n0 = g * 16;
566 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
567 for (int b = 0; b < blocks_per_row; ++b) {
568 block_q4_K_packed_meta_x16 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
569 pb->active = (uint8_t)active;
570 for (int lane = 0; lane < active; ++lane) {
571 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
572 pb->d[lane] = sb->d;
573 pb->dmin[lane] = sb->dmin;
574 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
575 memcpy(pb->qs[lane], sb->qs, sizeof(pb->qs[lane]));
576 }
577 }
578 }
579}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ pack_q4_k_to_packed_meta_x8()

void pack_q4_k_to_packed_meta_x8 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 525 of file gemm_kernels_q4k_q8k_vnni.c.

526{
527 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
528 return;
529 }
530 const block_q4_K *in = (const block_q4_K *)src;
531 block_q4_K_packed_meta_x8 *out = (block_q4_K_packed_meta_x8 *)dst;
532 const int blocks_per_row = K / QK_K;
533 const int groups = (N + 7) / 8;
534 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
535
536 for (int g = 0; g < groups; ++g) {
537 const int n0 = g * 8;
538 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
539 for (int b = 0; b < blocks_per_row; ++b) {
540 block_q4_K_packed_meta_x8 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
541 pb->active = (uint8_t)active;
542 for (int lane = 0; lane < active; ++lane) {
543 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
544 pb->d[lane] = sb->d;
545 pb->dmin[lane] = sb->dmin;
546 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
547 memcpy(pb->qs[lane], sb->qs, sizeof(pb->qs[lane]));
548 }
549 }
550 }
551}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

Referenced by ck_moe_q4k_llama_projection().

◆ pack_q4_k_to_packed_u8()

void pack_q4_k_to_packed_u8 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 479 of file gemm_kernels_q4k_q8k_vnni.c.

480{
481 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
482 return;
483 }
484 const block_q4_K *in = (const block_q4_K *)src;
485 block_q4_K_packed_u8 *out = (block_q4_K_packed_u8 *)dst;
486 const int blocks_per_row = K / QK_K;
487 for (int n = 0; n < N; ++n) {
488 for (int b = 0; b < blocks_per_row; ++b) {
489 const block_q4_K *sb = in + (size_t)n * (size_t)blocks_per_row + (size_t)b;
490 block_q4_K_packed_u8 *pb = out + (size_t)n * (size_t)blocks_per_row + (size_t)b;
491 pb->d = sb->d;
492 pb->dmin = sb->dmin;
493 unpack_q4_k_scales(sb->scales, pb->sc, pb->m);
494 for (int j = 0, q_offset = 0; j < QK_K; j += 64, q_offset += 32) {
495 const uint8_t *qs = &sb->qs[q_offset];
496 for (int l = 0; l < 32; ++l) {
497 pb->qs[j + l] = qs[l] & 0x0F;
498 pb->qs[j + 32 + l] = qs[l] >> 4;
499 }
500 }
501 }
502 }
503}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ pack_q4_k_to_packed_u8_x16()

void pack_q4_k_to_packed_u8_x16 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 581 of file gemm_kernels_q4k_q8k_vnni.c.

582{
583 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
584 return;
585 }
586 const block_q4_K *in = (const block_q4_K *)src;
587 block_q4_K_packed_u8_x16 *out = (block_q4_K_packed_u8_x16 *)dst;
588 const int blocks_per_row = K / QK_K;
589 const int groups = (N + 15) / 16;
590 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
591
592 for (int g = 0; g < groups; ++g) {
593 const int n0 = g * 16;
594 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
595 for (int b = 0; b < blocks_per_row; ++b) {
596 block_q4_K_packed_u8_x16 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
597 pb->active = (uint8_t)active;
598 for (int lane = 0; lane < active; ++lane) {
599 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
600 pb->d[lane] = sb->d;
601 pb->dmin[lane] = sb->dmin;
602 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
603 for (int j = 0, q_offset = 0; j < QK_K; j += 64, q_offset += 32) {
604 const uint8_t *qs = &sb->qs[q_offset];
605 for (int l = 0; l < 32; ++l) {
606 pb->qs[lane][j + l] = qs[l] & 0x0F;
607 pb->qs[lane][j + 32 + l] = qs[l] >> 4;
608 }
609 }
610 }
611 }
612 }
613}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ pack_q4_k_to_packed_vnni_x16()

void pack_q4_k_to_packed_vnni_x16 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 343 of file gemm_kernels_q4k_q8k_vnni.c.

344{
345 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
346 return;
347 }
348 const block_q4_K *in = (const block_q4_K *)src;
349 block_q4_K_packed_vnni_x16 *out = (block_q4_K_packed_vnni_x16 *)dst;
350 const int blocks_per_row = K / QK_K;
351 const int groups = (N + 15) / 16;
352 memset(out, 0,
353 (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
354
355 for (int group = 0; group < groups; ++group) {
356 const int n0 = group * 16;
357 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
358 for (int block = 0; block < blocks_per_row; ++block) {
359 block_q4_K_packed_vnni_x16 *packed =
360 out + (size_t)group * (size_t)blocks_per_row +
361 (size_t)block;
362 packed->active = (uint8_t)active;
363 for (int lane = 0; lane < active; ++lane) {
364 const block_q4_K *source =
365 in + (size_t)(n0 + lane) * (size_t)blocks_per_row +
366 (size_t)block;
367 uint8_t scales[8];
368 uint8_t mins[8];
369 packed->d[lane] = source->d;
370 packed->dmin[lane] = source->dmin;
371 unpack_q4_k_scales(source->scales, scales, mins);
372 for (int subblock = 0; subblock < 8; ++subblock) {
373 packed->sc[subblock][lane] = scales[subblock];
374 packed->m[subblock][lane] = mins[subblock];
375 }
376 for (int pair = 0; pair < QK_K / 64; ++pair) {
377 for (int segment = 0; segment < 8; ++segment) {
378 memcpy(packed->qs + (size_t)pair * 512u +
379 (size_t)segment * 64u +
380 (size_t)lane * 4u,
381 source->qs + (size_t)pair * 32u +
382 (size_t)segment * 4u,
383 4u);
384 }
385 }
386 }
387 }
388 }
389}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ pack_q4_k_to_packed_vnni_x8()

void pack_q4_k_to_packed_vnni_x8 ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 295 of file gemm_kernels_q4k_q8k_vnni.c.

296{
297 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
298 return;
299 }
300 const block_q4_K *in = (const block_q4_K *)src;
301 block_q4_K_packed_vnni_x8 *out = (block_q4_K_packed_vnni_x8 *)dst;
302 const int blocks_per_row = K / QK_K;
303 const int groups = (N + 7) / 8;
304 memset(out, 0,
305 (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
306
307 for (int group = 0; group < groups; ++group) {
308 const int n0 = group * 8;
309 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
310 for (int block = 0; block < blocks_per_row; ++block) {
311 block_q4_K_packed_vnni_x8 *packed =
312 out + (size_t)group * (size_t)blocks_per_row +
313 (size_t)block;
314 packed->active = (uint8_t)active;
315 for (int lane = 0; lane < active; ++lane) {
316 const block_q4_K *source =
317 in + (size_t)(n0 + lane) * (size_t)blocks_per_row +
318 (size_t)block;
319 uint8_t scales[8];
320 uint8_t mins[8];
321 packed->d[lane] = source->d;
322 packed->dmin[lane] = source->dmin;
323 unpack_q4_k_scales(source->scales, scales, mins);
324 for (int subblock = 0; subblock < 8; ++subblock) {
325 packed->sc[subblock][lane] = scales[subblock];
326 packed->m[subblock][lane] = mins[subblock];
327 }
328 for (int pair = 0; pair < QK_K / 64; ++pair) {
329 for (int segment = 0; segment < 8; ++segment) {
330 memcpy(packed->qs + (size_t)pair * 256u +
331 (size_t)segment * 32u +
332 (size_t)lane * 4u,
333 source->qs + (size_t)pair * 32u +
334 (size_t)segment * 4u,
335 4u);
336 }
337 }
338 }
339 }
340 }
341}

References block_q4_K::d, block_q4_K::dmin, QK_K, block_q4_K::qs, block_q4_K::scales, and unpack_q4_k_scales().

◆ q4_k_packed_meta_block_size()

size_t q4_k_packed_meta_block_size ( void  )

Definition at line 459 of file gemm_kernels_q4k_q8k_vnni.c.

460{
461 return sizeof(block_q4_K_packed_meta);
462}

◆ q4_k_packed_meta_x16_block_size()

size_t q4_k_packed_meta_x16_block_size ( void  )

Definition at line 469 of file gemm_kernels_q4k_q8k_vnni.c.

470{
471 return sizeof(block_q4_K_packed_meta_x16);
472}

◆ q4_k_packed_meta_x8_block_size()

size_t q4_k_packed_meta_x8_block_size ( void  )

Definition at line 464 of file gemm_kernels_q4k_q8k_vnni.c.

465{
466 return sizeof(block_q4_K_packed_meta_x8);
467}

Referenced by ck_moe_q4k_llama_projection_scratch_bytes().

◆ q4_k_packed_u8_block_size()

size_t q4_k_packed_u8_block_size ( void  )

Definition at line 454 of file gemm_kernels_q4k_q8k_vnni.c.

455{
456 return sizeof(block_q4_K_packed_u8);
457}

◆ q4_k_packed_u8_x16_block_size()

size_t q4_k_packed_u8_x16_block_size ( void  )

Definition at line 474 of file gemm_kernels_q4k_q8k_vnni.c.

475{
476 return sizeof(block_q4_K_packed_u8_x16);
477}

◆ q4_k_packed_vnni_x16_block_size()

size_t q4_k_packed_vnni_x16_block_size ( void  )

Definition at line 281 of file gemm_kernels_q4k_q8k_vnni.c.

282{
283 return sizeof(block_q4_K_packed_vnni_x16);
284}

◆ q4_k_packed_vnni_x8_block_size()

size_t q4_k_packed_vnni_x8_block_size ( void  )

Definition at line 267 of file gemm_kernels_q4k_q8k_vnni.c.

268{
269 return sizeof(block_q4_K_packed_vnni_x8);
270}

Referenced by ck_moe_swiglu_expert_forward_q4k_q5k_bucketed_impl().