FP32 Gated DeltaNet kernels for Qwen3.5-style recurrent attention. More...
#include "bf16_utils.h"#include "ckernel_engine.h"#include <dlfcn.h>#include <math.h>#include <pthread.h>#include <stdio.h>#include <stddef.h>#include <stdlib.h>#include <string.h>Go to the source code of this file.
Macros | |
| #define | CK_DELTANET_LLAMA_CHUNK_MAX_DIM 256 |
| #define | CK_DELTANET_LLAMA_CHUNK_SIZE 64 |
| #define | CK_DELTANET_MAX_STACK_DIM 4096 |
| #define | CK_DELTANET_NOINLINE |
Typedefs | |
| typedef float(* | ck_deltanet_libm_f32_fn) (float) |
| typedef void(* | ck_deltanet_mkl_vsexp_fn) (int, const float *, float *) |
Functions | |
| static void | ck_bind_deltanet_llama_libm (void) |
| static void | ck_bind_deltanet_pytorch_primitives (void) |
| static int | ck_deltanet_ceil_log2 (int value) |
| static int | ck_deltanet_force_ref (void) |
| static float | ck_deltanet_llama_sigmoidf (float x) |
| static void | ck_deltanet_pytorch_gate_values (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) |
| static float | ck_deltanet_sigmoidf (float x) |
| 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) |
| 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) |
| 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_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) |
| const char * | gated_deltanet_impl_name (void) |
| 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) |
| 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 | 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 | 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 | 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) |
| 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) |
| 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) |
| 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) |
| void | gated_deltanet_pytorch_gate_values_debug (const float *g, const float *beta, float *gate_values, float *beta_values, int num_heads) |
| 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_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 | 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_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) |
Variables | |
| static void * | ck_deltanet_libm_handle = NULL |
| static pthread_once_t | ck_deltanet_libm_once = PTHREAD_ONCE_INIT |
| static ck_deltanet_libm_f32_fn | ck_deltanet_llama_expf = NULL |
| static void * | ck_deltanet_mkl_handle = NULL |
| static pthread_once_t | ck_deltanet_pytorch_primitives_once = PTHREAD_ONCE_INIT |
| static ck_deltanet_mkl_vsexp_fn | ck_deltanet_pytorch_vsexp = NULL |
FP32 Gated DeltaNet kernels for Qwen3.5-style recurrent attention.
After changes: make test && make llamacpp-parity-full
This file implements the single-token recurrent update used by the qwen3next / Gated DeltaNet path in llama.cpp.
Per head, matching llama.cpp qwen35/qwen3next autoregressive DeltaNet: q_scaled = q / sqrt(state_dim) // q and k arrive pre-normalized k_hat = k beta_s = sigmoid(beta) gate = exp(g) S = gate * S_prev kv_mem = S^T * k_hat delta = (v - kv_mem) * beta_s S_new = S + outer(k_hat, delta) out = S_new^T * q_scaled
Design:
Definition in file deltanet_kernels.c.
| #define CK_DELTANET_LLAMA_CHUNK_MAX_DIM 256 |
Definition at line 54 of file deltanet_kernels.c.
| #define CK_DELTANET_LLAMA_CHUNK_SIZE 64 |
Definition at line 53 of file deltanet_kernels.c.
| #define CK_DELTANET_MAX_STACK_DIM 4096 |
Definition at line 52 of file deltanet_kernels.c.
| #define CK_DELTANET_NOINLINE |
Definition at line 59 of file deltanet_kernels.c.
| typedef float(* ck_deltanet_libm_f32_fn) (float) |
Definition at line 62 of file deltanet_kernels.c.
| typedef void(* ck_deltanet_mkl_vsexp_fn) (int, const float *, float *) |
Definition at line 99 of file deltanet_kernels.c.
|
static |
Definition at line 67 of file deltanet_kernels.c.
References ck_deltanet_libm_handle, and ck_deltanet_llama_expf.
Referenced by ck_deltanet_llama_sigmoidf(), gated_deltanet_llama_avx2_grouped_forward_impl(), and gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
|
static |
Definition at line 104 of file deltanet_kernels.c.
References ck_deltanet_mkl_handle, ck_deltanet_pytorch_vsexp, and RTLD_DEFAULT.
Referenced by ck_deltanet_pytorch_gate_values().
|
static |
Definition at line 812 of file deltanet_kernels.c.
Referenced by ck_deltanet_pytorch_outer_sum().
|
static |
Definition at line 1884 of file deltanet_kernels.c.
Referenced by gated_deltanet_autoregressive_forward(), and gated_deltanet_impl_name().
|
inlinestatic |
Definition at line 87 of file deltanet_kernels.c.
References ck_bind_deltanet_llama_libm(), ck_deltanet_libm_once, and ck_deltanet_llama_expf.
Referenced by gated_deltanet_llama_avx2_grouped_forward_impl(), and gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
|
static |
Definition at line 136 of file deltanet_kernels.c.
References bf16_to_float(), ck_bind_deltanet_pytorch_primitives(), ck_deltanet_pytorch_primitives_once, ck_deltanet_pytorch_vsexp, ck_deltanet_sigmoidf(), and float_to_bf16().
Referenced by gated_deltanet_pytorch_gate_values_debug(), and gated_deltanet_pytorch_grouped_bf16_forward_impl().
|
static |
Definition at line 826 of file deltanet_kernels.c.
References ck_deltanet_ceil_log2(), and mask.
Referenced by gated_deltanet_pytorch_grouped_bf16_forward_impl().
|
inlinestatic |
Definition at line 82 of file deltanet_kernels.c.
Referenced by ck_deltanet_pytorch_gate_values(), gated_deltanet_autoregressive_backward_ref(), gated_deltanet_autoregressive_forward_ref(), and gated_deltanet_llama_avx2_grouped_forward_impl().
| 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 | ||
| ) |
Definition at line 1988 of file deltanet_kernels.c.
References CK_DELTANET_MAX_STACK_DIM, and gated_deltanet_autoregressive_backward_ref().
| 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 | ||
| ) |
Definition at line 1339 of file deltanet_kernels.c.
References CK_DELTANET_MAX_STACK_DIM, and ck_deltanet_sigmoidf().
Referenced by gated_deltanet_autoregressive_backward().
| 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 | ||
| ) |
Definition at line 1906 of file deltanet_kernels.c.
References ck_deltanet_force_ref(), ck_strict_parity_enabled(), and gated_deltanet_autoregressive_forward_ref().
Referenced by ck_test_gated_deltanet_autoregressive(), and gated_deltanet_prefill_forward().
| 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 | ||
| ) |
Definition at line 1280 of file deltanet_kernels.c.
References ck_deltanet_sigmoidf().
Referenced by gated_deltanet_autoregressive_forward(), and gated_deltanet_llama_avx2_grouped_forward_impl().
| const char * gated_deltanet_impl_name | ( | void | ) |
Definition at line 1890 of file deltanet_kernels.c.
References ck_deltanet_force_ref(), and ck_strict_parity_enabled().
| 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 | ||
| ) |
Definition at line 767 of file deltanet_kernels.c.
References gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
Referenced by gated_deltanet_llama_avx2_prefill_forward().
| 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 | ||
| ) |
Definition at line 790 of file deltanet_kernels.c.
References gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
|
static |
Definition at line 574 of file deltanet_kernels.c.
References bf16_to_float(), ck_bind_deltanet_llama_libm(), ck_deltanet_libm_once, ck_deltanet_llama_expf, ck_deltanet_llama_sigmoidf(), CK_DELTANET_MAX_STACK_DIM, ck_deltanet_sigmoidf(), float_to_bf16(), and gated_deltanet_autoregressive_forward_ref().
|
static |
Definition at line 687 of file deltanet_kernels.c.
References ck_bind_deltanet_llama_libm(), ck_deltanet_libm_once, ck_deltanet_llama_expf, ck_deltanet_llama_sigmoidf(), and CK_DELTANET_MAX_STACK_DIM.
Referenced by gated_deltanet_llama_avx2_forward(), and gated_deltanet_llama_avx2_forward_head_range().
| 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 | ||
| ) |
Definition at line 1131 of file deltanet_kernels.c.
References gated_deltanet_llama_avx2_forward().
Referenced by gated_deltanet_llama_chunk64_prefill_forward().
| 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 | ||
| ) |
Definition at line 1202 of file deltanet_kernels.c.
References CK_DELTANET_LLAMA_CHUNK_MAX_DIM.
| 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 | ||
| ) |
Definition at line 1167 of file deltanet_kernels.c.
References CK_DELTANET_LLAMA_CHUNK_MAX_DIM, and gated_deltanet_llama_avx2_prefill_forward().
| 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 | ||
| ) |
Definition at line 1949 of file deltanet_kernels.c.
References gated_deltanet_autoregressive_forward().
| void gated_deltanet_pytorch_gate_values_debug | ( | const float * | g, |
| const float * | beta, | ||
| float * | gate_values, | ||
| float * | beta_values, | ||
| int | num_heads | ||
| ) |
Definition at line 178 of file deltanet_kernels.c.
References CK_DELTANET_MAX_STACK_DIM, and ck_deltanet_pytorch_gate_values().
| 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 | ||
| ) |
Definition at line 1089 of file deltanet_kernels.c.
References gated_deltanet_pytorch_grouped_bf16_forward_impl().
Referenced by gated_deltanet_pytorch_grouped_bf16_prefill_forward().
| 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 | ||
| ) |
Definition at line 1108 of file deltanet_kernels.c.
References gated_deltanet_pytorch_grouped_bf16_forward_impl().
|
static |
Definition at line 963 of file deltanet_kernels.c.
References bf16_to_float(), CK_DELTANET_MAX_STACK_DIM, ck_deltanet_pytorch_gate_values(), ck_deltanet_pytorch_outer_sum(), and float_to_bf16().
Referenced by gated_deltanet_pytorch_grouped_bf16_forward(), and gated_deltanet_pytorch_grouped_bf16_forward_debug().
| 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 | ||
| ) |
Definition at line 1243 of file deltanet_kernels.c.
References gated_deltanet_pytorch_grouped_bf16_forward().
|
static |
Definition at line 64 of file deltanet_kernels.c.
Referenced by ck_bind_deltanet_llama_libm().
|
static |
Definition at line 65 of file deltanet_kernels.c.
Referenced by ck_deltanet_llama_sigmoidf(), gated_deltanet_llama_avx2_grouped_forward_impl(), and gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
|
static |
Definition at line 63 of file deltanet_kernels.c.
Referenced by ck_bind_deltanet_llama_libm(), ck_deltanet_llama_sigmoidf(), gated_deltanet_llama_avx2_grouped_forward_impl(), and gated_deltanet_llama_avx2_grouped_forward_transposed_impl().
|
static |
Definition at line 101 of file deltanet_kernels.c.
Referenced by ck_bind_deltanet_pytorch_primitives().
|
static |
Definition at line 102 of file deltanet_kernels.c.
Referenced by ck_deltanet_pytorch_gate_values().
|
static |
Definition at line 100 of file deltanet_kernels.c.
Referenced by ck_bind_deltanet_pytorch_primitives(), and ck_deltanet_pytorch_gate_values().