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

AMX (Advanced Matrix Extensions) GEMM kernels. More...

#include <stdint.h>
#include <stddef.h>
#include <string.h>
#include <stdbool.h>
#include "ckernel_quant.h"

Go to the source code of this file.

Functions

bool amx_available (void)
 
void gemv_q4_k_q8_k_amx (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q4_k_q8_k_avx (float *y, const void *W, const void *x_q8, int M, int K)
 
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_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)
 

Detailed Description

AMX (Advanced Matrix Extensions) GEMM kernels.

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

Intel AMX provides dedicated matrix multiply hardware:

  • 8 tile registers (TMM0-TMM7), each up to 1KB
  • TDPBSSD: INT8 signed dot product (A signed, B signed)
  • TDPBSUD: INT8 mixed sign (A signed, B unsigned)
  • TDPBUSD: INT8 mixed sign (A unsigned, B signed)
  • TDPBUUD: INT8 unsigned dot product
  • TDPBF16PS: BF16 dot product to FP32

Tile dimensions:

  • Max: 16 rows x 64 bytes (1024 bytes per tile)
  • For INT8: 16x64 elements
  • For BF16: 16x32 elements

Performance:

  • AMX INT8: ~2000 INT8 ops/cycle (vs ~256 for AVX-512 VNNI)
  • AMX BF16: ~1000 BF16 ops/cycle
  • Expected 8-16x speedup over AVX-512 for large GEMM

Requirements:

  • Sapphire Rapids or newer (4th Gen Xeon)
  • Linux kernel 5.16+ with AMX support
  • Compiler: GCC 11+, Clang 12+, ICX 2022+

Definition in file gemm_kernels_amx.c.

Function Documentation

◆ amx_available()

bool amx_available ( void  )

Definition at line 271 of file gemm_kernels_amx.c.

271 {
272 return false;
273}

◆ gemv_q4_k_q8_k_amx()

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

Definition at line 256 of file gemm_kernels_amx.c.

256 {
257 /* No AMX support - cascade through fallbacks: AVX-512 VNNI → AVX2 → AVX → ref */
258#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
259 gemv_q4_k_q8_k_vnni(y, W, x_q8, M, K);
260#elif defined(__AVX2__)
261 gemv_q4_k_q8_k_avx2(y, W, x_q8, M, K);
262#elif defined(__AVX__)
263 gemv_q4_k_q8_k_avx(y, W, x_q8, M, K);
264#else
265 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
266#endif
267}
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_vnni(float *y, const void *W, const void *x_q8, int M, int K)
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_avx(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q4_k_q8_k_avx(), gemv_q4_k_q8_k_avx2(), gemv_q4_k_q8_k_ref(), and gemv_q4_k_q8_k_vnni().

◆ gemv_q4_k_q8_k_avx()

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

Definition at line 251 of file gemm_kernels_q4k_avx.c.

255{
256 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
257}
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q4_k_q8_k_ref().

Referenced by gemv_q4_k_q8_k_amx().

◆ 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}
#define QK_K
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q4_k_q8_k_ref(), and QK_K.

Referenced by gemv_q4_k_q8_k_amx().

◆ 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 gemv_q4_k_q8_k_amx().

◆ 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_amx().