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

Packed NVFP4 weight kernels for CPU inference. More...

#include <assert.h>
#include <math.h>
#include <stddef.h>
#include <stdint.h>
#include <string.h>
#include "ck_threadpool.h"
#include "ckernel_quant.h"

Go to the source code of this file.

Functions

static int ck_moe_swiglu_nvfp4_projection (const float *hidden, const void *gate, float gate_scale, const void *up, float up_scale, const void *down, float down_scale, float *result, int hidden_dim, int intermediate_dim, void *workspace)
 
static size_t ck_nvfp4_align64 (size_t value)
 
static void ck_nvfp4_gemv_rows (int begin, int end, void *opaque)
 
float ck_ue4m3_to_fp32 (uint8_t value)
 
static float ck_ue4m3_to_fp32_inline (uint8_t value)
 
void dequantize_row_nvfp4 (const void *weights, float *output, int k, float weight_scale)
 
void gemv_nvfp4_q8_0 (float *output, const void *weights, const float *weight_scales, const void *activations, int rows, int cols)
 
static void gemv_nvfp4_q8_0_uniform (float *output, const void *weights, float weight_scale, const void *activations, int rows, int cols)
 
int moe_swiglu_expert_forward_nvfp4_workspace (const float *hidden, const int *indices, const float *routing_weights, const void *expert_gate, const float *expert_gate_scales, const void *expert_up, const float *expert_up_scales, const void *expert_down, const float *expert_down_scales, float *output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void *workspace, size_t workspace_bytes)
 
size_t moe_swiglu_nvfp4_workspace_bytes (int hidden_dim, int intermediate_dim)
 
int moe_swiglu_shared_forward_nvfp4_workspace (const float *hidden, const float *routed, const void *shared_gate, const float *shared_gate_scale, const void *shared_up, const float *shared_up_scale, const void *shared_down, const float *shared_down_scale, float *output, int rows, int hidden_dim, int intermediate_dim, float combination_scale, void *workspace, size_t workspace_bytes)
 
void vec_dot_nvfp4_q8_0 (int n, float *output, const void *weights, const void *activations, float weight_scale)
 
void vec_dot_nvfp4_q8_0_ref (int n, float *output, const void *weights, const void *activations, float weight_scale)
 

Variables

static const int8_t ck_nvfp4_e2m1_x2 [16]
 
static const float ck_nvfp4_ue4m3 [128]
 

Detailed Description

Packed NVFP4 weight kernels for CPU inference.

The storage ABI keeps E2M1 weights and E4M3 block scales packed. The checkpoint's reciprocal tensor/expert scale is supplied separately so no weight expansion or scale re-encoding is required during conversion.

Definition in file gemm_kernels_nvfp4.c.

Function Documentation

◆ ck_moe_swiglu_nvfp4_projection()

static int ck_moe_swiglu_nvfp4_projection ( const float *  hidden,
const void *  gate,
float  gate_scale,
const void *  up,
float  up_scale,
const void *  down,
float  down_scale,
float *  result,
int  hidden_dim,
int  intermediate_dim,
void *  workspace 
)
static

Definition at line 289 of file gemm_kernels_nvfp4.c.

