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,
126 const void *activations,
float weight_scale)
129 assert(n >= 0 && n %
QK_NVFP4 == 0);
132 const int block_count = n /
QK_NVFP4;
133 const __m128i lut = _mm_loadu_si128(
135 const __m128i nibble_mask = _mm_set1_epi8(0x0f);
136 const __m256i ones = _mm256_set1_epi16(1);
137 __m256 accumulated = _mm256_setzero_ps();
139 for (
int block_index = 0; block_index < block_count; ++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));
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);
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);
172 const float q8_scale0 =
174 const float 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);
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;
290 const float *hidden,
const void *gate,
float gate_scale,
291 const void *up,
float up_scale,
const void *down,
float down_scale,
292 float *result,
int hidden_dim,
int intermediate_dim,
void *workspace)
294 uint8_t *cursor = (uint8_t *)workspace;
295 void *hidden_q8 = cursor;
298 float *gate_up = (
float *)cursor;
300 void *act_q8 = cursor;
304 intermediate_dim, hidden_dim);
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];
314 hidden_dim, intermediate_dim);
319 const float *hidden,
const int *indices,
const float *routing_weights,
320 const void *expert_gate,
const float *expert_gate_scales,
321 const void *expert_up,
const float *expert_up_scales,
322 const void *expert_down,
const float *expert_down_scales,
323 float *output,
int rows,
int hidden_dim,
int intermediate_dim,
324 int n_experts,
int top_k,
void *workspace,
size_t 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) {
339 2u * (
size_t)intermediate_dim *
sizeof(
float));
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 *
346 const size_t down_expert_bytes = (size_t)hidden_dim *
348 memset(output, 0, (
size_t)rows * (
size_t)hidden_dim *
sizeof(
float));
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) {
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];
378 const float *hidden,
const float *routed,
379 const void *shared_gate,
const float *shared_gate_scale,
380 const void *shared_up,
const float *shared_up_scale,
381 const void *shared_down,
const float *shared_down_scale,
382 float *output,
int rows,
int hidden_dim,
int intermediate_dim,
383 float combination_scale,
void *workspace,
size_t 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) {
395 2u * (
size_t)intermediate_dim *
sizeof(
float));
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));
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)
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)