← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
deltanet_kernels.c
Go to the documentation of this file.
1/**
2 * @file deltanet_kernels.c
3 * @brief FP32 Gated DeltaNet kernels for Qwen3.5-style recurrent attention.
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * This file implements the single-token recurrent update used by the
15 * qwen3next / Gated DeltaNet path in llama.cpp.
16 *
17 * Per head, matching llama.cpp qwen35/qwen3next autoregressive DeltaNet:
18 * q_scaled = q / sqrt(state_dim) // q and k arrive pre-normalized
19 * k_hat = k
20 * beta_s = sigmoid(beta)
21 * gate = exp(g)
22 * S = gate * S_prev
23 * kv_mem = S^T * k_hat
24 * delta = (v - kv_mem) * beta_s
25 * S_new = S + outer(k_hat, delta)
26 * out = S_new^T * q_scaled
27 *
28 * Design:
29 * - *_ref is the scalar reference implementation.
30 * - *_avx keeps a simple 1-row vector walk.
31 * - *_avx2 precomputes scaled q rows and unrolls the state sweep in
32 * row pairs to reduce loop overhead and layout churn.
33 * - The public dispatcher selects the best compiled ISA unless strict parity
34 * is enabled, in which case it falls back to *_ref.
35 */
36
37#include "bf16_utils.h"
38#include "ckernel_engine.h"
39
40#include <dlfcn.h>
41#include <math.h>
42#include <pthread.h>
43#include <stdio.h>
44#include <stddef.h>
45#include <stdlib.h>
46#include <string.h>
47
48#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
49#include <immintrin.h>
50#endif
51
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
55
56#if defined(__GNUC__) || defined(__clang__)
57#define CK_DELTANET_NOINLINE __attribute__((noinline))
58#else
59#define CK_DELTANET_NOINLINE
60#endif
61
62typedef float (*ck_deltanet_libm_f32_fn)(float);
64static void *ck_deltanet_libm_handle = NULL;
65static pthread_once_t ck_deltanet_libm_once = PTHREAD_ONCE_INIT;
66
68{
69 ck_deltanet_libm_handle = dlopen("libm.so.6", RTLD_NOW | RTLD_LOCAL);
73 }
75 fprintf(stderr,
76 "HARD KERNEL CONTRACT FAULT: llama.cpp DeltaNet requires "
77 "expf from libm.so.6\n");
78 abort();
79 }
80}
81
82static inline float ck_deltanet_sigmoidf(float x)
83{
84 return 1.0f / (1.0f + expf(-x));
85}
86
87static inline float ck_deltanet_llama_sigmoidf(float x)
88{
90 return 1.0f / (1.0f + ck_deltanet_llama_expf(-x));
91}
92
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;
97#endif
98
99typedef void (*ck_deltanet_mkl_vsexp_fn)(int, const float *, float *);
101static void *ck_deltanet_mkl_handle = NULL;
102static pthread_once_t ck_deltanet_pytorch_primitives_once = PTHREAD_ONCE_INIT;
103
105{
106 const char *mkl_library = getenv("CK_MKL_LIBRARY");
107 if (mkl_library && *mkl_library) {
108 ck_deltanet_mkl_handle = dlopen(mkl_library, RTLD_NOW | RTLD_LOCAL);
112 ck_deltanet_mkl_handle, "vsExp");
113 }
114 } else {
116 (ck_deltanet_mkl_vsexp_fn)dlsym(RTLD_DEFAULT, "vsExp");
117 }
118
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");
127 }
128 } else {
129 ck_deltanet_pytorch_expf16 =
130 (ck_deltanet_sleef_expf16_fn)dlsym(
131 RTLD_DEFAULT, "Sleef_expf16_u10");
132 }
133#endif
134}
135
136static void ck_deltanet_pytorch_gate_values(const float *g,
137 const float *beta,
138 float *gate_values,
139 float *beta_values,
140 int num_heads)
141{
142 pthread_once(
146 fprintf(stderr,
147 "HARD KERNEL CONTRACT FAULT: PyTorch DeltaNet requires "
148 "MKL vsExp; set CK_MKL_LIBRARY\n");
149 abort();
150 }
151 ck_deltanet_pytorch_vsexp(num_heads, g, gate_values);
152
153 int h = 0;
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) {
166 beta_values[h + lane] = bf16_to_float(
167 float_to_bf16(beta_lanes[lane]));
168 }
169 }
170 }
171#endif
172 for (; h < num_heads; ++h) {
173 beta_values[h] = bf16_to_float(float_to_bf16(
174 ck_deltanet_sigmoidf(beta[h])));
175 }
176}
177
179 const float *beta,
180 float *gate_values,
181 float *beta_values,
182 int num_heads)
183{
184 if (!g || !beta || !gate_values || !beta_values || num_heads <= 0 ||
185 num_heads > CK_DELTANET_MAX_STACK_DIM) {
186 return;
187 }
189 g, beta, gate_values, beta_values, num_heads);
190}
191
193 const float *k,
194 const float *v,
195 const float *g,
196 const float *beta,
197 const float *state_in,
198 float *state_out,
199 float *out,
200 int num_heads,
201 int state_dim,
202 float norm_eps);
203
204#if defined(__AVX2__)
205#if defined(__AVX512F__)
206static inline float ck_deltanet_gcc_reduce_add_ps(__m512 value)
207{
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)));
217}
218#endif
219
220static CK_DELTANET_NOINLINE float ck_deltanet_llama_avx2_dot(
221 const float *x, const float *y, int n)
222{
223#if defined(__AVX512F__)
224 const int np = n & ~63;
225 __m512 sum[4] = {
226 _mm512_setzero_ps(), _mm512_setzero_ps(),
227 _mm512_setzero_ps(), _mm512_setzero_ps()
228 };
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]);
234 }
235 }
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]);
240#else
241 const int np = n & ~31;
242 __m256 sum[4] = {
243 _mm256_setzero_ps(), _mm256_setzero_ps(),
244 _mm256_setzero_ps(), _mm256_setzero_ps()
245 };
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);
250#if defined(__FMA__)
251 sum[j] = _mm256_fmadd_ps(xv, yv, sum[j]);
252#else
253 sum[j] = _mm256_add_ps(_mm256_mul_ps(xv, yv), sum[j]);
254#endif
255 }
256 }
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));
264#endif
265 for (int i = np; i < n; ++i) {
266 result += x[i] * y[i];
267 }
268 return result;
269}
270
271static inline float ck_deltanet_llama_scale(int state_dim)
272{
273#if defined(__SSE__)
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)));
276#else
277 return 1.0f / sqrtf((float) state_dim);
278#endif
279}
280
281static inline __m256 ck_deltanet_fmadd8(__m256 a, __m256 b, __m256 acc)
282{
283#if defined(__FMA__)
284 return _mm256_fmadd_ps(a, b, acc);
285#else
286 return _mm256_add_ps(_mm256_mul_ps(a, b), acc);
287#endif
288}
289
290/*
291 * llama.cpp evaluates multi-token scalar-gate DeltaNet in 64-token chunks.
292 * This is algebraically equivalent to the recurrent update above, but its
293 * reduction tree is observably different after many recurrent layers.
294 *
295 * Layouts below are row-major:
296 * token vectors [chunk, state_dim]
297 * recurrent S [key_dim, value_dim]
298 * temporal mats [chunk, chunk]
299 *
300 * The fixed upper bound keeps the provider allocation-free. Qwen3.5/3.6 use
301 * state_dim=128; unsupported shapes retain the sequential fallback.
302 */
303static void gated_deltanet_llama_chunk64_head(
304 const float *q,
305 const float *k,
306 const float *v,
307 const float *g,
308 const float *beta,
309 const float *state_in,
310 float *state_out,
311 float *out,
312 int rows,
313 int num_heads,
314 int group_count,
315 int head,
316 int state_dim)
317{
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);
325
326 float gcum[C];
327 float beta_chunk[C];
328 float decay[C * C];
329 float transform[C * C];
330 float q_chunk[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
331 float k_chunk[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
332 float value_beta[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
333 float k_cumdecay[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
334 float v_new[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
335 float matrix_work[C * CK_DELTANET_LLAMA_CHUNK_MAX_DIM];
336 float work_row[C];
337 float gate_exp[C];
338
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];
344 }
345 }
346
347 for (int chunk_start = 0; chunk_start < rows; chunk_start += C) {
348 const int valid = rows - chunk_start < C ? rows - chunk_start : C;
349
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]
355 : 0.0f;
356 gcum[i] = gate_value + (i ? gcum[i - 1] : 0.0f);
357 gate_exp[i] = expf(gcum[i]);
358
359 float beta_value = 0.0f;
360 const float *q_src = NULL;
361 const float *k_src = NULL;
362 const float *v_src = NULL;
363 if (present) {
364 beta_value = ck_deltanet_sigmoidf(
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;
372 }
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;
379 k_cumdecay[offset] =
380 present ? k_src[d] * beta_value * gate_exp[i] : 0.0f;
381 }
382 }
383
384 /*
385 * kb[i,j] = beta_i * <k_i,k_j> * exp(gcum_i-gcum_j).
386 * llama solves (I + tril(kb,-1)) X = -tril(kb,-1), then adds I.
387 * Forward substitution is performed a column at a time to retain the
388 * same dependency order as ggml_solve_tri.
389 */
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
399 : 0.0f;
400 }
401 }
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];
405 }
406 for (int col = 0; col <= i; ++col) {
407 const float rhs = i == col ? 1.0f : 0.0f;
408 float solved = rhs;
409 for (int j = 0; j < i; ++j) {
410 solved -= work_row[j] * transform[(size_t)j * C + (size_t)col];
411 }
412 transform[(size_t)i * C + (size_t)col] = solved;
413 }
414 }
415
416 /*
417 * transformed V and cumulative-decay K use the solved temporal
418 * transform. Subtract the contribution of the incoming state to form
419 * V_new, matching llama's v_t_new node.
420 */
421 for (int i = 0; i < C; ++i) {
422 int d = 0;
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(
430 coefficient,
431 _mm256_loadu_ps(value_beta +
432 (size_t)j * (size_t)state_dim + (size_t)d),
433 value_sum);
434 key_sum = ck_deltanet_fmadd8(
435 coefficient,
436 _mm256_loadu_ps(k_cumdecay +
437 (size_t)j * (size_t)state_dim + (size_t)d),
438 key_sum);
439 }
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);
444 }
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];
454 }
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;
457 }
458 }
459 for (int i = 0; i < C; ++i) {
460 int d = 0;
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),
469 v_prime);
470 }
471 float *dst = v_new +
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));
474 }
475 for (; d < state_dim; ++d) {
476 float v_prime = 0.0f;
477 for (int r = 0; r < state_dim; ++r) {
478 v_prime +=
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];
481 }
482 v_new[(size_t)i * (size_t)state_dim + (size_t)d] -= v_prime;
483 }
484 }
485
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];
495 }
496 int d = 0;
497 for (; d + 7 < state_dim; d += 8) {
498 __m256 result = _mm256_setzero_ps();
499 for (int r = 0; r < state_dim; ++r) {
500 const float q_gate =
501 q_chunk[(size_t)i * (size_t)state_dim + (size_t)r] *
502 gate_exp[i];
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),
507 result);
508 }
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),
514 result);
515 }
516 _mm256_storeu_ps(out_token + d, result);
517 }
518 for (; d < state_dim; ++d) {
519 float result = 0.0f;
520 for (int r = 0; r < state_dim; ++r) {
521 result +=
522 q_chunk[(size_t)i * (size_t)state_dim + (size_t)r] *
523 gate_exp[i] *
524 state[(size_t)r * (size_t)state_dim + (size_t)d];
525 }
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];
529 }
530 out_token[d] = result;
531 }
532 }
533
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]);
537 }
538 for (int r = 0; r < state_dim; ++r) {
539 int d = 0;
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] *
548 gate_exp[i];
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),
553 updated);
554 }
555 _mm256_storeu_ps(state_row, updated);
556 }
557 for (; d < state_dim; ++d) {
558 float updated =
559 state[(size_t)r * (size_t)state_dim + (size_t)d] * last_decay;
560 for (int i = 0; i < C; ++i) {
561 updated +=
562 k_chunk[(size_t)i * (size_t)state_dim + (size_t)r] *
563 gate_exp[i] *
564 v_new[(size_t)i * (size_t)state_dim + (size_t)d];
565 }
566 state[(size_t)r * (size_t)state_dim + (size_t)d] = updated;
567 }
568 }
569 }
570}
571
572#endif
573
575 const float *q,
576 const float *k,
577 const float *v,
578 const float *g,
579 const float *beta,
580 const float *state_in,
581 float *state_out,
582 float *out,
583 int num_heads,
584 int group_count,
585 int state_dim,
586 float norm_eps,
587 int head_begin,
588 int head_end,
589 int pytorch_bf16_boundaries)
590{
591#if defined(__AVX2__)
592 (void) norm_eps;
593 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
594 num_heads <= 0 || num_heads > CK_DELTANET_MAX_STACK_DIM ||
595 group_count <= 0 || num_heads % group_count != 0 ||
596 state_dim <= 0 || state_dim > CK_DELTANET_MAX_STACK_DIM ||
597 head_begin < 0 || head_end < head_begin || head_end > num_heads) {
598 return;
599 }
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;
603 float column[CK_DELTANET_MAX_STACK_DIM];
604 float q_scaled[CK_DELTANET_MAX_STACK_DIM];
605
606 for (int h = head_begin; h < head_end; ++h) {
607 /* llama.cpp ggml_repeat_4d tiles compact Q/K heads (0..G-1 repeated),
608 * while the PyTorch Qwen3-Next reference uses repeat_interleave so
609 * each compact head owns H/G adjacent value heads. These layouts
610 * are distinct numerical contracts even though their buffer shapes
611 * are identical. */
612 const int group = pytorch_bf16_boundaries
613 ? h / (num_heads / group_count)
614 : h % 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;
621 float gate;
622 float beta_s;
623 if (pytorch_bf16_boundaries) {
624 gate = expf(g[h]);
625 beta_s = ck_deltanet_sigmoidf(beta[h]);
626 } else {
628 gate = ck_deltanet_llama_expf(g[h]);
629 beta_s = ck_deltanet_llama_sigmoidf(beta[h]);
630 }
631 if (pytorch_bf16_boundaries) {
632 beta_s = bf16_to_float(float_to_bf16(beta_s));
633 for (int row = 0; row < state_dim; ++row) {
634 q_scaled[row] = q_head[row] * scale;
635 }
636 }
637
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;
643 }
644 }
645
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];
649 }
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;
654#if defined(__FMA__)
655 const float updated = fmaf(k_head[row], delta, state_cur[offset]);
656#else
657 const float updated = state_cur[offset] + k_head[row] * delta;
658#endif
659 state_cur[offset] = updated;
660 column[row] = updated;
661 }
662 if (pytorch_bf16_boundaries) {
663 out_head[col] =
664 ck_deltanet_llama_avx2_dot(column, q_scaled, state_dim);
665 } else {
666 out_head[col] =
667 ck_deltanet_llama_avx2_dot(column, q_head, state_dim) * scale;
668 }
669 }
670 }
671#else
672 if (group_count != num_heads) {
673 return;
674 }
676 q, k, v, g, beta, state_in, state_out, out,
677 num_heads, state_dim, norm_eps);
678#endif
679}
680
681/*
682 * llama.cpp stores the recurrent matrix transposed so each logical state
683 * column is contiguous. Keeping that physical layout across prefill and
684 * decode removes the gather/scatter copy from every state update while
685 * preserving the same per-column dot and FMA order.
686 */
688 const float *q,
689 const float *k,
690 const float *v,
691 const float *g,
692 const float *beta,
693 const float *state_in,
694 float *state_out,
695 float *out,
696 int num_heads,
697 int group_count,
698 int state_dim,
699 float norm_eps,
700 int head_begin,
701 int head_end)
702{
703#if defined(__AVX2__)
704 (void) norm_eps;
705 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
706 num_heads <= 0 || group_count <= 0 || num_heads % group_count != 0 ||
707 state_dim <= 0 || state_dim > CK_DELTANET_MAX_STACK_DIM ||
708 head_begin < 0 || head_end < head_begin || head_end > num_heads) {
709 return;
710 }
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;
714
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;
724 const float gate = ck_deltanet_llama_expf(g[h]);
725 const float beta_s = ck_deltanet_llama_sigmoidf(beta[h]);
726
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;
730 int row = 0;
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);
736 }
737 for (; row < state_dim; ++row) {
738 cur_col[row] = prev_col[row] * gate;
739 }
740
741 const float memory =
742 ck_deltanet_llama_avx2_dot(cur_col, k_head, state_dim);
743 const float delta = (v_head[col] - memory) * beta_s;
744 row = 0;
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);
751 }
752 for (; row < state_dim; ++row) {
753 cur_col[row] = fmaf(k_head[row], delta, cur_col[row]);
754 }
755 out_head[col] =
756 ck_deltanet_llama_avx2_dot(cur_col, q_head, state_dim) * scale;
757 }
758 }
759#else
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;
764#endif
765}
766
768 const float *k,
769 const float *v,
770 const float *g,
771 const float *beta,
772 const float *state_in,
773 float *state_out,
774 float *out,
775 int num_heads,
776 int group_count,
777 int state_dim,
778 float norm_eps)
779{
781 q, k, v, g, beta, state_in, state_out, out,
782 num_heads, group_count, state_dim, norm_eps, 0, num_heads);
783}
784
785/*
786 * Orchestrator-facing range entry point. Heads own disjoint state/output
787 * slices, so a threadpool can partition them without changing any per-head
788 * reduction tree. Keep dispatch out of the numerical kernel itself.
789 */
791 const float *q,
792 const float *k,
793 const float *v,
794 const float *g,
795 const float *beta,
796 const float *state_in,
797 float *state_out,
798 float *out,
799 int num_heads,
800 int group_count,
801 int state_dim,
802 float norm_eps,
803 int head_begin,
804 int head_end)
805{
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);
810}
811
812static int ck_deltanet_ceil_log2(int value)
813{
814 int result = 0;
815 int power = 1;
816 while (power < value) {
817 power <<= 1;
818 ++result;
819 }
820 return result;
821}
822
823#if defined(__GNUC__) && !defined(__clang__)
824__attribute__((optimize("fp-contract=off")))
825#endif
826static void ck_deltanet_pytorch_outer_sum(const float *matrix,
827 const float *row_weights,
828 float *output,
829 int state_dim)
830{
831 const int num_levels = 4;
832 int level_power = ck_deltanet_ceil_log2(state_dim) / num_levels;
833 if (level_power < 4) {
834 level_power = 4;
835 }
836 const int level_step = 1 << level_power;
837 const int level_mask = level_step - 1;
838 int col = 0;
839
840#if defined(__AVX512F__)
841 /* PyTorch vectorized_outer_sum reduces four adjacent vectors together. */
842 for (; col + 63 < state_dim; col += 64) {
843 __m512 acc[4][4];
844 for (int level = 0; level < num_levels; ++level) {
845 for (int block = 0; block < 4; ++block) {
846 acc[level][block] = _mm512_setzero_ps();
847 }
848 }
849
850 int i = 0;
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);
859 }
860 }
861
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();
867 }
868 const int mask = level_mask << (level * level_power);
869 if ((i & mask) != 0) {
870 break;
871 }
872 }
873 }
874
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);
882 }
883 }
884
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]);
889 }
890 }
891 for (int block = 0; block < 4; ++block) {
892 _mm512_storeu_ps(output + col + block * 16, acc[0][block]);
893 }
894 }
895#elif defined(__AVX2__)
896 for (; col + 31 < state_dim; col += 32) {
897 __m256 acc[4][4];
898 for (int level = 0; level < num_levels; ++level) {
899 for (int block = 0; block < 4; ++block) {
900 acc[level][block] = _mm256_setzero_ps();
901 }
902 }
903
904 int i = 0;
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);
913 }
914 }
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();
920 }
921 const int mask = level_mask << (level * level_power);
922 if ((i & mask) != 0) {
923 break;
924 }
925 }
926 }
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);
934 }
935 }
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]);
940 }
941 }
942 for (int block = 0; block < 4; ++block) {
943 _mm256_storeu_ps(output + col + block * 8, acc[0][block]);
944 }
945 }
946#endif
947
948 /* Production Qwen dimensions are covered above. Keep a deterministic
949 * scalar fallback for uncommon tail widths. */
950 for (; col < state_dim; ++col) {
951 float sum = 0.0f;
952 for (int row = 0; row < state_dim; ++row) {
953 sum += matrix[(size_t)row * (size_t)state_dim + col] *
954 row_weights[row];
955 }
956 output[col] = sum;
957 }
958}
959
960#if defined(__GNUC__) && !defined(__clang__)
961__attribute__((optimize("fp-contract=off")))
962#endif
964 const float *q,
965 const float *k,
966 const float *v,
967 const float *g,
968 const float *beta,
969 const float *state_in,
970 float *state_out,
971 float *out,
972 float *debug_decayed_state,
973 float *debug_memory,
974 float *debug_delta,
975 int num_heads,
976 int group_count,
977 int state_dim,
978 float norm_eps)
979{
980 (void)norm_eps;
981 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out ||
982 num_heads <= 0 || group_count <= 0 || num_heads % group_count != 0 ||
983 state_dim <= 0 || state_dim > CK_DELTANET_MAX_STACK_DIM) {
984 return;
985 }
986
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);
991 float gate_values[CK_DELTANET_MAX_STACK_DIM];
992 float beta_values[CK_DELTANET_MAX_STACK_DIM];
993 float memory[CK_DELTANET_MAX_STACK_DIM];
994 float delta[CK_DELTANET_MAX_STACK_DIM];
995 float q_scaled[CK_DELTANET_MAX_STACK_DIM];
997 g, beta, gate_values, beta_values, num_heads);
998
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;
1007
1008 for (int col = 0; col < state_dim; ++col) {
1009 q_scaled[col] = q_head[col] / sqrt_dim;
1010 }
1011
1012 /* Materialize the decayed state before the separately ordered sum. */
1013 for (int row = 0; row < state_dim; ++row) {
1014 const size_t row_offset = (size_t)row * vector_stride;
1015 int col = 0;
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),
1021 gate8);
1022 _mm256_storeu_ps(state_cur + row_offset + (size_t)col, state);
1023 }
1024#endif
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;
1029 }
1030 }
1031
1033 state_cur, k_head, memory, state_dim);
1034
1035 if (debug_decayed_state) {
1036 memcpy(
1037 debug_decayed_state + (size_t)h * state_stride,
1038 state_cur,
1039 state_stride * sizeof(float));
1040 }
1041 if (debug_memory) {
1042 memcpy(
1043 debug_memory + (size_t)h * vector_stride,
1044 memory,
1045 vector_stride * sizeof(float));
1046 }
1047
1048 for (int col = 0; col < state_dim; ++col) {
1049 delta[col] = (v_head[col] - memory[col]) * beta_values[h];
1050 }
1051 if (debug_delta) {
1052 memcpy(
1053 debug_delta + (size_t)h * vector_stride,
1054 delta,
1055 vector_stride * sizeof(float));
1056 }
1057
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];
1061 int col = 0;
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),
1069 update);
1070 _mm256_storeu_ps(state_cur + row_offset + (size_t)col, state);
1071 }
1072#endif
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;
1077 }
1078 }
1079
1081 state_cur, q_scaled, out_head, state_dim);
1082
1083 for (int col = 0; col < state_dim; ++col) {
1084 out_head[col] = bf16_to_float(float_to_bf16(out_head[col]));
1085 }
1086 }
1087}
1088
1090 const float *k,
1091 const float *v,
1092 const float *g,
1093 const float *beta,
1094 const float *state_in,
1095 float *state_out,
1096 float *out,
1097 int num_heads,
1098 int group_count,
1099 int state_dim,
1100 float norm_eps)
1101{
1103 q, k, v, g, beta, state_in, state_out, out,
1104 NULL, NULL, NULL,
1105 num_heads, group_count, state_dim, norm_eps);
1106}
1107
1109 const float *q,
1110 const float *k,
1111 const float *v,
1112 const float *g,
1113 const float *beta,
1114 const float *state_in,
1115 float *state_out,
1116 float *out,
1117 float *decayed_state,
1118 float *memory,
1119 float *delta,
1120 int num_heads,
1121 int group_count,
1122 int state_dim,
1123 float norm_eps)
1124{
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);
1129}
1130
1132 const float *k,
1133 const float *v,
1134 const float *g,
1135 const float *beta,
1136 const float *state_in,
1137 float *state_out,
1138 float *out,
1139 int rows,
1140 int num_heads,
1141 int group_count,
1142 int state_dim,
1143 float norm_eps)
1144{
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) {
1148 return;
1149 }
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,
1161 state_out,
1162 out + (size_t) row * value_stride,
1163 num_heads, group_count, state_dim, norm_eps);
1164 }
1165}
1166
1168 const float *k,
1169 const float *v,
1170 const float *g,
1171 const float *beta,
1172 const float *state_in,
1173 float *state_out,
1174 float *out,
1175 int rows,
1176 int num_heads,
1177 int group_count,
1178 int state_dim,
1179 float norm_eps)
1180{
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) {
1184 return;
1185 }
1186#if defined(__AVX2__)
1187 (void)norm_eps;
1188 if (state_dim <= CK_DELTANET_LLAMA_CHUNK_MAX_DIM) {
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);
1193 }
1194 return;
1195 }
1196#endif
1198 q, k, v, g, beta, state_in, state_out, out,
1199 rows, num_heads, group_count, state_dim, norm_eps);
1200}
1201
1203 const float *k,
1204 const float *v,
1205 const float *g,
1206 const float *beta,
1207 const float *state_in,
1208 float *state_out,
1209 float *out,
1210 int rows,
1211 int num_heads,
1212 int group_count,
1213 int head,
1214 int state_dim)
1215{
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 ||
1220 state_dim <= 0 || state_dim > CK_DELTANET_LLAMA_CHUNK_MAX_DIM) {
1221 return;
1222 }
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);
1226#else
1227 (void)q;
1228 (void)k;
1229 (void)v;
1230 (void)g;
1231 (void)beta;
1232 (void)state_in;
1233 (void)state_out;
1234 (void)out;
1235 (void)rows;
1236 (void)num_heads;
1237 (void)group_count;
1238 (void)head;
1239 (void)state_dim;
1240#endif
1241}
1242
1244 const float *q,
1245 const float *k,
1246 const float *v,
1247 const float *g,
1248 const float *beta,
1249 const float *state_in,
1250 float *state_out,
1251 float *out,
1252 int rows,
1253 int num_heads,
1254 int group_count,
1255 int state_dim,
1256 float norm_eps)
1257{
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) {
1261 return;
1262 }
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,
1274 state_out,
1275 out + (size_t)row * value_stride,
1276 num_heads, group_count, state_dim, norm_eps);
1277 }
1278}
1279
1281 const float *k,
1282 const float *v,
1283 const float *g,
1284 const float *beta,
1285 const float *state_in,
1286 float *state_out,
1287 float *out,
1288 int num_heads,
1289 int state_dim,
1290 float norm_eps)
1291{
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;
1295
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;
1303
1304 const float beta_s = ck_deltanet_sigmoidf(beta[h]);
1305 const float gate = expf(g[h]);
1306
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;
1311 }
1312 }
1313
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;
1319 }
1320
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;
1325 }
1326 }
1327
1328 for (int col = 0; col < state_dim; ++col) {
1329 float acc = 0.0f;
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;
1333 }
1334 out_head[col] = acc;
1335 }
1336 }
1337}
1338
1340 const float *d_state_out,
1341 const float *q,
1342 const float *k,
1343 const float *v,
1344 const float *g,
1345 const float *beta,
1346 const float *state_in,
1347 const float *state_out,
1348 float *d_q,
1349 float *d_k,
1350 float *d_v,
1351 float *d_g,
1352 float *d_beta,
1353 float *d_state_in,
1354 int num_heads,
1355 int state_dim,
1356 float norm_eps)
1357{
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;
1361
1362 float q_hat[CK_DELTANET_MAX_STACK_DIM];
1363 float k_hat[CK_DELTANET_MAX_STACK_DIM];
1364 float kv_mem[CK_DELTANET_MAX_STACK_DIM];
1365 float delta[CK_DELTANET_MAX_STACK_DIM];
1366 float d_q_hat[CK_DELTANET_MAX_STACK_DIM];
1367 float d_k_hat[CK_DELTANET_MAX_STACK_DIM];
1368 float d_mem[CK_DELTANET_MAX_STACK_DIM];
1369
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;
1382
1383 const float beta_s = ck_deltanet_sigmoidf(beta[h]);
1384 const float gate = expf(g[h]);
1385
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;
1390
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];
1394 kv_mem[i] = 0.0f;
1395 d_q_hat[i] = 0.0f;
1396 d_k_hat[i] = 0.0f;
1397 d_mem[i] = 0.0f;
1398 d_v_head[i] = 0.0f;
1399 qk_dot += q_hat[i] * k_hat[i];
1400 }
1401
1402 for (int col = 0; col < state_dim; ++col) {
1403 float mem = 0.0f;
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];
1406 }
1407 kv_mem[col] = mem;
1408 delta[col] = (v_head[col] - mem) * beta_s;
1409 out_delta_dot += d_out_head[col] * delta[col];
1410 }
1411
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];
1420 }
1421 d_q_hat[row] = dq_acc;
1422 d_k_hat[row] = dk_acc;
1423 }
1424
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];
1429 }
1430
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]);
1434 }
1435
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];
1441 }
1442 d_k_hat[row] += s_dm_acc;
1443 }
1444
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];
1453 }
1454 }
1455
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];
1459 }
1460
1461 d_g[h] = gate_acc * gate;
1462 d_beta[h] = beta_acc * beta_s * (1.0f - beta_s);
1463 }
1464}
1465
1466#if defined(__AVX__)
1467static void ck_deltanet_scale_rows_avx(const float *src, float *dst, int dim, float scale)
1468{
1469 const __m256 scale_v = _mm256_set1_ps(scale);
1470 int i = 0;
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));
1474 }
1475 for (; i < dim; ++i) {
1476 dst[i] = src[i] * scale;
1477 }
1478}
1479
1480void gated_deltanet_autoregressive_forward_avx(const float *q,
1481 const float *k,
1482 const float *v,
1483 const float *g,
1484 const float *beta,
1485 const float *state_in,
1486 float *state_out,
1487 float *out,
1488 int num_heads,
1489 int state_dim,
1490 float norm_eps)
1491{
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;
1495
1496 float q_hat[CK_DELTANET_MAX_STACK_DIM];
1497 float k_hat[CK_DELTANET_MAX_STACK_DIM];
1498 float kv_mem[CK_DELTANET_MAX_STACK_DIM];
1499 float delta[CK_DELTANET_MAX_STACK_DIM];
1500
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;
1508
1509 const float gate = expf(g[h]);
1510 const float beta_s = ck_deltanet_sigmoidf(beta[h]);
1511
1512 /* q and k arrive pre-normalized by recurrent_qk_l2_norm. */
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);
1515
1516 const __m256 beta_v = _mm256_set1_ps(beta_s);
1517 const __m256 zero_v = _mm256_setzero_ps();
1518
1519 int col = 0;
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);
1523 }
1524 for (; col < state_dim; ++col) {
1525 kv_mem[col] = 0.0f;
1526 out_head[col] = 0.0f;
1527 }
1528
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);
1533
1534 col = 0;
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);
1542 }
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];
1547 }
1548 }
1549
1550 col = 0;
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);
1556 }
1557 for (; col < state_dim; ++col) {
1558 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1559 }
1560
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]);
1565
1566 col = 0;
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);
1575 }
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];
1580 }
1581 }
1582 }
1583}
1584#endif
1585
1586#if defined(__AVX2__)
1587static inline __m256 ck_deltanet_fmadd256(__m256 a, __m256 b, __m256 c)
1588{
1589#if defined(__FMA__)
1590 return _mm256_fmadd_ps(a, b, c);
1591#else
1592 return _mm256_add_ps(_mm256_mul_ps(a, b), c);
1593#endif
1594}
1595
1596void gated_deltanet_autoregressive_forward_avx2(const float *q,
1597 const float *k,
1598 const float *v,
1599 const float *g,
1600 const float *beta,
1601 const float *state_in,
1602 float *state_out,
1603 float *out,
1604 int num_heads,
1605 int state_dim,
1606 float norm_eps)
1607{
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;
1611
1612 float q_hat[CK_DELTANET_MAX_STACK_DIM];
1613 float k_hat[CK_DELTANET_MAX_STACK_DIM];
1614 float kv_mem[CK_DELTANET_MAX_STACK_DIM];
1615 float delta[CK_DELTANET_MAX_STACK_DIM];
1616
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;
1624
1625 const float gate = expf(g[h]);
1626 const float beta_s = ck_deltanet_sigmoidf(beta[h]);
1627
1628 /* q and k arrive pre-normalized by recurrent_qk_l2_norm. */
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);
1631
1632 const __m256 beta_v = _mm256_set1_ps(beta_s);
1633 const __m256 zero_v = _mm256_setzero_ps();
1634
1635 int col = 0;
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);
1639 }
1640 for (; col < state_dim; ++col) {
1641 kv_mem[col] = 0.0f;
1642 out_head[col] = 0.0f;
1643 }
1644
1645 int row = 0;
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);
1653
1654 col = 0;
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);
1666 }
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];
1673 }
1674 }
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);
1679 col = 0;
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);
1687 }
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];
1692 }
1693 }
1694
1695 col = 0;
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);
1701 }
1702 for (; col < state_dim; ++col) {
1703 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1704 }
1705
1706 row = 0;
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]);
1714
1715 col = 0;
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);
1728 }
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];
1735 }
1736 }
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]);
1741 col = 0;
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);
1750 }
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];
1755 }
1756 }
1757 }
1758}
1759#endif
1760
1761#if defined(__AVX512F__)
1762static inline __m512 ck_deltanet_madd512(__m512 a, __m512 b, __m512 c)
1763{
1764 return _mm512_add_ps(_mm512_mul_ps(a, b), c);
1765}
1766
1767static void ck_deltanet_scale_rows_avx512(const float *src, float *dst, int dim, float scale)
1768{
1769 const __m512 scale_v = _mm512_set1_ps(scale);
1770 int i = 0;
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));
1774 }
1775 for (; i < dim; ++i) {
1776 dst[i] = src[i] * scale;
1777 }
1778}
1779
1780void gated_deltanet_autoregressive_forward_avx512(const float *q,
1781 const float *k,
1782 const float *v,
1783 const float *g,
1784 const float *beta,
1785 const float *state_in,
1786 float *state_out,
1787 float *out,
1788 int num_heads,
1789 int state_dim,
1790 float norm_eps)
1791{
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;
1795
1796 float q_hat[CK_DELTANET_MAX_STACK_DIM];
1797 float k_hat[CK_DELTANET_MAX_STACK_DIM];
1798 float kv_mem[CK_DELTANET_MAX_STACK_DIM];
1799 float delta[CK_DELTANET_MAX_STACK_DIM];
1800
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;
1808
1809 const float gate = expf(g[h]);
1810 const float beta_s = ck_deltanet_sigmoidf(beta[h]);
1811
1812 /* q and k arrive pre-normalized by recurrent_qk_l2_norm. */
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);
1815
1816 const __m512 beta_v = _mm512_set1_ps(beta_s);
1817 const __m512 zero_v = _mm512_setzero_ps();
1818
1819 int col = 0;
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);
1823 }
1824 for (; col < state_dim; ++col) {
1825 kv_mem[col] = 0.0f;
1826 out_head[col] = 0.0f;
1827 }
1828
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);
1833 col = 0;
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);
1841 }
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];
1846 }
1847 }
1848
1849 col = 0;
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);
1855 }
1856 for (; col < state_dim; ++col) {
1857 delta[col] = (v_head[col] - kv_mem[col]) * beta_s;
1858 }
1859
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]);
1864 col = 0;
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);
1873 }
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];
1878 }
1879 }
1880 }
1881}
1882#endif
1883
1885{
1886 const char *env = getenv("CK_DELTANET_FORCE_REF");
1887 return env && atoi(env) != 0;
1888}
1889
1891{
1893 return "REF";
1894 }
1895#if defined(__AVX512F__)
1896 return "AVX512";
1897#elif defined(__AVX2__)
1898 return "AVX2";
1899#elif defined(__AVX__)
1900 return "AVX";
1901#else
1902 return "REF";
1903#endif
1904}
1905
1907 const float *k,
1908 const float *v,
1909 const float *g,
1910 const float *beta,
1911 const float *state_in,
1912 float *state_out,
1913 float *out,
1914 int num_heads,
1915 int state_dim,
1916 float norm_eps)
1917{
1918 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out) {
1919 return;
1920 }
1921 if (num_heads <= 0 || state_dim <= 0) {
1922 return;
1923 }
1924
1925 /*
1926 * q and k arrive pre-normalized by recurrent_qk_l2_norm, so the
1927 * ISA-specialized kernels can follow the same contract as the scalar ref.
1928 */
1931 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1932 return;
1933 }
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);
1943#else
1945 q, k, v, g, beta, state_in, state_out, out, num_heads, state_dim, norm_eps);
1946#endif
1947}
1948
1950 const float *k,
1951 const float *v,
1952 const float *g,
1953 const float *beta,
1954 const float *state_in,
1955 float *state_out,
1956 float *out,
1957 int rows,
1958 int num_heads,
1959 int state_dim,
1960 float norm_eps)
1961{
1962 if (!q || !k || !v || !g || !beta || !state_in || !state_out || !out) {
1963 return;
1964 }
1965 if (rows <= 0 || num_heads <= 0 || state_dim <= 0) {
1966 return;
1967 }
1968
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,
1979 row_state_in,
1980 state_out,
1981 out + (size_t)row * vector_stride,
1982 num_heads,
1983 state_dim,
1984 norm_eps);
1985 }
1986}
1987
1989 const float *d_state_out,
1990 const float *q,
1991 const float *k,
1992 const float *v,
1993 const float *g,
1994 const float *beta,
1995 const float *state_in,
1996 const float *state_out,
1997 float *d_q,
1998 float *d_k,
1999 float *d_v,
2000 float *d_g,
2001 float *d_beta,
2002 float *d_state_in,
2003 int num_heads,
2004 int state_dim,
2005 float norm_eps)
2006{
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) {
2009 return;
2010 }
2011 if (num_heads <= 0 || state_dim <= 0 || state_dim > CK_DELTANET_MAX_STACK_DIM) {
2012 return;
2013 }
2014
2016 d_out,
2017 d_state_out,
2018 q,
2019 k,
2020 v,
2021 g,
2022 beta,
2023 state_in,
2024 state_out,
2025 d_q,
2026 d_k,
2027 d_v,
2028 d_g,
2029 d_beta,
2030 d_state_in,
2031 num_heads,
2032 state_dim,
2033 norm_eps);
2034}
#define RTLD_DEFAULT
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90
static float bf16_to_float(uint16_t v)
Definition bf16_utils.h:38
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)
#define C(color)
Definition show_config.c:39
int32_t int32_t int32_t int32_t int32_t mask
Definition tokenizer.h:234
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)
const char * token
Definition tokenizer.h:307