293{
294 uint8_t *cursor = (uint8_t *)workspace;
295 void *hidden_q8 = cursor;
296 cursor += ck_nvfp4_align64(
297 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)hidden_dim));
298 float *gate_up = (float *)cursor;
299 cursor += ck_nvfp4_align64(2u * (size_t)intermediate_dim * sizeof(float));
300 void *act_q8 = cursor;
301
302 quantize_row_q8_0(hidden, hidden_q8, hidden_dim);
303 gemv_nvfp4_q8_0_uniform(gate_up, gate, gate_scale, hidden_q8,
304 intermediate_dim, hidden_dim);
305 gemv_nvfp4_q8_0_uniform(gate_up + intermediate_dim, up, up_scale,
306 hidden_q8, intermediate_dim, hidden_dim);
307 for (int i = 0; i < intermediate_dim; ++i) {
308 const float value = gate_up[i];
309 gate_up[i] = (value / (1.0f + expf(-value))) *
310 gate_up[intermediate_dim + i];
311 }
312 quantize_row_q8_0(gate_up, act_q8, intermediate_dim);
313 gemv_nvfp4_q8_0_uniform(result, down, down_scale, act_q8,
314 hidden_dim, intermediate_dim);
315 return 0;
316}
@ CK_DT_Q8_0
static size_t ck_dtype_row_bytes(CKDataType dt, size_t n_elements)
Calculate total bytes for n_elements of given dtype.
void quantize_row_q8_0(const float *x, void *y, int k)
Quantize FP32 to Q8_0 format (scalar reference)
static size_t ck_nvfp4_align64(size_t value)
static void gemv_nvfp4_q8_0_uniform(float *output, const void *weights, float weight_scale, const void *activations, int rows, int cols)

References CK_DT_Q8_0, ck_dtype_row_bytes(), ck_nvfp4_align64(), gemv_nvfp4_q8_0_uniform(), and quantize_row_q8_0().

Referenced by moe_swiglu_expert_forward_nvfp4_workspace(), and moe_swiglu_shared_forward_nvfp4_workspace().

◆ ck_nvfp4_align64()

static size_t ck_nvfp4_align64 ( size_t  value)
static

◆ ck_nvfp4_gemv_rows()

static void ck_nvfp4_gemv_rows ( int  begin,
int  end,
void *  opaque 
)
static

Definition at line 232 of file gemm_kernels_nvfp4.c.

233{
234 const ck_nvfp4_gemv_rows_args_t *args =
235 (const ck_nvfp4_gemv_rows_args_t *)opaque;
236 for (int row = begin; row < end; ++row) {
238 args->cols, &args->output[row],
239 args->weights + (size_t)row * args->row_bytes,
240 args->activations, args->weight_scale);
241 }
242}
void vec_dot_nvfp4_q8_0(int n, float *output, const void *weights, const void *activations, float weight_scale)
uint32_t end
Definition utf8.c:215

References end, and vec_dot_nvfp4_q8_0().

Referenced by gemv_nvfp4_q8_0_uniform().

◆ ck_ue4m3_to_fp32()

float ck_ue4m3_to_fp32 ( uint8_t  value)

Definition at line 53 of file gemm_kernels_nvfp4.c.

54{
55 return ck_ue4m3_to_fp32_inline(value);
56}
static float ck_ue4m3_to_fp32_inline(uint8_t value)

References ck_ue4m3_to_fp32_inline().

◆ ck_ue4m3_to_fp32_inline()

static float ck_ue4m3_to_fp32_inline ( uint8_t  value)
inlinestatic

Definition at line 48 of file gemm_kernels_nvfp4.c.

49{
50 return ck_nvfp4_ue4m3[value & UINT8_C(0x7f)];
51}
static const float ck_nvfp4_ue4m3[128]

References ck_nvfp4_ue4m3.

Referenced by ck_ue4m3_to_fp32(), dequantize_row_nvfp4(), vec_dot_nvfp4_q8_0(), and vec_dot_nvfp4_q8_0_ref().

◆ dequantize_row_nvfp4()

void dequantize_row_nvfp4 ( const void *  weights,
float *  output,
int  k,
float  weight_scale 
)

Definition at line 58 of file gemm_kernels_nvfp4.c.

