48#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
52#define CK_DELTANET_MAX_STACK_DIM 4096
53#define CK_DELTANET_LLAMA_CHUNK_SIZE 64
54#define CK_DELTANET_LLAMA_CHUNK_MAX_DIM 256
56#if defined(__GNUC__) || defined(__clang__)
57#define CK_DELTANET_NOINLINE __attribute__((noinline))
59#define CK_DELTANET_NOINLINE
76 "HARD KERNEL CONTRACT FAULT: llama.cpp DeltaNet requires "
77 "expf from libm.so.6\n");
84 return 1.0f / (1.0f + expf(-x));
93#if defined(__AVX512F__)
94typedef __m512 (*ck_deltanet_sleef_expf16_fn)(__m512);
95static ck_deltanet_sleef_expf16_fn ck_deltanet_pytorch_expf16 = NULL;
96static void *ck_deltanet_sleef_handle = NULL;
106 const char *mkl_library = getenv(
"CK_MKL_LIBRARY");
107 if (mkl_library && *mkl_library) {
119#if defined(__AVX512F__)
120 const char *library = getenv(
"CK_SLEEF_LIBRARY");
121 if (library && *library) {
122 ck_deltanet_sleef_handle = dlopen(library, RTLD_NOW | RTLD_LOCAL);
123 if (ck_deltanet_sleef_handle) {
124 ck_deltanet_pytorch_expf16 =
125 (ck_deltanet_sleef_expf16_fn)dlsym(
126 ck_deltanet_sleef_handle,
"Sleef_expf16_u10");
129 ck_deltanet_pytorch_expf16 =
130 (ck_deltanet_sleef_expf16_fn)dlsym(
147 "HARD KERNEL CONTRACT FAULT: PyTorch DeltaNet requires "
148 "MKL vsExp; set CK_MKL_LIBRARY\n");
154#if defined(__AVX512F__)
155 if (ck_deltanet_pytorch_expf16) {
156 const __m512 one = _mm512_set1_ps(1.0f);
157 for (; h + 15 < num_heads; h += 16) {
158 const __m512 bv = _mm512_loadu_ps(beta + h);
159 const __m512 beta_exp = ck_deltanet_pytorch_expf16(
160 _mm512_sub_ps(_mm512_setzero_ps(), bv));
161 const __m512 beta_sigmoid = _mm512_div_ps(
162 one, _mm512_add_ps(one, beta_exp));
163 float beta_lanes[16];
164 _mm512_storeu_ps(beta_lanes, beta_sigmoid);
165 for (
int lane = 0; lane < 16; ++lane) {
172 for (; h < num_heads; ++h) {
184 if (!g || !beta || !gate_values || !beta_values || num_heads <= 0 ||
189 g, beta, gate_values, beta_values, num_heads);
197 const float *state_in,
205#if defined(__AVX512F__)
206static inline float ck_deltanet_gcc_reduce_add_ps(__m512 value)
208 const __m256 hi = _mm512_extractf32x8_ps(value, 1);
209 const __m256 lo = _mm512_castps512_ps256(value);
210 const __m256 sum8 = _mm256_add_ps(hi, lo);
211 const __m128 hi4 = _mm256_extractf128_ps(sum8, 1);
212 const __m128 lo4 = _mm256_castps256_ps128(sum8);
213 const __m128 sum4 = _mm_add_ps(hi4, lo4);
214 const __m128 swapped = _mm_shuffle_ps(sum4, sum4, _MM_SHUFFLE(1, 0, 3, 2));
215 const __m128 sum2 = _mm_add_ps(sum4, swapped);
216 return _mm_cvtss_f32(_mm_add_ss(sum2, _mm_shuffle_ps(sum2, sum2, 1)));
221 const float *x,
const float *y,
int n)
223#if defined(__AVX512F__)
224 const int np = n & ~63;
226 _mm512_setzero_ps(), _mm512_setzero_ps(),
227 _mm512_setzero_ps(), _mm512_setzero_ps()
229 for (
int i = 0; i < np; i += 64) {
230 for (
int j = 0; j < 4; ++j) {
231 const __m512 xv = _mm512_loadu_ps(x + i + j * 16);
232 const __m512 yv = _mm512_loadu_ps(y + i + j * 16);
233 sum[j] = _mm512_fmadd_ps(xv, yv, sum[j]);
236 sum[0] = _mm512_add_ps(sum[0], sum[2]);
237 sum[1] = _mm512_add_ps(sum[1], sum[3]);
238 sum[0] = _mm512_add_ps(sum[0], sum[1]);
239 float result = ck_deltanet_gcc_reduce_add_ps(sum[0]);
241 const int np = n & ~31;
243 _mm256_setzero_ps(), _mm256_setzero_ps(),
244 _mm256_setzero_ps(), _mm256_setzero_ps()
246 for (
int i = 0; i < np; i += 32) {
247 for (
int j = 0; j < 4; ++j) {
248 const __m256 xv = _mm256_loadu_ps(x + i + j * 8);
249 const __m256 yv = _mm256_loadu_ps(y + i + j * 8);
251 sum[j] = _mm256_fmadd_ps(xv, yv, sum[j]);
253 sum[j] = _mm256_add_ps(_mm256_mul_ps(xv, yv), sum[j]);
257 sum[0] = _mm256_add_ps(sum[0], sum[2]);
258 sum[1] = _mm256_add_ps(sum[1], sum[3]);
259 sum[0] = _mm256_add_ps(sum[0], sum[1]);
260 const __m128 halves = _mm_add_ps(
261 _mm256_castps256_ps128(sum[0]), _mm256_extractf128_ps(sum[0], 1));
262 const __m128 pairs = _mm_hadd_ps(halves, halves);
263 float result = _mm_cvtss_f32(_mm_hadd_ps(pairs, pairs));
265 for (
int i = np; i < n; ++i) {
266 result += x[i] * y[i];
271static inline float ck_deltanet_llama_scale(
int state_dim)
274 const __m128 dim = _mm_set_ss((
float) state_dim);
275 return _mm_cvtss_f32(_mm_div_ss(_mm_set_ss(1.0f), _mm_sqrt_ss(dim)));
277 return 1.0f / sqrtf((
float) state_dim);
281static inline __m256 ck_deltanet_fmadd8(__m256 a, __m256 b, __m256 acc)
284 return _mm256_fmadd_ps(a, b, acc);
286 return _mm256_add_ps(_mm256_mul_ps(a, b), acc);
303static void gated_deltanet_llama_chunk64_head(
309 const float *state_in,
319 const int group = head % group_count;
320 const size_t qk_row_stride = (size_t)group_count * (
size_t)state_dim;
321 const size_t value_row_stride = (size_t)num_heads * (
size_t)state_dim;
322 const size_t gate_row_stride = (size_t)num_heads;
323 const size_t state_count = (size_t)state_dim * (
size_t)state_dim;
324 const float scale = 1.0f / sqrtf((
float)state_dim);
329 float transform[
C *
C];
339 float *state = state_out + (size_t)head * state_count;
340 const float *initial = state_in + (size_t)head * state_count;
341 if (state != initial) {
342 for (
size_t i = 0; i < state_count; ++i) {
343 state[i] = initial[i];
347 for (
int chunk_start = 0; chunk_start < rows; chunk_start +=
C) {
348 const int valid = rows - chunk_start <
C ? rows - chunk_start :
C;
350 for (
int i = 0; i <
C; ++i) {
351 const int token = chunk_start + i;
352 const int present = i < valid;
353 const float gate_value = present
354 ? g[(size_t)
token * gate_row_stride + (
size_t)head]
356 gcum[i] = gate_value + (i ? gcum[i - 1] : 0.0f);
357 gate_exp[i] = expf(gcum[i]);
359 float beta_value = 0.0f;
360 const float *q_src = NULL;
361 const float *k_src = NULL;
362 const float *v_src = NULL;
365 beta[(
size_t)
token * gate_row_stride + (
size_t)head]);
366 q_src = q + (size_t)
token * qk_row_stride +
367 (
size_t)group * (size_t)state_dim;
368 k_src = k + (size_t)
token * qk_row_stride +
369 (
size_t)group * (size_t)state_dim;
370 v_src = v + (size_t)
token * value_row_stride +
371 (
size_t)head * (size_t)state_dim;
373 beta_chunk[i] = beta_value;
374 for (
int d = 0; d < state_dim; ++d) {
375 const size_t offset = (size_t)i * (
size_t)state_dim + (size_t)d;
376 q_chunk[offset] = present ? q_src[d] * scale : 0.0f;
377 k_chunk[offset] = present ? k_src[d] : 0.0f;
378 value_beta[offset] = present ? v_src[d] * beta_value : 0.0f;
380 present ? k_src[d] * beta_value * gate_exp[i] : 0.0f;
390 for (
int i = 0; i <
C; ++i) {
391 for (
int j = 0; j <
C; ++j) {
392 const float d = j <= i ? expf(gcum[i] - gcum[j]) : 0.0f;
393 decay[(size_t)i *
C + (
size_t)j] = d;
394 transform[(size_t)i *
C + (
size_t)j] = j < i
395 ? ck_deltanet_llama_avx2_dot(
396 k_chunk + (
size_t)i * (
size_t)state_dim,
397 k_chunk + (
size_t)j * (
size_t)state_dim,
398 state_dim) * beta_chunk[i] * d
402 for (
int i = 0; i <
C; ++i) {
403 for (
int j = 0; j < i; ++j) {
404 work_row[j] = transform[(size_t)i *
C + (
size_t)j];
406 for (
int col = 0; col <= i; ++col) {
407 const float rhs = i == col ? 1.0f : 0.0f;
409 for (
int j = 0; j < i; ++j) {
410 solved -= work_row[j] * transform[(size_t)j *
C + (
size_t)col];
412 transform[(size_t)i *
C + (
size_t)col] = solved;
421 for (
int i = 0; i <
C; ++i) {
423 for (; d + 7 < state_dim; d += 8) {
424 __m256 value_sum = _mm256_setzero_ps();
425 __m256 key_sum = _mm256_setzero_ps();
426 for (
int j = 0; j <
C; ++j) {
427 const __m256 coefficient = _mm256_set1_ps(
428 transform[(
size_t)i *
C + (
size_t)j]);
429 value_sum = ck_deltanet_fmadd8(
431 _mm256_loadu_ps(value_beta +
432 (
size_t)j * (
size_t)state_dim + (
size_t)d),
434 key_sum = ck_deltanet_fmadd8(
436 _mm256_loadu_ps(k_cumdecay +
437 (
size_t)j * (
size_t)state_dim + (
size_t)d),
440 _mm256_storeu_ps(v_new +
441 (
size_t)i * (
size_t)state_dim + (
size_t)d, value_sum);
442 _mm256_storeu_ps(matrix_work +
443 (
size_t)i * (
size_t)state_dim + (
size_t)d, key_sum);
445 for (; d < state_dim; ++d) {
446 float value_sum = 0.0f;
447 float key_sum = 0.0f;
448 for (
int j = 0; j <
C; ++j) {
449 const float coefficient = transform[(size_t)i *
C + (
size_t)j];
450 value_sum += coefficient *
451 value_beta[(size_t)j * (
size_t)state_dim + (size_t)d];
452 key_sum += coefficient *
453 k_cumdecay[(size_t)j * (
size_t)state_dim + (size_t)d];
455 v_new[(size_t)i * (
size_t)state_dim + (size_t)d] = value_sum;
456 matrix_work[(size_t)i * (
size_t)state_dim + (size_t)d] = key_sum;
459 for (
int i = 0; i <
C; ++i) {
461 for (; d + 7 < state_dim; d += 8) {
462 __m256 v_prime = _mm256_setzero_ps();
463 for (
int r = 0; r < state_dim; ++r) {
464 v_prime = ck_deltanet_fmadd8(
465 _mm256_set1_ps(matrix_work[
466 (
size_t)i * (
size_t)state_dim + (
size_t)r]),
467 _mm256_loadu_ps(state +
468 (
size_t)r * (
size_t)state_dim + (
size_t)d),
472 (size_t)i * (
size_t)state_dim + (size_t)d;
473 _mm256_storeu_ps(dst, _mm256_sub_ps(_mm256_loadu_ps(dst), v_prime));
475 for (; d < state_dim; ++d) {
476 float v_prime = 0.0f;
477 for (
int r = 0; r < state_dim; ++r) {
479 matrix_work[(size_t)i * (
size_t)state_dim + (size_t)r] *
480 state[(
size_t)r * (size_t)state_dim + (
size_t)d];
482 v_new[(size_t)i * (
size_t)state_dim + (size_t)d] -= v_prime;
486 for (
int i = 0; i < valid; ++i) {
487 float *out_token = out +
488 (size_t)(chunk_start + i) * value_row_stride +
489 (size_t)head * (
size_t)state_dim;
490 for (
int j = 0; j <= i; ++j) {
491 work_row[j] = ck_deltanet_llama_avx2_dot(
492 q_chunk + (
size_t)i * (
size_t)state_dim,
493 k_chunk + (
size_t)j * (
size_t)state_dim,
494 state_dim) * decay[(size_t)i *
C + (
size_t)j];
497 for (; d + 7 < state_dim; d += 8) {
498 __m256 result = _mm256_setzero_ps();
499 for (
int r = 0; r < state_dim; ++r) {
501 q_chunk[(size_t)i * (
size_t)state_dim + (size_t)r] *
503 result = ck_deltanet_fmadd8(
504 _mm256_set1_ps(q_gate),
505 _mm256_loadu_ps(state +
506 (
size_t)r * (
size_t)state_dim + (
size_t)d),
509 for (
int j = 0; j <= i; ++j) {
510 result = ck_deltanet_fmadd8(
511 _mm256_set1_ps(work_row[j]),
512 _mm256_loadu_ps(v_new +
513 (
size_t)j * (
size_t)state_dim + (
size_t)d),
516 _mm256_storeu_ps(out_token + d, result);
518 for (; d < state_dim; ++d) {
520 for (
int r = 0; r < state_dim; ++r) {
522 q_chunk[(size_t)i * (
size_t)state_dim + (size_t)r] *
524 state[(
size_t)r * (size_t)state_dim + (
size_t)d];
526 for (
int j = 0; j <= i; ++j) {
527 result += work_row[j] *
528 v_new[(size_t)j * (
size_t)state_dim + (size_t)d];
530 out_token[d] = result;
534 const float last_decay = gate_exp[
C - 1];
535 for (
int i = 0; i <
C; ++i) {
536 gate_exp[i] = expf(gcum[
C - 1] - gcum[i]);
538 for (
int r = 0; r < state_dim; ++r) {
540 for (; d + 7 < state_dim; d += 8) {
541 float *state_row = state +
542 (size_t)r * (
size_t)state_dim + (size_t)d;
543 __m256 updated = _mm256_mul_ps(
544 _mm256_loadu_ps(state_row), _mm256_set1_ps(last_decay));
545 for (
int i = 0; i <
C; ++i) {
546 const float key_gate =
547 k_chunk[(size_t)i * (
size_t)state_dim + (size_t)r] *
549 updated = ck_deltanet_fmadd8(
550 _mm256_set1_ps(key_gate),
551 _mm256_loadu_ps(v_new +
552 (
size_t)i * (
size_t)state_dim + (
size_t)d),
555 _mm256_storeu_ps(state_row, updated);
557 for (; d < state_dim; ++d) {
559 state[(size_t)r * (
size_t)state_dim + (size_t)d] * last_decay;
560 for (
int i = 0; i <
C; ++i) {
562 k_chunk[(size_t)i * (
size_t)state_dim + (size_t)r] *
564 v_new[(
size_t)i * (size_t)state_dim + (
size_t)d];
566 state[(size_t)r * (
size_t)state_dim + (size_t)d] = updated;
580 const float *state_in,
589 int pytorch_bf16_boundaries)
593 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
595 group_count <= 0 || num_heads % group_count != 0 ||
597 head_begin < 0 || head_end < head_begin || head_end > num_heads) {
600 const float scale = ck_deltanet_llama_scale(state_dim);
601 const size_t vector_stride = (size_t) state_dim;
602 const size_t state_stride = (size_t) state_dim * (
size_t) state_dim;
606 for (
int h = head_begin; h < head_end; ++h) {
612 const int group = pytorch_bf16_boundaries
613 ? h / (num_heads / group_count)
615 const float *q_head = q + (size_t) group * vector_stride;
616 const float *k_head = k + (size_t) group * vector_stride;
617 const float *v_head = v + (size_t) h * vector_stride;
618 const float *state_prev = state_in + (size_t) h * state_stride;
619 float *state_cur = state_out + (size_t) h * state_stride;
620 float *out_head = out + (size_t) h * vector_stride;
623 if (pytorch_bf16_boundaries) {
631 if (pytorch_bf16_boundaries) {
633 for (
int row = 0; row < state_dim; ++row) {
634 q_scaled[row] = q_head[row] * scale;
638 for (
int row = 0; row < state_dim; ++row) {
639 const size_t row_offset = (size_t) row * (
size_t) state_dim;
640 for (
int col = 0; col < state_dim; ++col) {
641 state_cur[row_offset + (size_t) col] =
642 state_prev[row_offset + (
size_t) col] * gate;
646 for (
int col = 0; col < state_dim; ++col) {
647 for (
int row = 0; row < state_dim; ++row) {
648 column[row] = state_cur[(size_t) row * (
size_t) state_dim + (size_t) col];
650 const float memory = ck_deltanet_llama_avx2_dot(column, k_head, state_dim);
651 const float delta = (v_head[col] - memory) * beta_s;
652 for (
int row = 0; row < state_dim; ++row) {
653 const size_t offset = (size_t) row * (
size_t) state_dim + (size_t) col;
655 const float updated = fmaf(k_head[row], delta, state_cur[offset]);
657 const float updated = state_cur[offset] + k_head[row] * delta;
659 state_cur[offset] = updated;
660 column[row] = updated;
662 if (pytorch_bf16_boundaries) {
664 ck_deltanet_llama_avx2_dot(column, q_scaled, state_dim);
667 ck_deltanet_llama_avx2_dot(column, q_head, state_dim) * scale;
672 if (group_count != num_heads) {
676 q, k, v, g, beta, state_in, state_out, out,
677 num_heads, state_dim, norm_eps);
693 const float *state_in,
705 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
706 num_heads <= 0 || group_count <= 0 || num_heads % group_count != 0 ||
708 head_begin < 0 || head_end < head_begin || head_end > num_heads) {
711 const float scale = ck_deltanet_llama_scale(state_dim);
712 const size_t vector_stride = (size_t) state_dim;
713 const size_t state_stride = (size_t) state_dim * (
size_t) state_dim;
716 for (
int h = head_begin; h < head_end; ++h) {
717 const int group = h % group_count;
718 const float *q_head = q + (size_t) group * vector_stride;
719 const float *k_head = k + (size_t) group * vector_stride;
720 const float *v_head = v + (size_t) h * vector_stride;
721 const float *state_prev = state_in + (size_t) h * state_stride;
722 float *state_cur = state_out + (size_t) h * state_stride;
723 float *out_head = out + (size_t) h * vector_stride;
727 for (
int col = 0; col < state_dim; ++col) {
728 const float *prev_col = state_prev + (size_t) col * vector_stride;
729 float *cur_col = state_cur + (size_t) col * vector_stride;
731 const __m256 gate8 = _mm256_set1_ps(gate);
732 for (; row + 7 < state_dim; row += 8) {
733 const __m256 scaled = _mm256_mul_ps(
734 _mm256_loadu_ps(prev_col + row), gate8);
735 _mm256_storeu_ps(cur_col + row, scaled);
737 for (; row < state_dim; ++row) {
738 cur_col[row] = prev_col[row] * gate;
742 ck_deltanet_llama_avx2_dot(cur_col, k_head, state_dim);
743 const float delta = (v_head[col] - memory) * beta_s;
745 const __m256 delta8 = _mm256_set1_ps(delta);
746 for (; row + 7 < state_dim; row += 8) {
747 const __m256 updated = _mm256_fmadd_ps(
748 _mm256_loadu_ps(k_head + row), delta8,
749 _mm256_loadu_ps(cur_col + row));
750 _mm256_storeu_ps(cur_col + row, updated);
752 for (; row < state_dim; ++row) {
753 cur_col[row] = fmaf(k_head[row], delta, cur_col[row]);
756 ck_deltanet_llama_avx2_dot(cur_col, q_head, state_dim) * scale;
760 (void)q; (void)k; (void)v; (void)g; (void)beta;
761 (void)state_in; (void)state_out; (void)out;
762 (void)num_heads; (void)group_count; (void)state_dim; (void)norm_eps;
763 (void)head_begin; (void)head_end;
772 const float *state_in,
781 q, k, v, g, beta, state_in, state_out, out,
782 num_heads, group_count, state_dim, norm_eps, 0, num_heads);
796 const float *state_in,
807 q, k, v, g, beta, state_in, state_out, out,
808 num_heads, group_count, state_dim, norm_eps,
809 head_begin, head_end);
816 while (power < value) {
823#if defined(__GNUC__) && !defined(__clang__)
827 const float *row_weights,
831 const int num_levels = 4;
833 if (level_power < 4) {
836 const int level_step = 1 << level_power;
837 const int level_mask = level_step - 1;
840#if defined(__AVX512F__)
842 for (; col + 63 < state_dim; col += 64) {
844 for (
int level = 0; level < num_levels; ++level) {
845 for (
int block = 0; block < 4; ++block) {
846 acc[level][block] = _mm512_setzero_ps();
851 for (; i + level_step <= state_dim;) {
852 for (
int j = 0; j < level_step; ++j, ++i) {
853 const float *row = matrix + (size_t)i * (
size_t)state_dim + col;
854 const __m512 weight = _mm512_set1_ps(row_weights[i]);
855 for (
int block = 0; block < 4; ++block) {
856 const __m512 product = _mm512_mul_ps(
857 _mm512_loadu_ps(row + block * 16), weight);
858 acc[0][block] = _mm512_add_ps(acc[0][block], product);
862 for (
int level = 1; level < num_levels; ++level) {
863 for (
int block = 0; block < 4; ++block) {
864 acc[level][block] = _mm512_add_ps(
865 acc[level][block], acc[level - 1][block]);
866 acc[level - 1][block] = _mm512_setzero_ps();
868 const int mask = level_mask << (level * level_power);
869 if ((i &
mask) != 0) {
875 for (; i < state_dim; ++i) {
876 const float *row = matrix + (size_t)i * (
size_t)state_dim + col;
877 const __m512 weight = _mm512_set1_ps(row_weights[i]);
878 for (
int block = 0; block < 4; ++block) {
879 const __m512 product = _mm512_mul_ps(
880 _mm512_loadu_ps(row + block * 16), weight);
881 acc[0][block] = _mm512_add_ps(acc[0][block], product);
885 for (
int level = 1; level < num_levels; ++level) {
886 for (
int block = 0; block < 4; ++block) {
887 acc[0][block] = _mm512_add_ps(
888 acc[0][block], acc[level][block]);
891 for (
int block = 0; block < 4; ++block) {
892 _mm512_storeu_ps(output + col + block * 16, acc[0][block]);
895#elif defined(__AVX2__)
896 for (; col + 31 < state_dim; col += 32) {
898 for (
int level = 0; level < num_levels; ++level) {
899 for (
int block = 0; block < 4; ++block) {
900 acc[level][block] = _mm256_setzero_ps();
905 for (; i + level_step <= state_dim;) {
906 for (
int j = 0; j < level_step; ++j, ++i) {
907 const float *row = matrix + (size_t)i * (
size_t)state_dim + col;
908 const __m256 weight = _mm256_set1_ps(row_weights[i]);
909 for (
int block = 0; block < 4; ++block) {
910 const __m256 product = _mm256_mul_ps(
911 _mm256_loadu_ps(row + block * 8), weight);
912 acc[0][block] = _mm256_add_ps(acc[0][block], product);
915 for (
int level = 1; level < num_levels; ++level) {
916 for (
int block = 0; block < 4; ++block) {
917 acc[level][block] = _mm256_add_ps(
918 acc[level][block], acc[level - 1][block]);
919 acc[level - 1][block] = _mm256_setzero_ps();
921 const int mask = level_mask << (level * level_power);
922 if ((i &
mask) != 0) {
927 for (; i < state_dim; ++i) {
928 const float *row = matrix + (size_t)i * (
size_t)state_dim + col;
929 const __m256 weight = _mm256_set1_ps(row_weights[i]);
930 for (
int block = 0; block < 4; ++block) {
931 const __m256 product = _mm256_mul_ps(
932 _mm256_loadu_ps(row + block * 8), weight);
933 acc[0][block] = _mm256_add_ps(acc[0][block], product);
936 for (
int level = 1; level < num_levels; ++level) {
937 for (
int block = 0; block < 4; ++block) {
938 acc[0][block] = _mm256_add_ps(
939 acc[0][block], acc[level][block]);
942 for (
int block = 0; block < 4; ++block) {
943 _mm256_storeu_ps(output + col + block * 8, acc[0][block]);
950 for (; col < state_dim; ++col) {
952 for (
int row = 0; row < state_dim; ++row) {
953 sum += matrix[(size_t)row * (
size_t)state_dim + col] *
960#if defined(__GNUC__) && !defined(__clang__)
969 const float *state_in,
972 float *debug_decayed_state,
981 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
982 num_heads <= 0 || group_count <= 0 || num_heads % group_count != 0 ||
987 const size_t vector_stride = (size_t)state_dim;
988 const size_t state_stride = vector_stride * vector_stride;
989 const int heads_per_group = num_heads / group_count;
990 const float sqrt_dim = sqrtf((
float)state_dim);
997 g, beta, gate_values, beta_values, num_heads);
999 for (
int h = 0; h < num_heads; ++h) {
1000 const int group = h / heads_per_group;
1001 const float *q_head = q + (size_t)group * vector_stride;
1002 const float *k_head = k + (size_t)group * vector_stride;
1003 const float *v_head = v + (size_t)h * vector_stride;
1004 const float *state_prev = state_in + (size_t)h * state_stride;
1005 float *state_cur = state_out + (size_t)h * state_stride;
1006 float *out_head = out + (size_t)h * vector_stride;
1008 for (
int col = 0; col < state_dim; ++col) {
1009 q_scaled[col] = q_head[col] / sqrt_dim;
1013 for (
int row = 0; row < state_dim; ++row) {
1014 const size_t row_offset = (size_t)row * vector_stride;
1016#if defined(__AVX2__)
1017 const __m256 gate8 = _mm256_set1_ps(gate_values[h]);
1018 for (; col + 7 < state_dim; col += 8) {
1019 const __m256 state = _mm256_mul_ps(
1020 _mm256_loadu_ps(state_prev + row_offset + (
size_t)col),
1022 _mm256_storeu_ps(state_cur + row_offset + (
size_t)col, state);
1025 for (; col < state_dim; ++col) {
1026 const size_t offset = row_offset + (size_t)col;
1027 const float state = state_prev[offset] * gate_values[h];
1028 state_cur[offset] = state;
1033 state_cur, k_head, memory, state_dim);
1035 if (debug_decayed_state) {
1037 debug_decayed_state + (
size_t)h * state_stride,
1039 state_stride *
sizeof(
float));
1043 debug_memory + (
size_t)h * vector_stride,
1045 vector_stride *
sizeof(
float));
1048 for (
int col = 0; col < state_dim; ++col) {
1049 delta[col] = (v_head[col] - memory[col]) * beta_values[h];
1053 debug_delta + (
size_t)h * vector_stride,
1055 vector_stride *
sizeof(
float));
1058 for (
int row = 0; row < state_dim; ++row) {
1059 const size_t row_offset = (size_t)row * vector_stride;
1060 const float key = k_head[row];
1062#if defined(__AVX2__)
1063 const __m256 key8 = _mm256_set1_ps(key);
1064 for (; col + 7 < state_dim; col += 8) {
1065 const __m256 update = _mm256_mul_ps(
1066 key8, _mm256_loadu_ps(delta + col));
1067 const __m256 state = _mm256_add_ps(
1068 _mm256_loadu_ps(state_cur + row_offset + (
size_t)col),
1070 _mm256_storeu_ps(state_cur + row_offset + (
size_t)col, state);
1073 for (; col < state_dim; ++col) {
1074 const size_t offset = row_offset + (size_t)col;
1075 const float state = state_cur[offset] + key * delta[col];
1076 state_cur[offset] = state;
1081 state_cur, q_scaled, out_head, state_dim);
1083 for (
int col = 0; col < state_dim; ++col) {
1094 const float *state_in,
1103 q, k, v, g, beta, state_in, state_out, out,
1105 num_heads, group_count, state_dim, norm_eps);
1114 const float *state_in,
1117 float *decayed_state,
1126 q, k, v, g, beta, state_in, state_out, out,
1127 decayed_state, memory, delta,
1128 num_heads, group_count, state_dim, norm_eps);
1136 const float *state_in,
1145 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
1146 rows <= 0 || num_heads <= 0 || group_count <= 0 ||
1147 num_heads % group_count != 0 || state_dim <= 0) {
1150 const size_t qk_stride = (size_t) group_count * (
size_t) state_dim;
1151 const size_t value_stride = (size_t) num_heads * (
size_t) state_dim;
1152 const size_t gate_stride = (size_t) num_heads;
1153 for (
int row = 0; row < rows; ++row) {
1155 q + (
size_t) row * qk_stride,
1156 k + (
size_t) row * qk_stride,
1157 v + (
size_t) row * value_stride,
1158 g + (
size_t) row * gate_stride,
1159 beta + (
size_t) row * gate_stride,
1160 row == 0 ? state_in : state_out,
1162 out + (size_t) row * value_stride,
1163 num_heads, group_count, state_dim, norm_eps);
1172 const float *state_in,
1181 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
1182 rows <= 0 || num_heads <= 0 || group_count <= 0 ||
1183 num_heads % group_count != 0 || state_dim <= 0) {
1186#if defined(__AVX2__)
1189 for (
int head = 0; head < num_heads; ++head) {
1190 gated_deltanet_llama_chunk64_head(
1191 q, k, v, g, beta, state_in, state_out, out,
1192 rows, num_heads, group_count, head, state_dim);
1198 q, k, v, g, beta, state_in, state_out, out,
1199 rows, num_heads, group_count, state_dim, norm_eps);
1207 const float *state_in,
1216#if defined(__AVX2__)
1217 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
1218 rows <= 0 || num_heads <= 0 || group_count <= 0 ||
1219 num_heads % group_count != 0 || head < 0 || head >= num_heads ||
1223 gated_deltanet_llama_chunk64_head(
1224 q, k, v, g, beta, state_in, state_out, out,
1225 rows, num_heads, group_count, head, state_dim);
1249 const float *state_in,
1258 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
1259 rows <= 0 || num_heads <= 0 || group_count <= 0 ||
1260 num_heads % group_count != 0 || state_dim <= 0) {
1263 const size_t qk_stride = (size_t)group_count * (
size_t)state_dim;
1264 const size_t value_stride = (size_t)num_heads * (
size_t)state_dim;
1265 const size_t gate_stride = (size_t)num_heads;
1266 for (
int row = 0; row < rows; ++row) {
1268 q + (
size_t)row * qk_stride,
1269 k + (
size_t)row * qk_stride,
1270 v + (
size_t)row * value_stride,
1271 g + (
size_t)row * gate_stride,
1272 beta + (
size_t)row * gate_stride,
1273 row == 0 ? state_in : state_out,
1275 out + (size_t)row * value_stride,
1276 num_heads, group_count, state_dim, norm_eps);
1285 const float *state_in,
1292 const float q_scale = 1.0f / sqrtf((
float)state_dim);
1293 const size_t vec_stride = (size_t)state_dim;
1294 const size_t state_stride = (size_t)state_dim * (
size_t)state_dim;
1296 for (
int h = 0; h < num_heads; ++h) {
1297 const float *q_head = q + (size_t)h * vec_stride;
1298 const float *k_head = k + (size_t)h * vec_stride;
1299 const float *v_head = v + (size_t)h * vec_stride;
1300 const float *state_prev = state_in + (size_t)h * state_stride;
1301 float *state_cur = state_out + (size_t)h * state_stride;
1302 float *out_head = out + (size_t)h * vec_stride;
1305 const float gate = expf(g[h]);
1307 for (
int row = 0; row < state_dim; ++row) {
1308 const size_t row_off = (size_t)row * (
size_t)state_dim;
1309 for (
int col = 0; col < state_dim; ++col) {
1310 state_cur[row_off + (size_t)col] = state_prev[row_off + (
size_t)col] * gate;
1314 for (
int col = 0; col < state_dim; ++col) {
1315 float kv_mem = 0.0f;
1316 for (
int row = 0; row < state_dim; ++row) {
1317 const float k_hat = k_head[row];
1318 kv_mem += state_cur[(size_t)row * (
size_t)state_dim + (size_t)col] * k_hat;
1321 const float delta = (v_head[col] - kv_mem) * beta_s;
1322 for (
int row = 0; row < state_dim; ++row) {
1323 const float k_hat = k_head[row];
1324 state_cur[(size_t)row * (
size_t)state_dim + (size_t)col] += k_hat * delta;
1328 for (
int col = 0; col < state_dim; ++col) {
1330 for (
int row = 0; row < state_dim; ++row) {
1331 const float q_hat = q_head[row] * q_scale;
1332 acc += state_cur[(size_t)row * (
size_t)state_dim + (size_t)col] * q_hat;
1334 out_head[col] = acc;
1340 const float *d_state_out,
1346 const float *state_in,
1347 const float *state_out,
1358 const float q_scale = 1.0f / sqrtf((
float)state_dim);
1359 const size_t vec_stride = (size_t)state_dim;
1360 const size_t state_stride = (size_t)state_dim * (
size_t)state_dim;
1370 for (
int h = 0; h < num_heads; ++h) {
1371 const float *d_out_head = d_out + (size_t)h * vec_stride;
1372 const float *d_state_out_head = d_state_out + (size_t)h * state_stride;
1373 const float *q_head = q + (size_t)h * vec_stride;
1374 const float *k_head = k + (size_t)h * vec_stride;
1375 const float *v_head = v + (size_t)h * vec_stride;
1376 const float *state_prev = state_in + (size_t)h * state_stride;
1377 const float *state_cur = state_out + (size_t)h * state_stride;
1378 float *d_q_head = d_q + (size_t)h * vec_stride;
1379 float *d_k_head = d_k + (size_t)h * vec_stride;
1380 float *d_v_head = d_v + (size_t)h * vec_stride;
1381 float *d_state_prev = d_state_in + (size_t)h * state_stride;
1384 const float gate = expf(g[h]);
1386 float qk_dot = 0.0f;
1387 float out_delta_dot = 0.0f;
1388 float beta_acc = 0.0f;
1389 float gate_acc = 0.0f;
1391 for (
int i = 0; i < state_dim; ++i) {
1392 q_hat[i] = q_head[i] * q_scale;
1393 k_hat[i] = k_head[i];
1399 qk_dot += q_hat[i] * k_hat[i];
1402 for (
int col = 0; col < state_dim; ++col) {
1404 for (
int row = 0; row < state_dim; ++row) {
1405 mem += (state_prev[(size_t)row * (
size_t)state_dim + (size_t)col] * gate) * k_hat[row];
1408 delta[col] = (v_head[col] - mem) * beta_s;
1409 out_delta_dot += d_out_head[col] * delta[col];
1412 for (
int row = 0; row < state_dim; ++row) {
1413 const size_t row_off = (size_t)row * (
size_t)state_dim;
1414 float dq_acc = 0.0f;
1415 float dk_acc = q_hat[row] * out_delta_dot;
1416 for (
int col = 0; col < state_dim; ++col) {
1417 const float d_state_direct = d_state_out_head[row_off + (size_t)col];
1418 dq_acc += state_cur[row_off + (size_t)col] * d_out_head[col];
1419 dk_acc += d_state_direct * delta[col];
1421 d_q_hat[row] = dq_acc;
1422 d_k_hat[row] = dk_acc;
1425 for (
int col = 0; col < state_dim; ++col) {
1426 float d_delta_acc = d_out_head[col] * qk_dot;
1427 for (
int row = 0; row < state_dim; ++row) {
1428 d_delta_acc += d_state_out_head[(size_t)row * (
size_t)state_dim + (size_t)col] * k_hat[row];
1431 d_v_head[col] = beta_s * d_delta_acc;
1432 d_mem[col] = -beta_s * d_delta_acc;
1433 beta_acc += d_delta_acc * (v_head[col] - kv_mem[col]);
1436 for (
int row = 0; row < state_dim; ++row) {
1437 const size_t row_off = (size_t)row * (
size_t)state_dim;
1438 float s_dm_acc = 0.0f;
1439 for (
int col = 0; col < state_dim; ++col) {
1440 s_dm_acc += (state_prev[row_off + (size_t)col] * gate) * d_mem[col];
1442 d_k_hat[row] += s_dm_acc;
1445 for (
int row = 0; row < state_dim; ++row) {
1446 const size_t row_off = (size_t)row * (
size_t)state_dim;
1447 for (
int col = 0; col < state_dim; ++col) {
1448 const float d_state_total = d_state_out_head[row_off + (size_t)col]
1449 + q_hat[row] * d_out_head[col]
1450 + k_hat[row] * d_mem[col];
1451 d_state_prev[row_off + (size_t)col] = gate * d_state_total;
1452 gate_acc += d_state_total * state_prev[row_off + (size_t)col];
1456 for (
int i = 0; i < state_dim; ++i) {
1457 d_q_head[i] = d_q_hat[i] * q_scale;
1458 d_k_head[i] = d_k_hat[i];
1461 d_g[h] = gate_acc * gate;
1462 d_beta[h] = beta_acc * beta_s * (1.0f - beta_s);
1467static void ck_deltanet_scale_rows_avx(
const float *src,
float *dst,
int dim,
float scale)
1469 const __m256 scale_v = _mm256_set1_ps(scale);
1471 for (; i + 8 <= dim; i += 8) {
1472 __m256 x = _mm256_loadu_ps(src + i);
1473 _mm256_storeu_ps(dst + i, _mm256_mul_ps(x, scale_v));
1475 for (; i < dim; ++i) {
1476 dst[i] = src[i] * scale;
1480void gated_deltanet_autoregressive_forward_avx(
const float *q,
1485 const float *state_in,
1492 const float q_scale = 1.0f / sqrtf((
float)state_dim);
1493 const size_t vec_stride = (size_t)state_dim;
1494 const size_t state_stride = (size_t)state_dim * (
size_t)state_dim;
1501 for (
int h = 0; h < num_heads; ++h) {
1502 const float *q_head = q + (size_t)h * vec_stride;
1503 const float *k_head = k + (size_t)h * vec_stride;
1504 const float *v_head = v + (size_t)h * vec_stride;
1505 const float *state_prev = state_in + (size_t)h * state_stride;
1506 float *state_cur = state_out + (size_t)h * state_stride;
1507 float *out_head = out + (size_t)h * vec_stride;
1509 const float gate = expf(g[h]);
1513 ck_deltanet_scale_rows_avx(q_head, q_hat, state_dim, q_scale);
1514 ck_deltanet_scale_rows_avx(k_head, k_hat, state_dim, 1.0f);
1516 const __m256 beta_v = _mm256_set1_ps(beta_s);
1517 const __m256 zero_v = _mm256_setzero_ps();
1520 for (; col + 8 <= state_dim; col += 8) {
1521 _mm256_storeu_ps(kv_mem + col, zero_v);
1522 _mm256_storeu_ps(out_head + col, zero_v);
1524 for (; col < state_dim; ++col) {
1526 out_head[col] = 0.0f;
1529 for (
int row = 0; row < state_dim; ++row) {
1530 const size_t row_off = (size_t)row * (
size_t)state_dim;
1531 const __m256 k_hat_v = _mm256_set1_ps(k_hat[row]);
1532 const __m256 gate_v = _mm256_set1_ps(gate);
1535 for (; col + 8 <= state_dim; col += 8) {
1536 __m256 prev_v = _mm256_loadu_ps(state_prev + row_off + (
size_t)col);
1537 __m256 cur_v = _mm256_mul_ps(prev_v, gate_v);
1538 __m256 kv_v = _mm256_loadu_ps(kv_mem + col);
1539 kv_v = _mm256_add_ps(kv_v, _mm256_mul_ps(cur_v, k_hat_v));
1540 _mm256_storeu_ps(state_cur + row_off + (
size_t)col, cur_v);
1541 _mm256_storeu_ps(kv_mem + col, kv_v);
1543 for (; col < state_dim; ++col) {
1544 const float cur = state_prev[row_off + (size_t)col] * gate;
1545 state_cur[row_off + (size_t)col] = cur;
1546 kv_mem[col] += cur * k_hat[row];
1551 for (; col + 8 <= state_dim; col += 8) {
1552 __m256 v_v = _mm256_loadu_ps(v_head + col);
1553 __m256 kv_v = _mm256_loadu_ps(kv_mem + col);
1554 __m256 delta_v = _mm256_mul_ps(_mm256_sub_ps(v_v, kv_v), beta_v);
1555 _mm256_storeu_ps(delta + col, delta_v);
1557 for (; col < state_dim; ++col) {
1558 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1561 for (
int row = 0; row < state_dim; ++row) {
1562 const size_t row_off = (size_t)row * (
size_t)state_dim;
1563 const __m256 k_hat_v = _mm256_set1_ps(k_hat[row]);
1564 const __m256 q_hat_v = _mm256_set1_ps(q_hat[row]);
1567 for (; col + 8 <= state_dim; col += 8) {
1568 __m256 cur_v = _mm256_loadu_ps(state_cur + row_off + (
size_t)col);
1569 __m256 delta_v = _mm256_loadu_ps(delta + col);
1570 __m256 out_v = _mm256_loadu_ps(out_head + col);
1571 __m256 updated_v = _mm256_add_ps(cur_v, _mm256_mul_ps(k_hat_v, delta_v));
1572 out_v = _mm256_add_ps(out_v, _mm256_mul_ps(updated_v, q_hat_v));
1573 _mm256_storeu_ps(state_cur + row_off + (
size_t)col, updated_v);
1574 _mm256_storeu_ps(out_head + col, out_v);
1576 for (; col < state_dim; ++col) {
1577 const float updated = state_cur[row_off + (size_t)col] + k_hat[row] * delta[col];
1578 state_cur[row_off + (size_t)col] = updated;
1579 out_head[col] += updated * q_hat[row];
1586#if defined(__AVX2__)
1587static inline __m256 ck_deltanet_fmadd256(__m256 a, __m256 b, __m256 c)
1590 return _mm256_fmadd_ps(a, b, c);
1592 return _mm256_add_ps(_mm256_mul_ps(a, b), c);
1596void gated_deltanet_autoregressive_forward_avx2(
const float *q,
1601 const float *state_in,
1608 const float q_scale = 1.0f / sqrtf((
float)state_dim);
1609 const size_t vec_stride = (size_t)state_dim;
1610 const size_t state_stride = (size_t)state_dim * (
size_t)state_dim;
1617 for (
int h = 0; h < num_heads; ++h) {
1618 const float *q_head = q + (size_t)h * vec_stride;
1619 const float *k_head = k + (size_t)h * vec_stride;
1620 const float *v_head = v + (size_t)h * vec_stride;
1621 const float *state_prev = state_in + (size_t)h * state_stride;
1622 float *state_cur = state_out + (size_t)h * state_stride;
1623 float *out_head = out + (size_t)h * vec_stride;
1625 const float gate = expf(g[h]);
1629 ck_deltanet_scale_rows_avx(q_head, q_hat, state_dim, q_scale);
1630 ck_deltanet_scale_rows_avx(k_head, k_hat, state_dim, 1.0f);
1632 const __m256 beta_v = _mm256_set1_ps(beta_s);
1633 const __m256 zero_v = _mm256_setzero_ps();
1636 for (; col + 8 <= state_dim; col += 8) {
1637 _mm256_storeu_ps(kv_mem + col, zero_v);
1638 _mm256_storeu_ps(out_head + col, zero_v);
1640 for (; col < state_dim; ++col) {
1642 out_head[col] = 0.0f;
1646 for (; row + 2 <= state_dim; row += 2) {
1647 const size_t row0_off = (size_t)row * (
size_t)state_dim;
1648 const size_t row1_off = (size_t)(row + 1) * (size_t)state_dim;
1649 const __m256 k0_v = _mm256_set1_ps(k_hat[row]);
1650 const __m256 k1_v = _mm256_set1_ps(k_hat[row + 1]);
1651 const __m256 gate0_v = _mm256_set1_ps(gate);
1652 const __m256 gate1_v = _mm256_set1_ps(gate);
1655 for (; col + 8 <= state_dim; col += 8) {
1656 __m256 prev0_v = _mm256_loadu_ps(state_prev + row0_off + (
size_t)col);
1657 __m256 prev1_v = _mm256_loadu_ps(state_prev + row1_off + (
size_t)col);
1658 __m256 cur0_v = _mm256_mul_ps(prev0_v, gate0_v);
1659 __m256 cur1_v = _mm256_mul_ps(prev1_v, gate1_v);
1660 __m256 kv_v = _mm256_loadu_ps(kv_mem + col);
1661 kv_v = ck_deltanet_fmadd256(cur0_v, k0_v, kv_v);
1662 kv_v = ck_deltanet_fmadd256(cur1_v, k1_v, kv_v);
1663 _mm256_storeu_ps(state_cur + row0_off + (
size_t)col, cur0_v);
1664 _mm256_storeu_ps(state_cur + row1_off + (
size_t)col, cur1_v);
1665 _mm256_storeu_ps(kv_mem + col, kv_v);
1667 for (; col < state_dim; ++col) {
1668 const float cur0 = state_prev[row0_off + (size_t)col] * gate;
1669 const float cur1 = state_prev[row1_off + (size_t)col] * gate;
1670 state_cur[row0_off + (size_t)col] = cur0;
1671 state_cur[row1_off + (size_t)col] = cur1;
1672 kv_mem[col] += cur0 * k_hat[row] + cur1 * k_hat[row + 1];
1675 for (; row < state_dim; ++row) {
1676 const size_t row_off = (size_t)row * (
size_t)state_dim;
1677 const __m256 k_hat_v = _mm256_set1_ps(k_hat[row]);
1678 const __m256 gate_v = _mm256_set1_ps(gate);
1680 for (; col + 8 <= state_dim; col += 8) {
1681 __m256 prev_v = _mm256_loadu_ps(state_prev + row_off + (
size_t)col);
1682 __m256 cur_v = _mm256_mul_ps(prev_v, gate_v);
1683 __m256 kv_v = _mm256_loadu_ps(kv_mem + col);
1684 kv_v = ck_deltanet_fmadd256(cur_v, k_hat_v, kv_v);
1685 _mm256_storeu_ps(state_cur + row_off + (
size_t)col, cur_v);
1686 _mm256_storeu_ps(kv_mem + col, kv_v);
1688 for (; col < state_dim; ++col) {
1689 const float cur = state_prev[row_off + (size_t)col] * gate;
1690 state_cur[row_off + (size_t)col] = cur;
1691 kv_mem[col] += cur * k_hat[row];
1696 for (; col + 8 <= state_dim; col += 8) {
1697 __m256 v_v = _mm256_loadu_ps(v_head + col);
1698 __m256 kv_v = _mm256_loadu_ps(kv_mem + col);
1699 __m256 delta_v = _mm256_mul_ps(_mm256_sub_ps(v_v, kv_v), beta_v);
1700 _mm256_storeu_ps(delta + col, delta_v);
1702 for (; col < state_dim; ++col) {
1703 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1707 for (; row + 2 <= state_dim; row += 2) {
1708 const size_t row0_off = (size_t)row * (
size_t)state_dim;
1709 const size_t row1_off = (size_t)(row + 1) * (size_t)state_dim;
1710 const __m256 k0_v = _mm256_set1_ps(k_hat[row]);
1711 const __m256 k1_v = _mm256_set1_ps(k_hat[row + 1]);
1712 const __m256 q0_v = _mm256_set1_ps(q_hat[row]);
1713 const __m256 q1_v = _mm256_set1_ps(q_hat[row + 1]);
1716 for (; col + 8 <= state_dim; col += 8) {
1717 __m256 cur0_v = _mm256_loadu_ps(state_cur + row0_off + (
size_t)col);
1718 __m256 cur1_v = _mm256_loadu_ps(state_cur + row1_off + (
size_t)col);
1719 __m256 delta_v = _mm256_loadu_ps(delta + col);
1720 __m256 out_v = _mm256_loadu_ps(out_head + col);
1721 __m256 upd0_v = ck_deltanet_fmadd256(k0_v, delta_v, cur0_v);
1722 __m256 upd1_v = ck_deltanet_fmadd256(k1_v, delta_v, cur1_v);
1723 out_v = ck_deltanet_fmadd256(upd0_v, q0_v, out_v);
1724 out_v = ck_deltanet_fmadd256(upd1_v, q1_v, out_v);
1725 _mm256_storeu_ps(state_cur + row0_off + (
size_t)col, upd0_v);
1726 _mm256_storeu_ps(state_cur + row1_off + (
size_t)col, upd1_v);
1727 _mm256_storeu_ps(out_head + col, out_v);
1729 for (; col < state_dim; ++col) {
1730 const float upd0 = state_cur[row0_off + (size_t)col] + k_hat[row] * delta[col];
1731 const float upd1 = state_cur[row1_off + (size_t)col] + k_hat[row + 1] * delta[col];
1732 state_cur[row0_off + (size_t)col] = upd0;
1733 state_cur[row1_off + (size_t)col] = upd1;
1734 out_head[col] += upd0 * q_hat[row] + upd1 * q_hat[row + 1];
1737 for (; row < state_dim; ++row) {
1738 const size_t row_off = (size_t)row * (
size_t)state_dim;
1739 const __m256 k_hat_v = _mm256_set1_ps(k_hat[row]);
1740 const __m256 q_hat_v = _mm256_set1_ps(q_hat[row]);
1742 for (; col + 8 <= state_dim; col += 8) {
1743 __m256 cur_v = _mm256_loadu_ps(state_cur + row_off + (
size_t)col);
1744 __m256 delta_v = _mm256_loadu_ps(delta + col);
1745 __m256 out_v = _mm256_loadu_ps(out_head + col);
1746 __m256 updated_v = ck_deltanet_fmadd256(k_hat_v, delta_v, cur_v);
1747 out_v = ck_deltanet_fmadd256(updated_v, q_hat_v, out_v);
1748 _mm256_storeu_ps(state_cur + row_off + (
size_t)col, updated_v);
1749 _mm256_storeu_ps(out_head + col, out_v);
1751 for (; col < state_dim; ++col) {
1752 const float updated = state_cur[row_off + (size_t)col] + k_hat[row] * delta[col];
1753 state_cur[row_off + (size_t)col] = updated;
1754 out_head[col] += updated * q_hat[row];
1761#if defined(__AVX512F__)
1762static inline __m512 ck_deltanet_madd512(__m512 a, __m512 b, __m512 c)
1764 return _mm512_add_ps(_mm512_mul_ps(a, b), c);
1767static void ck_deltanet_scale_rows_avx512(
const float *src,
float *dst,
int dim,
float scale)
1769 const __m512 scale_v = _mm512_set1_ps(scale);
1771 for (; i + 16 <= dim; i += 16) {
1772 __m512 x = _mm512_loadu_ps(src + i);
1773 _mm512_storeu_ps(dst + i, _mm512_mul_ps(x, scale_v));
1775 for (; i < dim; ++i) {
1776 dst[i] = src[i] * scale;
1780void gated_deltanet_autoregressive_forward_avx512(
const float *q,
1785 const float *state_in,
1792 const float q_scale = 1.0f / sqrtf((
float)state_dim);
1793 const size_t vec_stride = (size_t)state_dim;
1794 const size_t state_stride = (size_t)state_dim * (
size_t)state_dim;
1801 for (
int h = 0; h < num_heads; ++h) {
1802 const float *q_head = q + (size_t)h * vec_stride;
1803 const float *k_head = k + (size_t)h * vec_stride;
1804 const float *v_head = v + (size_t)h * vec_stride;
1805 const float *state_prev = state_in + (size_t)h * state_stride;
1806 float *state_cur = state_out + (size_t)h * state_stride;
1807 float *out_head = out + (size_t)h * vec_stride;
1809 const float gate = expf(g[h]);
1813 ck_deltanet_scale_rows_avx512(q_head, q_hat, state_dim, q_scale);
1814 ck_deltanet_scale_rows_avx512(k_head, k_hat, state_dim, 1.0f);
1816 const __m512 beta_v = _mm512_set1_ps(beta_s);
1817 const __m512 zero_v = _mm512_setzero_ps();
1820 for (; col + 16 <= state_dim; col += 16) {
1821 _mm512_storeu_ps(kv_mem + col, zero_v);
1822 _mm512_storeu_ps(out_head + col, zero_v);
1824 for (; col < state_dim; ++col) {
1826 out_head[col] = 0.0f;
1829 for (
int row = 0; row < state_dim; ++row) {
1830 const size_t row_off = (size_t)row * (
size_t)state_dim;
1831 const __m512 k_hat_v = _mm512_set1_ps(k_hat[row]);
1832 const __m512 gate_v = _mm512_set1_ps(gate);
1834 for (; col + 16 <= state_dim; col += 16) {
1835 __m512 prev_v = _mm512_loadu_ps(state_prev + row_off + (
size_t)col);
1836 __m512 cur_v = _mm512_mul_ps(prev_v, gate_v);
1837 __m512 kv_v = _mm512_loadu_ps(kv_mem + col);
1838 kv_v = ck_deltanet_madd512(cur_v, k_hat_v, kv_v);
1839 _mm512_storeu_ps(state_cur + row_off + (
size_t)col, cur_v);
1840 _mm512_storeu_ps(kv_mem + col, kv_v);
1842 for (; col < state_dim; ++col) {
1843 const float cur = state_prev[row_off + (size_t)col] * gate;
1844 state_cur[row_off + (size_t)col] = cur;
1845 kv_mem[col] += cur * k_hat[row];
1850 for (; col + 16 <= state_dim; col += 16) {
1851 __m512 v_v = _mm512_loadu_ps(v_head + col);
1852 __m512 kv_v = _mm512_loadu_ps(kv_mem + col);
1853 __m512 delta_v = _mm512_mul_ps(_mm512_sub_ps(v_v, kv_v), beta_v);
1854 _mm512_storeu_ps(delta + col, delta_v);
1856 for (; col < state_dim; ++col) {
1857 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1860 for (
int row = 0; row < state_dim; ++row) {
1861 const size_t row_off = (size_t)row * (
size_t)state_dim;
1862 const __m512 k_hat_v = _mm512_set1_ps(k_hat[row]);
1863 const __m512 q_hat_v = _mm512_set1_ps(q_hat[row]);
1865 for (; col + 16 <= state_dim; col += 16) {
1866 __m512 cur_v = _mm512_loadu_ps(state_cur + row_off + (
size_t)col);
1867 __m512 delta_v = _mm512_loadu_ps(delta + col);
1868 __m512 out_v = _mm512_loadu_ps(out_head + col);
1869 __m512 updated_v = ck_deltanet_madd512(k_hat_v, delta_v, cur_v);
1870 out_v = ck_deltanet_madd512(updated_v, q_hat_v, out_v);
1871 _mm512_storeu_ps(state_cur + row_off + (
size_t)col, updated_v);
1872 _mm512_storeu_ps(out_head + col, out_v);
1874 for (; col < state_dim; ++col) {
1875 const float updated = state_cur[row_off + (size_t)col] + k_hat[row] * delta[col];
1876 state_cur[row_off + (size_t)col] = updated;
1877 out_head[col] += updated * q_hat[row];
1886 const char *env = getenv(
"CK_DELTANET_FORCE_REF");
1887 return env && atoi(env) != 0;
1895#if defined(__AVX512F__)
1897#elif defined(__AVX2__)
1899#elif defined(__AVX__)
1911 const float *state_in,
1918 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out) {
1921 if (num_heads <= 0 || state_dim <= 0) {
1931 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1934#if defined(__AVX512F__)
1935 gated_deltanet_autoregressive_forward_avx512(
1936 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1937#elif defined(__AVX2__)
1938 gated_deltanet_autoregressive_forward_avx2(
1939 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1940#elif defined(__AVX__)
1941 gated_deltanet_autoregressive_forward_avx(
1942 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1945 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1954 const float *state_in,
1962 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out) {
1965 if (rows <= 0 || num_heads <= 0 || state_dim <= 0) {
1969 const size_t vector_stride = (size_t)num_heads * (
size_t)state_dim;
1970 const size_t gate_stride = (size_t)num_heads;
1971 for (
int row = 0; row < rows; ++row) {
1972 const float *row_state_in = row == 0 ? state_in : state_out;
1974 q + (
size_t)row * vector_stride,
1975 k + (
size_t)row * vector_stride,
1976 v + (
size_t)row * vector_stride,
1977 g + (
size_t)row * gate_stride,
1978 beta + (
size_t)row * gate_stride,
1981 out + (
size_t)row * vector_stride,
1989 const float *d_state_out,
1995 const float *state_in,
1996 const float *state_out,
2007 if (!d_out || !d_state_out || !q || !k || !v || !g || !beta || !state_in || !state_out ||
2008 !d_q || !d_k || !d_v || !d_g || !d_beta || !d_state_in) {
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
int ck_strict_parity_enabled(void)
#define CK_DELTANET_LLAMA_CHUNK_SIZE
#define CK_DELTANET_NOINLINE
void gated_deltanet_autoregressive_backward(const float *d_out, const float *d_state_out, const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, const float *state_out, float *d_q, float *d_k, float *d_v, float *d_g, float *d_beta, float *d_state_in, int num_heads, int state_dim, float norm_eps)
static ck_deltanet_mkl_vsexp_fn ck_deltanet_pytorch_vsexp
static pthread_once_t ck_deltanet_libm_once
void gated_deltanet_pytorch_gate_values_debug(const float *g, const float *beta, float *gate_values, float *beta_values, int num_heads)
static void ck_deltanet_pytorch_outer_sum(const float *matrix, const float *row_weights, float *output, int state_dim)
void gated_deltanet_pytorch_grouped_bf16_forward_debug(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, float *decayed_state, float *memory, float *delta, int num_heads, int group_count, int state_dim, float norm_eps)
static void ck_bind_deltanet_pytorch_primitives(void)
void gated_deltanet_autoregressive_forward_ref(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int state_dim, float norm_eps)
static pthread_once_t ck_deltanet_pytorch_primitives_once
static int ck_deltanet_ceil_log2(int value)
static void gated_deltanet_pytorch_grouped_bf16_forward_impl(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, float *debug_decayed_state, float *debug_memory, float *debug_delta, int num_heads, int group_count, int state_dim, float norm_eps)
void gated_deltanet_llama_avx2_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
static ck_deltanet_libm_f32_fn ck_deltanet_llama_expf
static void gated_deltanet_llama_avx2_grouped_forward_transposed_impl(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps, int head_begin, int head_end)
void(* ck_deltanet_mkl_vsexp_fn)(int, const float *, float *)
void gated_deltanet_llama_avx2_forward_head_range(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps, int head_begin, int head_end)
static void * ck_deltanet_mkl_handle
void gated_deltanet_llama_avx2_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps)
#define CK_DELTANET_LLAMA_CHUNK_MAX_DIM
static void * ck_deltanet_libm_handle
static void gated_deltanet_llama_avx2_grouped_forward_impl(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps, int head_begin, int head_end, int pytorch_bf16_boundaries)
static void ck_deltanet_pytorch_gate_values(const float *g, const float *beta, float *gate_values, float *beta_values, int num_heads)
static float ck_deltanet_sigmoidf(float x)
void gated_deltanet_pytorch_grouped_bf16_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps)
void gated_deltanet_pytorch_grouped_bf16_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
void gated_deltanet_autoregressive_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int state_dim, float norm_eps)
void gated_deltanet_llama_chunk64_head_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int head, int state_dim)
#define CK_DELTANET_MAX_STACK_DIM
void gated_deltanet_autoregressive_backward_ref(const float *d_out, const float *d_state_out, const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, const float *state_out, float *d_q, float *d_k, float *d_v, float *d_g, float *d_beta, float *d_state_in, int num_heads, int state_dim, float norm_eps)
static void ck_bind_deltanet_llama_libm(void)
static float ck_deltanet_llama_sigmoidf(float x)
void gated_deltanet_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int state_dim, float norm_eps)
static int ck_deltanet_force_ref(void)
void gated_deltanet_llama_chunk64_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
const char * gated_deltanet_impl_name(void)
float(* ck_deltanet_libm_f32_fn)(float)
int32_t int32_t int32_t int32_t int32_t mask
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)