60{
61 assert(k >= 0 && k % QK_NVFP4 == 0);
62 const block_nvfp4 *blocks = (const block_nvfp4 *)weights;
63 const int block_count = k / QK_NVFP4;
64
65 for (int block_index = 0; block_index < block_count; ++block_index) {
66 const block_nvfp4 *block = &blocks[block_index];
67 for (int sub = 0; sub < QK_NVFP4 / QK_NVFP4_SUB; ++sub) {
68 const float scale = ck_ue4m3_to_fp32_inline(block->d[sub]) * weight_scale;
69 const uint8_t *packed = &block->qs[sub * (QK_NVFP4_SUB / 2)];
70 float *dst = &output[block_index * QK_NVFP4 + sub * QK_NVFP4_SUB];
71 for (int lane = 0; lane < QK_NVFP4_SUB / 2; ++lane) {
72 const uint8_t pair = packed[lane];
73 dst[lane] = scale * (float)ck_nvfp4_e2m1_x2[pair & 0x0f];
74 dst[lane + QK_NVFP4_SUB / 2] =
75 scale * (float)ck_nvfp4_e2m1_x2[pair >> 4];
76 }
77 }
78 }
79}
#define QK_NVFP4_SUB
#define QK_NVFP4
static const int8_t ck_nvfp4_e2m1_x2[16]
uint8_t qs[64/2]
uint8_t d[64/16]

References ck_nvfp4_e2m1_x2, ck_ue4m3_to_fp32_inline(), block_nvfp4::d, QK_NVFP4, QK_NVFP4_SUB, and block_nvfp4::qs.

◆ gemv_nvfp4_q8_0()

void gemv_nvfp4_q8_0 ( float *  output,
const void *  weights,
const float *  weight_scales,
const void *  activations,
int  rows,
int  cols 
)

Definition at line 203 of file gemm_kernels_nvfp4.c.

206{
207 assert(rows >= 0 && cols >= 0 && cols % QK_NVFP4 == 0);
208 const size_t row_bytes = (size_t)(cols / QK_NVFP4) * sizeof(block_nvfp4);
209 const uint8_t *weight_bytes = (const uint8_t *)weights;
210 for (int row = 0; row < rows; ++row) {
211 const float scale = weight_scales ? weight_scales[row] : 1.0f;
212 vec_dot_nvfp4_q8_0(cols, &output[row],
213 weight_bytes + (size_t)row * row_bytes,
214 activations, scale);
215 }
216}

References QK_NVFP4, and vec_dot_nvfp4_q8_0().

◆ gemv_nvfp4_q8_0_uniform()

static void gemv_nvfp4_q8_0_uniform ( float *  output,
const void *  weights,
float  weight_scale,
const void *  activations,
int  rows,
int  cols 
)
static

Definition at line 244 of file gemm_kernels_nvfp4.c.

247{
248 ck_nvfp4_gemv_rows_args_t args = {
249 .output = output,
250 .weights = (const uint8_t *)weights,
251 .activations = activations,
252 .weight_scale = weight_scale,
253 .row_bytes = ck_dtype_row_bytes(CK_DT_NVFP4, (size_t)cols),
254 .cols = cols,
255 };
256 ck_threadpool_t *pool = ck_threadpool_global();
257 int active_threads = pool ? ck_threadpool_n_threads(pool) : 1;
258 if (active_threads > rows) {
259 active_threads = rows;
260 }
261 if (!pool || active_threads <= 1 || rows < 64) {
262 ck_nvfp4_gemv_rows(0, rows, &args);
263 return;
264 }
265
266 int grain = rows / (active_threads * 4);
267 if (grain < 8) {
268 grain = 8;
269 }
271 pool, active_threads, 0, rows, grain, ck_nvfp4_gemv_rows, &args);
272}
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)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
@ CK_DT_NVFP4
static void ck_nvfp4_gemv_rows(int begin, int end, void *opaque)

References CK_DT_NVFP4, ck_dtype_row_bytes(), ck_nvfp4_gemv_rows(), ck_threadpool_global(), ck_threadpool_n_threads(), and ck_threadpool_parallel_for_n().

Referenced by ck_moe_swiglu_nvfp4_projection().

◆ moe_swiglu_expert_forward_nvfp4_workspace()

int moe_swiglu_expert_forward_nvfp4_workspace ( const float *  hidden,
const int *  indices,
const float *  routing_weights,
const void *  expert_gate,
const float *  expert_gate_scales,
const void *  expert_up,
const float *  expert_up_scales,
const void *  expert_down,
const float *  expert_down_scales,
float *  output,
int  rows,
int  hidden_dim,
int  intermediate_dim,
int  n_experts,
int  top_k,
void *  workspace,
size_t  workspace_bytes 
)

Definition at line 318 of file gemm_kernels_nvfp4.c.

325{
326 const size_t required = moe_swiglu_nvfp4_workspace_bytes(
327 hidden_dim, intermediate_dim);
328 if (!hidden || !indices || !routing_weights || !expert_gate ||
329 !expert_gate_scales || !expert_up || !expert_up_scales ||
330 !expert_down || !expert_down_scales || !output || !workspace ||
331 required == 0 || workspace_bytes < required || rows <= 0 ||
332 n_experts <= 0 || top_k <= 0 || top_k > n_experts) {
333 return -1;
334 }
335
336 const size_t hidden_q8_bytes = ck_nvfp4_align64(
337 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)hidden_dim));
338 const size_t gate_up_bytes = ck_nvfp4_align64(
339 2u * (size_t)intermediate_dim * sizeof(float));
340 const size_t act_q8_bytes = ck_nvfp4_align64(
341 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)intermediate_dim));
342 float *expert_output = (float *)((uint8_t *)workspace + hidden_q8_bytes +
343 gate_up_bytes + act_q8_bytes);
344 const size_t up_expert_bytes = (size_t)intermediate_dim *
345 ck_dtype_row_bytes(CK_DT_NVFP4, (size_t)hidden_dim);
346 const size_t down_expert_bytes = (size_t)hidden_dim *
347 ck_dtype_row_bytes(CK_DT_NVFP4, (size_t)intermediate_dim);
348 memset(output, 0, (size_t)rows * (size_t)hidden_dim * sizeof(float));
349
350 for (int row = 0; row < rows; ++row) {
351 const float *x = hidden + (size_t)row * (size_t)hidden_dim;
352 float *y = output + (size_t)row * (size_t)hidden_dim;
353 for (int slot = 0; slot < top_k; ++slot) {
354 const size_t route = (size_t)row * (size_t)top_k + (size_t)slot;
355 const int expert = indices[route];
356 if (expert < 0 || expert >= n_experts) {
357 return -2;
358 }
360 x,
361 (const uint8_t *)expert_gate + (size_t)expert * up_expert_bytes,
362 expert_gate_scales[expert],
363 (const uint8_t *)expert_up + (size_t)expert * up_expert_bytes,
364 expert_up_scales[expert],
365 (const uint8_t *)expert_down + (size_t)expert * down_expert_bytes,
366 expert_down_scales[expert], expert_output,
367 hidden_dim, intermediate_dim, workspace);
368 const float route_weight = routing_weights[route];
369 for (int h = 0; h < hidden_dim; ++h) {
370 y[h] += route_weight * expert_output[h];
371 }
372 }
373 }
374 return 0;
375}
static int ck_moe_swiglu_nvfp4_projection(const float *hidden, const void *gate, float gate_scale, const void *up, float up_scale, const void *down, float down_scale, float *result, int hidden_dim, int intermediate_dim, void *workspace)
size_t moe_swiglu_nvfp4_workspace_bytes(int hidden_dim, int intermediate_dim)

References CK_DT_NVFP4, CK_DT_Q8_0, ck_dtype_row_bytes(), ck_moe_swiglu_nvfp4_projection(), ck_nvfp4_align64(), and moe_swiglu_nvfp4_workspace_bytes().

◆ moe_swiglu_nvfp4_workspace_bytes()

size_t moe_swiglu_nvfp4_workspace_bytes ( int  hidden_dim,
int  intermediate_dim 
)

Definition at line 274 of file gemm_kernels_nvfp4.c.

275{
276 if (hidden_dim <= 0 || intermediate_dim <= 0 || hidden_dim % 64 != 0 ||
277 intermediate_dim % 64 != 0) {
278 return 0;
279 }
280 size_t bytes = ck_nvfp4_align64(
281 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)hidden_dim));
282 bytes += ck_nvfp4_align64(2u * (size_t)intermediate_dim * sizeof(float));
283 bytes += ck_nvfp4_align64(
284 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)intermediate_dim));
285 bytes += ck_nvfp4_align64((size_t)hidden_dim * sizeof(float));
286 return bytes;
287}

References CK_DT_Q8_0, ck_dtype_row_bytes(), and ck_nvfp4_align64().

Referenced by moe_swiglu_expert_forward_nvfp4_workspace(), and moe_swiglu_shared_forward_nvfp4_workspace().

◆ moe_swiglu_shared_forward_nvfp4_workspace()

int moe_swiglu_shared_forward_nvfp4_workspace ( const float *  hidden,
const float *  routed,
const void *  shared_gate,
const float *  shared_gate_scale,
const void *  shared_up,
const float *  shared_up_scale,
const void *  shared_down,
const float *  shared_down_scale,
float *  output,
int  rows,
int  hidden_dim,
int  intermediate_dim,
float  combination_scale,
void *  workspace,
size_t  workspace_bytes 
)

Definition at line 377 of file gemm_kernels_nvfp4.c.

384{
385 const size_t required = moe_swiglu_nvfp4_workspace_bytes(
386 hidden_dim, intermediate_dim);
387 if (!hidden || !shared_gate || !shared_gate_scale || !shared_up ||
388 !shared_up_scale || !shared_down || !shared_down_scale || !output ||
389 !workspace || required == 0 || workspace_bytes < required || rows <= 0) {
390 return -1;
391 }
392 const size_t hidden_q8_bytes = ck_nvfp4_align64(
393 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)hidden_dim));
394 const size_t gate_up_bytes = ck_nvfp4_align64(
395 2u * (size_t)intermediate_dim * sizeof(float));
396 const size_t act_q8_bytes = ck_nvfp4_align64(
397 ck_dtype_row_bytes(CK_DT_Q8_0, (size_t)intermediate_dim));
398 float *shared_output = (float *)((uint8_t *)workspace + hidden_q8_bytes +
399 gate_up_bytes + act_q8_bytes);
400 for (int row = 0; row < rows; ++row) {
402 hidden + (size_t)row * (size_t)hidden_dim,
403 shared_gate, shared_gate_scale[0], shared_up, shared_up_scale[0],
404 shared_down, shared_down_scale[0], shared_output,
405 hidden_dim, intermediate_dim, workspace);
406 float *y = output + (size_t)row * (size_t)hidden_dim;
407 const float *route = routed ? routed + (size_t)row * (size_t)hidden_dim : NULL;
408 for (int h = 0; h < hidden_dim; ++h) {
409 y[h] = combination_scale *
410 (shared_output[h] + (route ? route[h] : 0.0f));
411 }
412 }
413 return 0;
414}

References CK_DT_Q8_0, ck_dtype_row_bytes(), ck_moe_swiglu_nvfp4_projection(), ck_nvfp4_align64(), and moe_swiglu_nvfp4_workspace_bytes().

◆ vec_dot_nvfp4_q8_0()

void vec_dot_nvfp4_q8_0 ( int  n,
float *  output,
const void *  weights,
const void *  activations,
float  weight_scale 
)

Definition at line 125 of file gemm_kernels_nvfp4.c.

127{
128#if defined(__AVX2__)
129 assert(n >= 0 && n % QK_NVFP4 == 0);
130 const block_nvfp4 *w = (const block_nvfp4 *)weights;
131 const block_q8_0 *x = (const block_q8_0 *)activations;
132 const int block_count = n / QK_NVFP4;
133 const __m128i lut = _mm_loadu_si128(
134 (const __m128i *)ck_nvfp4_e2m1_x2);
135 const __m128i nibble_mask = _mm_set1_epi8(0x0f);
136 const __m256i ones = _mm256_set1_epi16(1);
137 __m256 accumulated = _mm256_setzero_ps();
138
139 for (int block_index = 0; block_index < block_count; ++block_index) {
140 const block_nvfp4 *block = &w[block_index];
141 const __m128i packed01 = _mm_loadu_si128(
142 (const __m128i *)(block->qs + 0));
143 const __m128i packed23 = _mm_loadu_si128(
144 (const __m128i *)(block->qs + 16));
145 const __m128i low01 = _mm_shuffle_epi8(
146 lut, _mm_and_si128(packed01, nibble_mask));
147 const __m128i high01 = _mm_shuffle_epi8(
148 lut, _mm_and_si128(_mm_srli_epi16(packed01, 4), nibble_mask));
149 const __m128i low23 = _mm_shuffle_epi8(
150 lut, _mm_and_si128(packed23, nibble_mask));
151 const __m128i high23 = _mm_shuffle_epi8(
152 lut, _mm_and_si128(_mm_srli_epi16(packed23, 4), nibble_mask));
153
154 __m256i values01 = _mm256_castsi128_si256(
155 _mm_unpacklo_epi64(low01, high01));
156 values01 = _mm256_inserti128_si256(
157 values01, _mm_unpackhi_epi64(low01, high01), 1);
158 __m256i values23 = _mm256_castsi128_si256(
159 _mm_unpacklo_epi64(low23, high23));
160 values23 = _mm256_inserti128_si256(
161 values23, _mm_unpackhi_epi64(low23, high23), 1);
162
163 const __m256i q8_01 = _mm256_loadu_si256(
164 (const __m256i *)x[2 * block_index + 0].qs);
165 const __m256i q8_23 = _mm256_loadu_si256(
166 (const __m256i *)x[2 * block_index + 1].qs);
167 const __m256i dot01 = _mm256_madd_epi16(
168 ck_nvfp4_mul_add_i8_avx2(values01, q8_01), ones);
169 const __m256i dot23 = _mm256_madd_epi16(
170 ck_nvfp4_mul_add_i8_avx2(values23, q8_23), ones);
171
172 const float q8_scale0 =
173 CK_FP16_TO_FP32(x[2 * block_index + 0].d);
174 const float q8_scale1 =
175 CK_FP16_TO_FP32(x[2 * block_index + 1].d);
176 const float scale0 = ck_ue4m3_to_fp32_inline(block->d[0]) * q8_scale0;
177 const float scale1 = ck_ue4m3_to_fp32_inline(block->d[1]) * q8_scale0;
178 const float scale2 = ck_ue4m3_to_fp32_inline(block->d[2]) * q8_scale1;
179 const float scale3 = ck_ue4m3_to_fp32_inline(block->d[3]) * q8_scale1;
180 const __m256 scales01 = _mm256_insertf128_ps(
181 _mm256_castps128_ps256(_mm_set1_ps(scale0)),
182 _mm_set1_ps(scale1), 1);
183 const __m256 scales23 = _mm256_insertf128_ps(
184 _mm256_castps128_ps256(_mm_set1_ps(scale2)),
185 _mm_set1_ps(scale3), 1);
186 accumulated = _mm256_fmadd_ps(
187 scales01, _mm256_cvtepi32_ps(dot01), accumulated);
188 accumulated = _mm256_fmadd_ps(
189 scales23, _mm256_cvtepi32_ps(dot23), accumulated);
190 }
191
192 __m128 sum4 = _mm_add_ps(
193 _mm256_castps256_ps128(accumulated),
194 _mm256_extractf128_ps(accumulated, 1));
195 sum4 = _mm_hadd_ps(sum4, sum4);
196 sum4 = _mm_hadd_ps(sum4, sum4);
197 *output = _mm_cvtss_f32(sum4) * weight_scale;
198#else
199 vec_dot_nvfp4_q8_0_ref(n, output, weights, activations, weight_scale);
200#endif
201}
#define CK_FP16_TO_FP32(x)
void vec_dot_nvfp4_q8_0_ref(int n, float *output, const void *weights, const void *activations, float weight_scale)

References CK_FP16_TO_FP32, ck_nvfp4_e2m1_x2, ck_ue4m3_to_fp32_inline(), block_nvfp4::d, QK_NVFP4, block_nvfp4::qs, and vec_dot_nvfp4_q8_0_ref().

Referenced by ck_nvfp4_gemv_rows(), and gemv_nvfp4_q8_0().

◆ vec_dot_nvfp4_q8_0_ref()

void vec_dot_nvfp4_q8_0_ref ( int  n,
float *  output,
const void *  weights,
const void *  activations,
float  weight_scale 
)

Definition at line 81 of file gemm_kernels_nvfp4.c.

83{
84 assert(n >= 0 && n % QK_NVFP4 == 0);
85 const block_nvfp4 *w = (const block_nvfp4 *)weights;
86 const block_q8_0 *x = (const block_q8_0 *)activations;
87 const int block_count = n / QK_NVFP4;
88 float sum = 0.0f;
89
90 for (int block_index = 0; block_index < block_count; ++block_index) {
91 for (int sub = 0; sub < QK_NVFP4 / QK_NVFP4_SUB; ++sub) {
92 const int q8_block = sub / 2;
93 const int q8_offset = (sub % 2) * QK_NVFP4_SUB;
94 const float scale = ck_ue4m3_to_fp32_inline(w[block_index].d[sub]) *
95 CK_FP16_TO_FP32(x[2 * block_index + q8_block].d) *
96 weight_scale;
97 const uint8_t *packed =
98 &w[block_index].qs[sub * (QK_NVFP4_SUB / 2)];
99 const int8_t *q8 = &x[2 * block_index + q8_block].qs[q8_offset];
100 int integer_sum = 0;
101 for (int lane = 0; lane < QK_NVFP4_SUB / 2; ++lane) {
102 const uint8_t pair = packed[lane];
103 integer_sum += (int)q8[lane] *
104 (int)ck_nvfp4_e2m1_x2[pair & 0x0f];
105 integer_sum += (int)q8[lane + QK_NVFP4_SUB / 2] *
106 (int)ck_nvfp4_e2m1_x2[pair >> 4];
107 }
108 sum += scale * (float)integer_sum;
109 }
110 }
111 *output = sum;
112}
int8_t qs[32]

References CK_FP16_TO_FP32, ck_nvfp4_e2m1_x2, ck_ue4m3_to_fp32_inline(), QK_NVFP4, QK_NVFP4_SUB, block_q8_0::qs, and block_nvfp4::qs.

Referenced by vec_dot_nvfp4_q8_0().

Variable Documentation

◆ ck_nvfp4_e2m1_x2

const int8_t ck_nvfp4_e2m1_x2[16]
static
Initial value:
= {
0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12,
}

Definition at line 24 of file gemm_kernels_nvfp4.c.

24 {
25 0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12,
26};

Referenced by dequantize_row_nvfp4(), vec_dot_nvfp4_q8_0(), and vec_dot_nvfp4_q8_0_ref().

◆ ck_nvfp4_ue4m3

const float ck_nvfp4_ue4m3[128]
static
Initial value:
= {
0.0f, 0.0009765625f, 0.001953125f, 0.0029296875f, 0.00390625f, 0.0048828125f, 0.005859375f, 0.0068359375f,
0.0078125f, 0.0087890625f, 0.009765625f, 0.0107421875f, 0.01171875f, 0.0126953125f, 0.013671875f, 0.0146484375f,
0.015625f, 0.017578125f, 0.01953125f, 0.021484375f, 0.0234375f, 0.025390625f, 0.02734375f, 0.029296875f,
0.03125f, 0.03515625f, 0.0390625f, 0.04296875f, 0.046875f, 0.05078125f, 0.0546875f, 0.05859375f,
0.0625f, 0.0703125f, 0.078125f, 0.0859375f, 0.09375f, 0.1015625f, 0.109375f, 0.1171875f,
0.125f, 0.140625f, 0.15625f, 0.171875f, 0.1875f, 0.203125f, 0.21875f, 0.234375f,
0.25f, 0.28125f, 0.3125f, 0.34375f, 0.375f, 0.40625f, 0.4375f, 0.46875f,
0.5f, 0.5625f, 0.625f, 0.6875f, 0.75f, 0.8125f, 0.875f, 0.9375f,
1.0f, 1.125f, 1.25f, 1.375f, 1.5f, 1.625f, 1.75f, 1.875f,
2.0f, 2.25f, 2.5f, 2.75f, 3.0f, 3.25f, 3.5f, 3.75f,
4.0f, 4.5f, 5.0f, 5.5f, 6.0f, 6.5f, 7.0f, 7.5f,
8.0f, 9.0f, 10.0f, 11.0f, 12.0f, 13.0f, 14.0f, 15.0f,
16.0f, 18.0f, 20.0f, 22.0f, 24.0f, 26.0f, 28.0f, 30.0f,
32.0f, 36.0f, 40.0f, 44.0f, 48.0f, 52.0f, 56.0f, 60.0f,
64.0f, 72.0f, 80.0f, 88.0f, 96.0f, 104.0f, 112.0f, 120.0f,
128.0f, 144.0f, 160.0f, 176.0f, 192.0f, 208.0f, 224.0f, 0.0f,
}

Definition at line 29 of file gemm_kernels_nvfp4.c.

29 {
30 0.0f, 0.0009765625f, 0.001953125f, 0.0029296875f, 0.00390625f, 0.0048828125f, 0.005859375f, 0.0068359375f,
31 0.0078125f, 0.0087890625f, 0.009765625f, 0.0107421875f, 0.01171875f, 0.0126953125f, 0.013671875f, 0.0146484375f,
32 0.015625f, 0.017578125f, 0.01953125f, 0.021484375f, 0.0234375f, 0.025390625f, 0.02734375f, 0.029296875f,
33 0.03125f, 0.03515625f, 0.0390625f, 0.04296875f, 0.046875f, 0.05078125f, 0.0546875f, 0.05859375f,
34 0.0625f, 0.0703125f, 0.078125f, 0.0859375f, 0.09375f, 0.1015625f, 0.109375f, 0.1171875f,
35 0.125f, 0.140625f, 0.15625f, 0.171875f, 0.1875f, 0.203125f, 0.21875f, 0.234375f,
36 0.25f, 0.28125f, 0.3125f, 0.34375f, 0.375f, 0.40625f, 0.4375f, 0.46875f,
37 0.5f, 0.5625f, 0.625f, 0.6875f, 0.75f, 0.8125f, 0.875f, 0.9375f,
38 1.0f, 1.125f, 1.25f, 1.375f, 1.5f, 1.625f, 1.75f, 1.875f,
39 2.0f, 2.25f, 2.5f, 2.75f, 3.0f, 3.25f, 3.5f, 3.75f,
40 4.0f, 4.5f, 5.0f, 5.5f, 6.0f, 6.5f, 7.0f, 7.5f,
41 8.0f, 9.0f, 10.0f, 11.0f, 12.0f, 13.0f, 14.0f, 15.0f,
42 16.0f, 18.0f, 20.0f, 22.0f, 24.0f, 26.0f, 28.0f, 30.0f,
43 32.0f, 36.0f, 40.0f, 44.0f, 48.0f, 52.0f, 56.0f, 60.0f,
44 64.0f, 72.0f, 80.0f, 88.0f, 96.0f, 104.0f, 112.0f, 120.0f,
45 128.0f, 144.0f, 160.0f, 176.0f, 192.0f, 208.0f, 224.0f, 0.0f,
46};

Referenced by ck_ue4m3_to_fp32_inline().