API Reference

Complete API documentation for C-Kernel-Engine kernels. All functions are exported from libckernel_engine.so and can be called from C or via Python ctypes.

Auto-generated from source
This documentation is extracted from the C header files using Doxygen. Functions marked Forward compute activations, and Backward compute gradients.

Quick Reference

Include Header

#include "ckernel_engine.h"

Link Library

-lckernel_engine

Python ctypes

lib = ctypes.CDLL("libckernel_engine.so")

Memory Layouts

All kernels use consistent memory layouts optimized for cache efficiency:

Buffer Layout Description
input/output [B, T, D] Batch × Tokens × Embedding dimension
Q [H, T, d_k] num_heads × Tokens × head_dim (head-major)
K, V [H_kv, T, d_k] num_kv_heads × Tokens × head_dim (for GQA)
scores [H, T, T] num_heads × query_tokens × key_tokens
weights [out, in] Row-major weight matrices

Kernel Functions

GEMM (Matrix Multiplication) (258)
attention_forward_causal_head_major_gqa_flash_strided_gemma4 Forward
void attention_forward_causal_head_major_gqa_flash_strided_gemma4(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_gemma4_token_output Forward
void attention_forward_causal_head_major_gqa_flash_strided_gemma4_token_output(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4 Forward
void attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl Forward
void attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window, int output_token_major)

Flash attention decode with sliding window Testtest_attention.py::TestAttentionForward::test_sliding_window_decode

attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_token_output Forward
void attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_token_output(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window)

Forward pass computation

attention_forward_causal_head_major_shared_kv_gemma4 Forward
void attention_forward_causal_head_major_shared_kv_gemma4(const float * q, float * output, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_shared_kv_sliding_gemma4 Forward
void attention_forward_causal_head_major_shared_kv_sliding_gemma4(const float * q, float * output, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window)

Forward pass computation

attention_forward_chunk_head_major_gqa_flash_gemma4 Forward
void attention_forward_chunk_head_major_gqa_flash_gemma4(const float * q_chunk, const float * k_cache, const float * v_cache, float * out_chunk, int num_heads, int num_kv_heads, int q_tokens, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_gemma4 Forward
void attention_forward_decode_head_major_gqa_flash_gemma4(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_sliding_gemma4 Forward
void attention_forward_decode_head_major_gqa_flash_sliding_gemma4(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window)

Forward pass computation

attention_forward_decode_head_major_shared_kv_gemma4 Forward
void attention_forward_decode_head_major_shared_kv_gemma4(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_decode_head_major_shared_kv_sliding_gemma4 Forward
void attention_forward_decode_head_major_shared_kv_sliding_gemma4(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window)

Forward pass computation

attention_forward_full_head_major_gqa_flash_strided_gemma4 Forward
void attention_forward_full_head_major_gqa_flash_strided_gemma4(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4 Forward
void attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int visual_start, int visual_tokens)

Forward pass computation

attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4_impl Forward
void attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4_impl(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int visual_start, int visual_tokens, int output_token_major)

Forward pass computation

attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4_token_output Forward
void attention_forward_mixed_visual_chunk_head_major_gqa_flash_strided_gemma4_token_output(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int visual_start, int visual_tokens)

Forward pass computation

ck_gemm_add_bias Forward
void ck_gemm_add_bias(float * C, const float * bias, int M, int N)

ck_gemm_bf16_amx_available Forward
int ck_gemm_bf16_amx_available(void)

ck_gemm_bf16_amx_work Forward
void ck_gemm_bf16_amx_work(int ith, int nth, void * opaque)

ck_gemm_bf16_fp32out_amx_raw Forward
int ck_gemm_bf16_fp32out_amx_raw(const uint16_t * A, const uint16_t * B, float * C, int M, int N, int K, int accumulate)

ck_gemm_bf16_native_work Forward
void ck_gemm_bf16_native_work(int ith, int nth, void * opaque)

ck_gemm_dynamic_schedule_enabled Forward
int ck_gemm_dynamic_schedule_enabled(void)

Return non-zero when independent GEMM tiles should use dynamic claiming.

ck_gemm_f16_input_fp16_work Forward
void ck_gemm_f16_input_fp16_work(int ith, int nth, void * opaque)

ck_gemm_f16_pick_active_threads Forward
int ck_gemm_f16_pick_active_threads(const ck_threadpool_t * pool, int M, int N, int K)

ck_gemm_f16_threadpool_enabled Forward
int ck_gemm_f16_threadpool_enabled(int M, int N, int K)

ck_gemm_nn_impl_probe_enabled Forward
int ck_gemm_nn_impl_probe_enabled(void)

ck_gemm_nt_bf16_exact_rows Forward
void ck_gemm_nt_bf16_exact_rows(int begin, int end, void * opaque)

ck_gemm_nt_bf16_storage_exact_rows Forward
void ck_gemm_nt_bf16_storage_exact_rows(int begin, int end, void * opaque)

ck_gemm_nt_f16_ggml_oracle Forward
int ck_gemm_nt_f16_ggml_oracle(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

ck_gemm_nt_f16_simd_lanes Forward
int ck_gemm_nt_f16_simd_lanes(void)

ck_gemm_nt_f32_llama_production_output Forward
void ck_gemm_nt_f32_llama_production_output(const float * A, const float * B, const float * bias, float * C, int M, int N, int K, int index)

ck_gemm_nt_fp32_exact_rows Forward
void ck_gemm_nt_fp32_exact_rows(int begin, int end, void * opaque)

ck_gemm_nt_head_major_q5_0 Forward
void ck_gemm_nt_head_major_q5_0(const float * attn_out, const void * wo, const float * bias, float * output, int tokens, int embed_dim, int num_heads, int head_dim)

Output projection from head-major attention (auto-dispatch)

ck_gemm_nt_head_major_q8_0 Forward
void ck_gemm_nt_head_major_q8_0(const float * attn_out, const void * wo, const float * bias, float * output, int tokens, int embed_dim, int num_heads, int head_dim)

Output projection from head-major attention (Q8_0 weights)

ck_gemm_nt_quant Forward
void ck_gemm_nt_quant(const float * A, const void * B, const float * bias, float * C, int M, int N, int K, CKDataType dtype)

ck_gemma4_dequant_q5_k_block Forward
void ck_gemma4_dequant_q5_k_block(const ck_gemma4_block_q5_K * block, float * out)

ck_gemma4_embed_range Forward
void ck_gemma4_embed_range(int begin, int end, void * opaque)

ck_gemma4_gelu Forward
float ck_gemma4_gelu(float x)

ck_gemma4_prepare_bf16_range Forward
void ck_gemma4_prepare_bf16_range(int begin, int end, void * opaque)

ck_gemma4_prepare_parallel Forward
void ck_gemma4_prepare_parallel(int tokens, ck_range_fn_t fn, ck_gemma4_prepare_args_t * args)

ck_gemma4_prepare_q5_range Forward
void ck_gemma4_prepare_q5_range(int begin, int end, void * opaque)

ck_gemma4_q5_k_value Forward
uint8_t ck_gemma4_q5_k_value(const ck_gemma4_block_q5_K * block, int subblock, int i)

ck_gemma4_rmsnorm_tmp Forward
void ck_gemma4_rmsnorm_tmp(const float * x, const float * gamma, float * out, int n, float eps)

ck_gemma4_unpack_q5_k_scales Forward
void ck_gemma4_unpack_q5_k_scales(const uint8_t * scales, uint8_t * sc, uint8_t * m)

ck_get_gemm_schedule Forward
int ck_get_gemm_schedule(void)

Return the configured process-wide GEMM scheduling policy.

ck_llama_regular_gemm_f16 Forward
float ck_llama_regular_gemm_f16(const float * probability, const float * value_column, int count)

ck_set_gemm_schedule Forward
int ck_set_gemm_schedule(int policy)

Set the process-wide GEMM tile scheduling policy.

ck_strict_consume_next_gemm_a Forward
const float * ck_strict_consume_next_gemm_a(size_t elems)

ck_strict_store_next_gemm_a Forward
void ck_strict_store_next_gemm_a(const float * data, size_t elems)

ck_test_gemm_q4_k Forward
void ck_test_gemm_q4_k(const void * weight_q4k, const float * input_f32, float * output, int rows, int cols, int n_tokens)

Q4_K GEMM - batched matrix multiply with quantized weights.

ck_test_gemm_q5_0 Forward
void ck_test_gemm_q5_0(const void * weight_q5_0, const float * input_f32, float * output, int rows, int cols, int n_tokens)

Test Q5_0 x Q8_0 GEMM (batch matrix multiply)

ck_test_gemm_q6_k Forward
void ck_test_gemm_q6_k(const void * weight_q6k, const float * input_f32, float * output, int rows, int cols, int n_tokens)

Test Q6_K x Q8_K GEMM (batch matrix multiply)

ck_test_gemm_q8_0 Forward
void ck_test_gemm_q8_0(const void * weight_q8_0, const float * input_f32, float * output, int rows, int cols, int n_tokens)

Test Q8_0 x Q8_0 GEMM (batch matrix multiply)

ck_train_gemm_backward_serial Backward
void ck_train_gemm_backward_serial(const float * d_output, const float * input, const float * W, float * d_input, float * d_W, float * d_b, int T, int aligned_in, int aligned_out)

Backward pass / gradient computation

ck_train_gemm_backward_work Backward
void ck_train_gemm_backward_work(int ith, int nth, void * argp)

Backward pass / gradient computation

ck_train_gemm_nn_compute_rows Forward
void ck_train_gemm_nn_compute_rows(const float * A, const float * B, const float * bias, float * C, int row_start, int row_end, int N, int K)

ck_train_gemm_nn_work Forward
void ck_train_gemm_nn_work(int ith, int nth, void * argp)

ck_train_gemm_nt_compute_rows Forward
void ck_train_gemm_nt_compute_rows(const float * A, const float * B, const float * bias, float * C, int row_start, int row_end, int N, int K)

ck_train_gemm_tn_compute_rows Forward
void ck_train_gemm_tn_compute_rows(const float * A, const float * B, const float * bias, float * C, int row_start, int row_end, int M, int N, int K)

ck_train_gemm_work Forward
void ck_train_gemm_work(int ith, int nth, void * argp)

ckernel_sgemm_native Forward
void ckernel_sgemm_native(int M, int N, int K, const float * A, int lda, const float * B, int ldb, const float * bias, float * C, int ldc)

Native GEMM backend that directly reuses the C-Transformer GEMM kernel.

compute_gemm_params Forward
void compute_gemm_params(const CPUInfo * cpu, GEMMParams * params)

fused_rmsnorm_gemm_2d_tiled Forward
void fused_rmsnorm_gemm_2d_tiled(const float * x, const float * gamma, const float * W, float * output, int seq_len, int hidden, int out_dim, float eps, float * x_norm_scratch)

Fused RMSNorm + single GEMM with 2D tiling (weight reuse)

gemm_avx512_parallel Forward
void gemm_avx512_parallel(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_backward_bf16_mixed Backward
void gemm_backward_bf16_mixed(const uint16_t * d_output, const uint16_t * input, const uint16_t * weight, float * d_input, float * d_weight, float * d_bias, int tokens, int in_dim, int out_dim)

Backward pass / gradient computation

gemm_backward_f32_train_parallel_dispatch Backward
void gemm_backward_f32_train_parallel_dispatch(const float * d_output, const float * input, const float * W, float * d_input, float * d_W, float * d_b, int T, int aligned_in, int aligned_out, int num_threads)

Backward pass / gradient computation

gemm_backward_f32_train_parallel_dispatch_v2 Backward
void gemm_backward_f32_train_parallel_dispatch_v2(const float * d_output, const float * input, const float * W, float * d_input, float * d_W, float * d_b, int T, int aligned_in, int aligned_out, int num_threads)

Backward pass / gradient computation

gemm_batch_int8_impl_name Forward
const char * gemm_batch_int8_impl_name(void)

Get the best implementation name for logging/debugging.

gemm_bf16_fp32out Forward
void gemm_bf16_fp32out(const uint16_t * A, const uint16_t * B, const float * bias, float * C, int M, int N, int K)

gemm_bias_gelu_fused Forward
void gemm_bias_gelu_fused(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_bias_relu_fused Forward
void gemm_bias_relu_fused(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_bias_silu_fused Forward
void gemm_bias_silu_fused(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_blocked_serial Forward
void gemm_blocked_serial(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_blocked_serial_bf16 Forward
void gemm_blocked_serial_bf16(const uint16_t * A, const uint16_t * B, const uint16_t * bias, uint16_t * C, int M, int N, int K)

gemm_blocked_serial_train_parallel_dispatch Forward
void gemm_blocked_serial_train_parallel_dispatch(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_f16 Forward
void gemm_f16(float * Y, const uint16_t * W, const float * X, int M, int N, int K)

Auto-dispatch GEMM based on available SIMD.

gemm_f16_backward Backward
void gemm_f16_backward(float * dX, const uint16_t * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_f16_input_fp16_ref Forward
void gemm_f16_input_fp16_ref(float * Y, const uint16_t * W, const float * X, int M, int N, int K)

gemm_f16_input_fp16_threadpool Forward
int gemm_f16_input_fp16_threadpool(float * Y, const uint16_t * W, const float * X, int M, int N, int K)

gemm_f16_ref Forward
void gemm_f16_ref(float * Y, const uint16_t * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with FP16 weights (scalar reference)

gemm_fine_grained_parallel Forward
void gemm_fine_grained_parallel(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_get_backend
const char * gemm_get_backend(void)

gemm_init_threads Forward
void gemm_init_threads(void)

gemm_microkernel Forward
void gemm_microkernel(const float * A, const float * B, float * C, int M, int N, int K, int B_transposed)

gemm_microkernel_blocked Forward
void gemm_microkernel_blocked(const float * A, const float * B, float * C, int M, int N, int K)

gemm_microkernel_blocked_bt Forward
void gemm_microkernel_blocked_bt(const float * A, const float * B, float * C, int M, int N, int K)

gemm_microkernel_edge Forward
void gemm_microkernel_edge(int m, int n, int K, const float * A, int lda, const float * B, int ldb, float * C, int ldc, int first_k)

gemm_microkernel_packed Forward
void gemm_microkernel_packed(const float * A, const float * B, float * C, int M, int N, int K)

gemm_microkernel_sequential Forward
void gemm_microkernel_sequential(const float * A, const float * B, float * C, int M, int N, int K)

gemm_naive_parallel Forward
void gemm_naive_parallel(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_naive_serial_double Forward
void gemm_naive_serial_double(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_naive_serial_float Forward
void gemm_naive_serial_float(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_avx512 Forward
void gemm_nn_avx512(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_avx512_probe Forward
void gemm_nn_avx512_probe(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_bf16 Forward
void gemm_nn_bf16(const uint16_t * A, const uint16_t * B, const uint16_t * bias, uint16_t * C, int M, int N, int K)

gemm_nn_blocked Forward
void gemm_nn_blocked(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_parallel Forward
void gemm_nn_parallel(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_serial_double Forward
void gemm_nn_serial_double(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nn_simd Forward
void gemm_nn_simd(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nt Forward
void gemm_nt(const float * input, const float * weight, float * output, int rows, int cols, int common)

gemm_nt_bf16 Forward
void gemm_nt_bf16(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_amx_bf16_storage Forward
void gemm_nt_bf16_amx_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_amx_bf16_storage_workspace Forward
void gemm_nt_bf16_amx_bf16_storage_workspace(const float * A, const void * B, const float * bias, float * C, int M, int N, int K, uint16_t * a_bf16, size_t a_bf16_bytes)

gemm_nt_bf16_bf16_storage Forward
void gemm_nt_bf16_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_bf16_storage_parallel_dispatch Forward
void gemm_nt_bf16_bf16_storage_parallel_dispatch(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_bf16_storage_row_range Forward
void gemm_nt_bf16_bf16_storage_row_range(const float * A, const void * B, const float * bias, float * C, int M, int N, int K, int row_begin, int row_end)

gemm_nt_bf16_native_bf16_storage Forward
void gemm_nt_bf16_native_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_parallel_dispatch Forward
void gemm_nt_bf16_parallel_dispatch(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_prefill_shape_safe_bf16_storage Forward
void gemm_nt_bf16_prefill_shape_safe_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_prefill_shape_safe_bf16_storage_workspace Forward
void gemm_nt_bf16_prefill_shape_safe_bf16_storage_workspace(const float * A, const void * B, const float * bias, float * C, int M, int N, int K, uint16_t * a_bf16, size_t a_bf16_bytes)

gemm_nt_bf16_pytorch_onednn_3_12_brgemm_bf16_storage Forward
void gemm_nt_bf16_pytorch_onednn_3_12_brgemm_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage Forward
void gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage_impl Forward
void gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage_impl(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_bf16_row_range Forward
void gemm_nt_bf16_row_range(const float * A, const void * B, const float * bias, float * C, int M, int N, int K, int row_begin, int row_end)

gemm_nt_f16 Forward
void gemm_nt_f16(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

NT GEMM wrapper for FP16 weights with the engine's standard ABI.

gemm_nt_f16_clipped Forward
void gemm_nt_f16_clipped(const float * A, const void * B, const float * bias, const float * input_min, const float * input_max, const float * output_min, const float * output_max, float * C, int M, int N, int K)

gemm_nt_f16_ggml_strict Forward
int gemm_nt_f16_ggml_strict(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_f32_llama_production Forward
void gemm_nt_f32_llama_production(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_f32_llama_production_output_range Forward
void gemm_nt_f32_llama_production_output_range(const float * A, const float * B, const float * bias, float * C, int M, int N, int K, int output_begin, int output_end)

gemm_nt_fp32_exact_parallel_dispatch Forward
void gemm_nt_fp32_exact_parallel_dispatch(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_matvec_parallel Forward
void gemm_nt_matvec_parallel(const float * A, const float * B, const float * bias, float * C, int N, int K)

gemm_nt_q4_0 Forward
void gemm_nt_q4_0(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.

gemm_nt_q4_1 Forward
void gemm_nt_q4_1(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

GEMM with transposed Q4_1 weights: C = A @ B^T.

gemm_nt_q4_k Forward
void gemm_nt_q4_k(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_q8_k Forward
void gemm_nt_q4_k_packed_meta_q8_k(const void * A_q8, const void * B_packed, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_q8_k_threaded Forward
void gemm_nt_q4_k_packed_meta_q8_k_threaded(const void * A_q8, const void * B_packed, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit Forward
void gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit(const void * A_q8, const void * B_packed, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_meta_q8_k_tile Forward
void gemm_nt_q4_k_packed_meta_q8_k_tile(const void * A_q8, const void * B_packed, const float * bias, float * C, int M, int N, int K, int m0, int m1, int n0, int n1)

gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni Forward
void gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni(const void * A_q8, const void * B_packed_x16, const float * bias, float * C, int M, int D, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order Forward
void gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse Forward
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse(const void * A_q8, const void * B_packed_x16, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile Forward
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile(const void * A_q8, const void * B_packed_x16, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit Forward
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit(const void * A_q8, const void * B_packed_x8, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_u8_q8_k Forward
void gemm_nt_q4_k_packed_u8_q8_k(const void * A_q8, const void * B_packed, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile Forward
void gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile(const void * A_q8, const void * B_packed_u8_x16, const float * bias, float * C, int M, int N, int K, int tile_m, int active_threads)

gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order Forward
void gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order(const void * A_q8, const void * B_packed_x16, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m Forward
void gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m(const void * A_q8, const void * B_packed_vnni_x16, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m Forward
void gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m(const void * A_q8, const void * B_packed_vnni_x8, const float * bias, float * C, int M, int N, int K, int active_threads)

gemm_nt_q4_k_q8_k Forward
void gemm_nt_q4_k_q8_k(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni Forward
void gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni(const void * A_q8, const void * B_gate_up, const float * bias, float * C, int M, int D, int K, int threads)

gemm_nt_q4_k_q8_k_pairwise_split_min_parallel_dispatch Forward
void gemm_nt_q4_k_q8_k_pairwise_split_min_parallel_dispatch(const void * input, const void * weight, const float * bias, float * output, int rows, int output_dim, int input_dim)

gemm_nt_q5_0 Forward
void gemm_nt_q5_0(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_q8_0 Forward
void gemm_nt_q5_0_q8_0(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K)

Batch GEMM with Q5_0 weights and Q8_0 activations for prefill.

gemm_nt_q5_0_q8_0_m2n4 Forward
void gemm_nt_q5_0_q8_0_m2n4(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_q8_0_m2n4_tile Forward
void gemm_nt_q5_0_q8_0_m2n4_tile(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K, int ldc)

gemm_nt_q5_0_q8_0_m4n2 Forward
void gemm_nt_q5_0_q8_0_m4n2(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_q8_0_m4n2_tile Forward
void gemm_nt_q5_0_q8_0_m4n2_tile(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K, int ldc)

gemm_nt_q5_0_q8_0_parallel_dispatch Forward
void gemm_nt_q5_0_q8_0_parallel_dispatch(const void * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_q8_0_ref Forward
void gemm_nt_q5_0_q8_0_ref(const void * A, const void * B, float * C, int M, int N, int K)

Dispatcher for gemm_nt_q8_0_q8_0.

gemm_nt_q5_0_q8_0_unroll_avx Forward
void gemm_nt_q5_0_q8_0_unroll_avx(const void * A_q8, const void * B_q5, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_ref Forward
void gemm_nt_q5_0_ref(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

GEMM with transposed Q5_0 weights: C = A @ B^T.

gemm_nt_q5_0_sse Forward
void gemm_nt_q5_0_sse(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_0_sse_v2 Forward
void gemm_nt_q5_0_sse_v2(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_1 Forward
void gemm_nt_q5_1(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

GEMM with transposed Q5_1 weights: C = A @ B^T.

gemm_nt_q5_1_q8_1 Forward
void gemm_nt_q5_1_q8_1(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_1_q8_1_m4 Forward
void gemm_nt_q5_1_q8_1_m4(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_1_q8_1_m8 Forward
void gemm_nt_q5_1_q8_1_m8(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_1_q8_1_ref Forward
void gemm_nt_q5_1_q8_1_ref(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k Forward
void gemm_nt_q5_k(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_prepared Forward
void gemm_nt_q5_k_prepared(const float * A, const void * B_prepared, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_prepared_m4 Forward
void gemm_nt_q5_k_prepared_m4(const float * A, const void * B_prepared, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_prepared_q8_m4_nrange Forward
void gemm_nt_q5_k_prepared_q8_m4_nrange(const void * A_q8, const void * B_prepared, const float * bias, float * C, int M, int N, int K, int n_begin, int n_end)

gemm_nt_q5_k_q8_k Forward
void gemm_nt_q5_k_q8_k(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_q8_k_ref Forward
void gemm_nt_q5_k_q8_k_ref(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_ref Forward
void gemm_nt_q5_k_ref(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q5_k_ref_fp32 Forward
void gemm_nt_q5_k_ref_fp32(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k Forward
void gemm_nt_q6_k(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k_q8_k Forward
void gemm_nt_q6_k_q8_k(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

NT GEMM: C = A @ B^T where A is Q8_K and B is Q6_K.

gemm_nt_q6_k_q8_k_m4_tile Forward
void gemm_nt_q6_k_q8_k_m4_tile(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, int m0, int m1, int n0, int n1)

gemm_nt_q6_k_q8_k_parallel_dispatch Forward
void gemm_nt_q6_k_q8_k_parallel_dispatch(const void * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k_q8_k_prepared Forward
void gemm_nt_q6_k_q8_k_prepared(const void * A_q8, const void * B_prepared, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k_q8_k_prepared_avx512_vnni Forward
void gemm_nt_q6_k_q8_k_prepared_avx512_vnni(const void * A_q8, const void * B_prepared, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k_q8_k_prepared_tile Forward
void gemm_nt_q6_k_q8_k_prepared_tile(const void * A_q8, const void * B_prepared, const float * bias, float * C, int M, int N, int K, int m0, int m1, int n0, int n1)

gemm_nt_q6_k_q8_k_prepared_tile_impl Forward
void gemm_nt_q6_k_q8_k_prepared_tile_impl(const void * A_q8, const void * B_prepared, const float * bias, float * C, int M, int N, int K, int m0, int m1, int n0, int n1, int use_avx512_vnni)

gemm_nt_q6_k_q8_k_tile Forward
void gemm_nt_q6_k_q8_k_tile(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, int m0, int m1, int n0, int n1)

Compute one C[m0:m1, n0:n1] tile for Q8_K activations x Q6_K weights.

gemm_nt_q6_k_q8_k_tiled Forward
void gemm_nt_q6_k_q8_k_tiled(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

Experimental single-thread tiled NT GEMM wrapper.

gemm_nt_q6_k_ref Forward
void gemm_nt_q6_k_ref(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q6_k_sse Forward
void gemm_nt_q6_k_sse(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0 Forward
void gemm_nt_q8_0(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0_dispatch Forward
void gemm_nt_q8_0_dispatch(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, CKDataType dt)

gemm_nt_q8_0_mlp_dispatch Forward
void gemm_nt_q8_0_mlp_dispatch(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, CKDataType dt)

gemm_nt_q8_0_q8_0 Forward
void gemm_nt_q8_0_q8_0(const void * A_q8, const void * B_q8, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0_q8_0 with optional bias (matches header signature)

gemm_nt_q8_0_q8_0_contract Forward
void gemm_nt_q8_0_q8_0_contract(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0_q8_0_ggml_strict Forward
int gemm_nt_q8_0_q8_0_ggml_strict(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0_q8_0_m2n4 Forward
void gemm_nt_q8_0_q8_0_m2n4(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K)

gemm_nt_q8_0_q8_0_m2n4_tile Forward
void gemm_nt_q8_0_q8_0_m2n4_tile(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, int ldc)

gemm_nt_q8_0_q8_0_ref Forward
void gemm_nt_q8_0_q8_0_ref(const void * A, const void * B, float * C, int M, int N, int K)

Scalar reference: gemm_nt_q8_0_q8_0.

gemm_nt_q8_0_rowloop Forward
void gemm_nt_q8_0_rowloop(const float * A, const void * B, const float * bias, float * C, int M, int N, int K)

Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.

gemm_nt_q8_k_mlp_dispatch Forward
void gemm_nt_q8_k_mlp_dispatch(const void * A_q8, const void * B, const float * bias, float * C, int M, int N, int K, CKDataType dt)

gemm_nt_q8_k_qkv_dispatch Forward
void gemm_nt_q8_k_qkv_dispatch(const void * A_q8k, const void * B, const float * bias, float * C, int M, int N, int K, CKDataType dt)

gemm_q4_0 Forward
void gemm_q4_0(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q4_0 weights.

gemm_q4_0_backward Backward
void gemm_q4_0_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q4_1 Forward
void gemm_q4_1(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q4_1 weights.

gemm_q4_1_backward Backward
void gemm_q4_1_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q4_gateup_swiglu_thread_fn Forward
void gemm_q4_gateup_swiglu_thread_fn(int ith, int nth, void * args)

gemm_q4_gateup_swiglu_x16_thread_fn Forward
void gemm_q4_gateup_swiglu_x16_thread_fn(int ith, int nth, void * args)

gemm_q4_k Forward
void gemm_q4_k(float * Y, const void * W, const float * X, int M, int N, int K)

Auto-dispatch GEMM based on available SIMD.

gemm_q4_k_backward Backward
void gemm_q4_k_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q4_k_q8_k Forward
void gemm_q4_k_q8_k(float * Y, const void * W, const void * X_q8, int M, int N, int K)

gemm_q4_k_q8_k_compact_rows4 Forward
void gemm_q4_k_q8_k_compact_rows4(float * output, int output_stride, const void * weights, const void *const input_rows, int rows, int output_dim, int input_dim)

gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4 Forward
void gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4(float * output, const void * weights_packed, const void * input_q8, int rows, int output_dim, int input_dim)

gemm_q4_k_q8_k_ref Forward
void gemm_q4_k_q8_k_ref(float * Y, const void * W, const void * X_q8, int M, int N, int K)

gemm_q4_k_q8_k_thread_fn Forward
void gemm_q4_k_q8_k_thread_fn(int ith, int nth, void * args)

gemm_q4_k_ref Forward
void gemm_q4_k_ref(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q4_K weights (scalar reference)

gemm_q4_packed_meta_nsplit_thread_fn Forward
void gemm_q4_packed_meta_nsplit_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_thread_fn Forward
void gemm_q4_packed_meta_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x16_mreuse_process_job Forward
void gemm_q4_packed_meta_x16_mreuse_process_job(const gemm_q4_packed_meta_x16_thread_work_t * a, int job, int mt, int tile_m)

gemm_q4_packed_meta_x16_mreuse_thread_fn Forward
void gemm_q4_packed_meta_x16_mreuse_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x16_mtile_thread_fn Forward
void gemm_q4_packed_meta_x16_mtile_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_mreuse_thread_fn Forward
void gemm_q4_packed_meta_x8_mreuse_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_mtile_thread_fn Forward
void gemm_q4_packed_meta_x8_mtile_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_nsplit_thread_fn Forward
void gemm_q4_packed_meta_x8_nsplit_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_split_min_4m_thread_fn Forward
void gemm_q4_packed_meta_x8_split_min_4m_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_split_min_8m_thread_fn Forward
void gemm_q4_packed_meta_x8_split_min_8m_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn Forward
void gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_u8_x16_mtile_thread_fn Forward
void gemm_q4_packed_u8_x16_mtile_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_vnni_x16_q8k_16m_thread_fn Forward
void gemm_q4_packed_vnni_x16_q8k_16m_thread_fn(int ith, int nth, void * args)

gemm_q4_packed_vnni_x8_q8k_4m_job Forward
void gemm_q4_packed_vnni_x8_q8k_4m_job(gemm_q4_packed_vnni_x8_thread_work_t * a, int job, int row_tiles)

gemm_q4_packed_vnni_x8_q8k_4m_range_fn Forward
void gemm_q4_packed_vnni_x8_q8k_4m_range_fn(int begin, int end, void * args)

gemm_q4_packed_vnni_x8_q8k_4m_thread_fn Forward
void gemm_q4_packed_vnni_x8_q8k_4m_thread_fn(int ith, int nth, void * args)

gemm_q5_0 Forward
void gemm_q5_0(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q5_0 weights.

gemm_q5_0_backward Backward
void gemm_q5_0_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q5_1 Forward
void gemm_q5_1(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q5_1 weights.

gemm_q5_1_backward Backward
void gemm_q5_1_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q5_k_q8_k_compact_rows4 Forward
void gemm_q5_k_q8_k_compact_rows4(float * output, int output_stride, const void * weights, const void *const input_rows, int rows, int output_dim, int input_dim)

gemm_q6_k Forward
void gemm_q6_k(float * Y, const void * W, const float * X, int M, int N, int K)

gemm_q6_k_q8_k Forward
void gemm_q6_k_q8_k(float * Y, const void * W, const void * X_q8, int M, int N, int K)

GEMM: Y = W @ X^T where W is Q6_K and X is Q8_K.

gemm_q8_0 Forward
void gemm_q8_0(float * Y, const void * W, const float * X, int M, int N, int K)

Matrix-matrix multiply with Q8_0 weights.

gemm_q8_0_backward Backward
void gemm_q8_0_backward(float * dX, const void * W, const float * dY, int M, int N, int K)

Batched backward pass.

gemm_q8_0_q8_0_m2n4 Forward
void gemm_q8_0_q8_0_m2n4(float * C, const void * W, const void * A_q8, int M, int N, int K)

gemm_q8_0_q8_0_m2n4_strided Forward
void gemm_q8_0_q8_0_m2n4_strided(float * C, int ldc, const void * W, const void * A_q8, int M, int N, int K)

gemm_swiglu_fused Forward
void gemm_swiglu_fused(const float * x, const float * W_gate, const float * W_up, const float * b_gate, const float * b_up, float * output, int M, int N, int K)

gemm_tile_nt_strided Forward
void gemm_tile_nt_strided(const float * A, const float * B_tile, float * C, int tile_m, int tile_n, int K, int C_stride)

GEMM tile with N-dimension tiling (weight reuse)

gemm_tn_avx512 Forward
void gemm_tn_avx512(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_tn_bf16 Forward
void gemm_tn_bf16(const uint16_t * A, const uint16_t * B, const uint16_t * bias, uint16_t * C, int M, int N, int K)

gemm_tn_blocked Forward
void gemm_tn_blocked(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_tn_parallel Forward
void gemm_tn_parallel(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemm_tn_serial_double Forward
void gemm_tn_serial_double(const float * A, const float * B, const float * bias, float * C, int M, int N, int K)

gemma4_final_logit_softcap_forward Forward
void gemma4_final_logit_softcap_forward(float * logits, int tokens, int vocab_size, float cap)

Forward pass computation

gemma4_per_layer_embed_forward Forward
void gemma4_per_layer_embed_forward(float * hidden, const float * per_layer_input, const float * inp_gate, const float * proj, const float * post_norm, const float * out_scale, int tokens, int layer, int num_layers, int embed_dim, int per_layer_dim, float eps)

Forward pass computation

gemma4_per_layer_prepare_bf16_forward Forward
void gemma4_per_layer_prepare_bf16_forward(float * per_layer_input, const float * hidden, const int32_t * token_ids, const uint16_t * per_layer_token_emb, const uint16_t * per_layer_model_proj, const float * per_layer_proj_norm, int tokens, int num_layers, int embed_dim, int per_layer_dim, int vocab_size, float eps)

Forward pass computation

gemma4_per_layer_prepare_forward Forward
void gemma4_per_layer_prepare_forward(float * per_layer_input, const float * hidden, const int32_t * token_ids, const void * per_layer_token_emb, const uint16_t * per_layer_model_proj, const float * per_layer_proj_norm, int tokens, int num_layers, int embed_dim, int per_layer_dim, int vocab_size, float eps)

Forward pass computation

gemma4_v_norm_forward Forward
void gemma4_v_norm_forward(const float * input, float * output, float * rstd_cache, int tokens, int num_kv_heads, int head_dim, float eps)

Forward pass computation

gemma4_v_norm_forward_parallel_dispatch Forward
void gemma4_v_norm_forward_parallel_dispatch(const float * input, float * output, float * rstd_cache, int tokens, int num_kv_heads, int head_dim, float eps)

Forward pass computation

gemma4_vision_projector_prep_forward Forward
void gemma4_vision_projector_prep_forward(const float * input, float * output, int tokens, int dim, float scale, float eps)

Forward pass computation

get_gemm_params Forward
const GEMMParams * get_gemm_params(void)

position_embeddings_add_gemma4v_xy Forward
void position_embeddings_add_gemma4v_xy(float * x, const float * position_embd, int grid_h, int grid_w, int embed_dim, int source_grid_size)

rope_forward_gemma4v_vision_xy_one Forward
void rope_forward_gemma4v_vision_xy_one(float * x, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int grid_w, int rotary_dim, float freq_base)

Forward pass computation

rope_forward_qk_gemma4_direct Forward
void rope_forward_qk_gemma4_direct(float * q, float * k, const float * freq_factors, int use_freq_factors, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base)

Forward pass computation

rope_forward_qk_gemma4v_vision_xy Forward
void rope_forward_qk_gemma4v_vision_xy(float * q, float * k, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int grid_w, int rotary_dim, float freq_base)

Forward pass computation

Layer Normalization (13)
layernorm_backward_kernel Backward
void layernorm_backward_kernel(const float * d_output, const float * input, const float * gamma, const float * mean, const float * rstd, float * d_input, float * d_gamma, float * d_beta, int tokens, int d_model, int aligned_embed_dim)

Backward pass / gradient computation

layernorm_backward_kernel_bf16 Backward
void layernorm_backward_kernel_bf16(const uint16_t * d_output, const uint16_t * input, const float * gamma, const float * mean, const float * rstd, uint16_t * d_input, float * d_gamma, float * d_beta, int tokens, int d_model, int aligned_embed_dim, float * scratch_d_output, float * scratch_input, float * scratch_d_input)

Backward pass / gradient computation

layernorm_forward_ggml_exact Forward
void layernorm_forward_ggml_exact(const float * input, const float * gamma, const float * beta, float * output, float * mean_cache, float * rstd_cache, int tokens, int d_model, int input_stride, int output_stride, int aligned_embed_dim, float eps)

Forward pass computation

layernorm_forward_rolled_slice Forward
void layernorm_forward_rolled_slice(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

layernorm_forward_rolled_slice_bf16 Forward
void layernorm_forward_rolled_slice_bf16(const uint16_t *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, uint16_t *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, int aligned_embed_dim, float eps, float * scratch_input, float * scratch_output)

Forward pass computation

layernorm_forward_unrolled_slice Forward
void layernorm_forward_unrolled_slice(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, float eps)

Forward pass computation

layernorm_forward_unrolled_slice_bf16 Forward
void layernorm_forward_unrolled_slice_bf16(const uint16_t *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, uint16_t *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, float eps, float * scratch_input, float * scratch_output)

Forward pass computation

layernorm_forward_unrolled_slice_scalar Forward
void layernorm_forward_unrolled_slice_scalar(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, float eps)

Forward pass computation

layernorm_naive_serial Forward
void layernorm_naive_serial(const float * input, const float * gamma, const float * beta, float * output, float * mean_cache, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

layernorm_naive_serial_bf16_storage Forward
void layernorm_naive_serial_bf16_storage(const float * input, const float * gamma, const float * beta, float * output, float * mean_cache, float * rstd_cache, int tokens, int d_model, float eps)

layernorm_naive_serial_matched_precision Forward
void layernorm_naive_serial_matched_precision(const float * input, const float * gamma, const float * beta, float * output, float * mean_cache, float * rstd_cache, int tokens, int d_model, float eps)

layernorm_pytorch_welford_bf16_storage Forward
void layernorm_pytorch_welford_bf16_storage(const float * input, const float * gamma, const float * beta, float * output, float * mean_cache, float * rstd_cache, int tokens, int d_model, float eps)

zero_layernorm_padding Forward
void zero_layernorm_padding(float * out_ptr, int d_model, int aligned_embed_dim)

RMS Normalization (60)
ck_layer_backward_rmsnorm_swiglu Backward
void ck_layer_backward_rmsnorm_swiglu(const CKLayerBackwardParams * p)

Backward pass / gradient computation

ck_layer_forward_rmsnorm_swiglu Forward
void ck_layer_forward_rmsnorm_swiglu(const CKLayerForwardParams * p)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode Forward
void ck_layer_forward_rmsnorm_swiglu_decode(const CKLayerForwardParams * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_fused Forward
void ck_layer_forward_rmsnorm_swiglu_decode_fused(const CKLayerForwardParams * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_fused_attn Forward
void ck_layer_forward_rmsnorm_swiglu_decode_fused_attn(const CKLayerForwardParams * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_fused_attn_impl Forward
void ck_layer_forward_rmsnorm_swiglu_decode_fused_attn_impl(const CKLayerForwardParams * p, int token_index, int cache_capacity, int fuse_mlp)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_fused_attn_mlp Forward
void ck_layer_forward_rmsnorm_swiglu_decode_fused_attn_mlp(const CKLayerForwardParams * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_q4_k Forward
void ck_layer_forward_rmsnorm_swiglu_decode_q4_k(const CKLayerForwardParamsQ4K * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_decode_quant Forward
void ck_layer_forward_rmsnorm_swiglu_decode_quant(const CKLayerForwardParamsQ4K * p, int token_index, int cache_capacity)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_q4_k Forward
void ck_layer_forward_rmsnorm_swiglu_q4_k(const CKLayerForwardParamsQ4K * p)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_quant Forward
void ck_layer_forward_rmsnorm_swiglu_quant(const CKLayerForwardParamsQ4K * p)

Forward pass computation

ck_layer_forward_rmsnorm_swiglu_ref Forward
void ck_layer_forward_rmsnorm_swiglu_ref(const CKLayerForwardParams * p)

Forward pass computation

ck_test_rmsnorm Forward
void ck_test_rmsnorm(const float * input, const float * weight, float * output, int n_tokens, int dim, float eps)

RMSNorm.

fused_rmsnorm Forward
void fused_rmsnorm(const float * input, const float * gamma, const float * beta, float * output, int hidden, float eps)

Fused RMSNorm - writes to pre-allocated buffer.

fused_rmsnorm_linear_q4k Forward
void fused_rmsnorm_linear_q4k(float * y, const float * x, const float * gamma, const void * W_q4k, int M, int K, float eps)

Fused RMSNorm + Q4_K Linear projection.

fused_rmsnorm_qkv Forward
void fused_rmsnorm_qkv(const float * input, const float * gamma, const float * W_qkv, const float * b_qkv, float * q_out, float * k_out, float * v_out, int hidden, int num_heads, int num_kv_heads, int head_dim, float eps)

Fused RMSNorm with fused QKV projection.

fused_rmsnorm_qkv_prefill Forward
void fused_rmsnorm_qkv_prefill(const float * x, const float * gamma, const float * Wq, const float * Wk, const float * Wv, float * Q, float * K, float * V, int seq_len, int hidden, int q_dim, int kv_dim, float eps, float * scratch)

Fused RMSNorm + QKV projection for prefill.

fused_rmsnorm_qkv_prefill_head_major Forward
void fused_rmsnorm_qkv_prefill_head_major(const float * x, const float * gamma, const float * Wq, const float * Bq, const float * Wk, const float * Bk, const float * Wv, const float * Bv, float * Q, float * K, float * V, int seq_len, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int kv_stride_tokens, float eps, float * scratch)

Fused RMSNorm + QKV projection for prefill (head-major outputs)

fused_rmsnorm_qkv_prefill_head_major_quant Forward
void fused_rmsnorm_qkv_prefill_head_major_quant(const float * x, const float * gamma, const void * Wq, const float * Bq, CKDataType wq_dt, const void * Wk, const float * Bk, CKDataType wk_dt, const void * Wv, const float * Bv, CKDataType wv_dt, float * Q, float * K, float * V, int seq_len, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int kv_stride_tokens, float eps, void * scratch)

Fused RMSNorm + QKV projection for prefill (head-major, Q8 activations)

fused_rmsnorm_qkv_prefill_head_major_quant_scratch_size Forward
size_t fused_rmsnorm_qkv_prefill_head_major_quant_scratch_size(int aligned_embed_dim)

Get scratch buffer size for fused_rmsnorm_qkv_prefill_head_major_quant.

fused_rmsnorm_qkv_scratch_size Forward
size_t fused_rmsnorm_qkv_scratch_size(int hidden)

Get scratch buffer size for fused_rmsnorm_qkv_prefill.

mamba2_rmsnorm_gate_f32 Forward
void mamba2_rmsnorm_gate_f32(const float * x, const float * gate, const float * weight, float * out, int rows, int inner_dim, int group_size, float eps)

mega_fuse_rmsnorm_qkv Forward
void mega_fuse_rmsnorm_qkv(float * q_out, float * k_out, float * v_out, const float * input, const float * gamma, const float * W_qkv, const float * b_qkv, int hidden, int num_heads, int num_kv_heads, int head_dim, float eps)

Phase 1: Fused RMSNorm + QKV (intermediates in registers)

mega_fuse_rmsnorm_qkv_avx Forward
void mega_fuse_rmsnorm_qkv_avx(float * q_out, float * k_out, float * v_out, const float * input, const float * gamma, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, float eps)

Fused RMSNorm + QKV for decode (single token)

mega_fuse_rmsnorm_qkv_rope Forward
void mega_fuse_rmsnorm_qkv_rope(float * q_out, float * k_out, float * v_out, const float * input, const float * gamma, const float * W_qkv, const float * b_qkv, const float * rope_cos, const float * rope_sin, int pos, int hidden, int num_heads, int num_kv_heads, int head_dim, int max_seq, float eps)

Phase 2: Fused RMSNorm + QKV + RoPE.

qwen4_group_rmsnorm_llama Forward
void qwen4_group_rmsnorm_llama(const float * input, const float * weight, float * output, int groups, int hidden_dim, float eps)

qwen4_group_rmsnorm_pytorch_bf16 Forward
void qwen4_group_rmsnorm_pytorch_bf16(const float * input, const float * weight, float * output, int groups, int hidden_dim, float eps)

qwen4_shared_head_rmsnorm Forward
void qwen4_shared_head_rmsnorm(const float * input, const float * weight, float * output, int heads, int head_dim, float eps)

rmsnorm_backward Backward
void rmsnorm_backward(const float * d_output, const float * input, const float * gamma, const float * rstd_cache, float * d_input, float * d_gamma, int tokens, int d_model, int aligned_embed_dim)

RMSNorm backward pass Testtest_rmsnorm.py::TestRMSNormBackward::test_backward_tokens test_rmsnorm.py::TestRMSNormBackward::test_backward_single test_parity.py::test_rmsnorm_backward_parity

rmsnorm_backward_bf16 Backward
void rmsnorm_backward_bf16(const uint16_t * d_output, const uint16_t * input, const float * gamma, const float * rstd_cache, uint16_t * d_input, float * d_gamma, int tokens, int d_model, int aligned_embed_dim)

Backward pass / gradient computation

rmsnorm_backward_int4 Backward
void rmsnorm_backward_int4(const uint8_t * d_output, const uint8_t * input, const float * gamma, const float * rstd_cache, uint8_t * d_input, float * d_gamma, int tokens, int d_model, int aligned_embed_dim, float * scratch_d_output, float * scratch_input, float * scratch_d_input)

Backward pass / gradient computation

rmsnorm_backward_int8 Backward
void rmsnorm_backward_int8(const int8_t * d_output, const int8_t * input, const float * gamma, const float * rstd_cache, int8_t * d_input, float * d_gamma, int tokens, int d_model, int aligned_embed_dim, float * scratch_d_output, float * scratch_input, float * scratch_d_input)

Backward pass / gradient computation

rmsnorm_backward_strict_scalar Backward
void rmsnorm_backward_strict_scalar(const float * d_output, const float * input, const float * gamma, const float * rstd_cache, float * d_input, float * d_gamma, int tokens, int d_model, int aligned_embed_dim)

Backward pass / gradient computation

rmsnorm_forward Forward
void rmsnorm_forward(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_bf16 Forward
void rmsnorm_forward_bf16(const uint16_t * input, const float * gamma, uint16_t * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_fp32_square_fp64_sum Forward
void rmsnorm_forward_fp32_square_fp64_sum(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_fp64_sum Forward
void rmsnorm_forward_fp64_sum(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_int4 Forward
void rmsnorm_forward_int4(const uint8_t * input, const float * gamma, uint8_t * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps, float * scratch_input, float * scratch_output)

Forward pass computation

rmsnorm_forward_int8 Forward
void rmsnorm_forward_int8(const int8_t * input, const float * gamma, int8_t * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps, float * scratch_input, float * scratch_output)

Forward pass computation

rmsnorm_forward_kv_lora Forward
void rmsnorm_forward_kv_lora(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_llama_production Forward
void rmsnorm_forward_llama_production(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_no_weight Forward
void rmsnorm_forward_no_weight(const float * input, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_parallel_dispatch Forward
void rmsnorm_forward_parallel_dispatch(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_pytorch_bf16_storage Forward
void rmsnorm_forward_pytorch_bf16_storage(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_pytorch_bf16_storage_impl Forward
void rmsnorm_forward_pytorch_bf16_storage_impl(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int input_stride, int output_stride, float eps, int qwen3next_weight_order)

Forward pass computation

rmsnorm_forward_qwen3next_pytorch_bf16_storage Forward
void rmsnorm_forward_qwen3next_pytorch_bf16_storage(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)

Forward pass computation

rmsnorm_forward_strict_scalar Forward
void rmsnorm_forward_strict_scalar(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int input_stride, int output_stride, float eps)

Forward pass computation

rmsnorm_forward_strided_f32 Forward
void rmsnorm_forward_strided_f32(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int input_stride, int output_stride, float eps)

RMSNorm forward pass Testtest_rmsnorm.py::TestRMSNormForward::test_fp32_tokens test_rmsnorm.py::TestRMSNormForward::test_fp32_single test_rmsnorm.py::TestRMSNormForward::test_perf_rolled test_layernorm.py::TestLayerNormForward::test_rmsnorm_compat test_parity.py::test_rmsnorm_parity

rmsnorm_forward_strided_pytorch_bf16_storage Forward
void rmsnorm_forward_strided_pytorch_bf16_storage(const float * input, const float * gamma, float * output, float * rstd_cache, int tokens, int d_model, int input_stride, int output_stride, float eps)

Forward pass computation

rmsnorm_llama_production_rstd Forward
float rmsnorm_llama_production_rstd(float mean_eps)

rmsnorm_q8_k_fused Forward
void rmsnorm_q8_k_fused(const float * input, const float * gamma, void * vy, int tokens, int d_model, int aligned_embed_dim, float eps)

rmsnorm_qkv_fp32_fused Forward
void rmsnorm_qkv_fp32_fused(const float * x, const float * rms_weight, const float * wq, const float * wk, const float * wv, float * q_out, float * k_out, float * v_out, int embed_dim, int q_dim, int kv_dim, float eps)

rmsnorm_qkv_fp32_fused_v2 Forward
void rmsnorm_qkv_fp32_fused_v2(const float * x, const float * rms_weight, const float * wq, const float * wk, const float * wv, float * q_out, float * k_out, float * v_out, int embed_dim, int q_dim, int kv_dim, float eps)

rmsnorm_qkv_fp32_fused_v3 Forward
void rmsnorm_qkv_fp32_fused_v3(const float * x, const float * rms_weight, const float * wq, const float * wk, const float * wv, float * q_out, float * k_out, float * v_out, int embed_dim, int q_dim, int kv_dim, float eps)

rmsnorm_qkv_q4k_fused Forward
void rmsnorm_qkv_q4k_fused(const float * x, const float * rms_weight, const void * wq, const void * wk, const void * wv, float * q_out, float * k_out, float * v_out, int embed_dim, int q_dim, int kv_dim, float eps)

rmsnorm_qkv_separate_fp32 Forward
void rmsnorm_qkv_separate_fp32(const float * x, const float * rms_weight, const float * wq, const float * wk, const float * wv, float * normed, float * q_out, float * k_out, float * v_out, int embed_dim, int q_dim, int kv_dim, float eps)

rmsnorm_tile Forward
void rmsnorm_tile(const float * input, const float * gamma, float * output, int tile_m, int embed_dim, int aligned_embed_dim, float eps)

Compute RMSNorm for a tile of tokens.

simple_rmsnorm Forward
void simple_rmsnorm(const float * input, const float * gamma, float * output, int tokens, int d_model, float eps)

unfused_rmsnorm_linear_q4k_ref Forward
void unfused_rmsnorm_linear_q4k_ref(float * y, const float * x, const float * gamma, const void * W_q4k, int M, int K, float eps)

Reference (unfused) implementation for correctness testing.

unfused_rmsnorm_qkv_prefill Forward
void unfused_rmsnorm_qkv_prefill(const float * x, const float * gamma, const float * Wq, const float * Wk, const float * Wv, float * x_norm, float * Q, float * K, float * V, int seq_len, int hidden, int q_dim, int kv_dim, float eps)

Unfused version for benchmarking comparison.

GELU Activation (27)
ck_gelu_ggml_runtime_init Forward
void ck_gelu_ggml_runtime_init(void)

ck_gelu_ggml_table_init Forward
void ck_gelu_ggml_table_init(void)

ck_gelu_reference_math_init Forward
void ck_gelu_reference_math_init(void)

ck_gelu_system_erf Forward
ck_gelu_math_f64_fn ck_gelu_system_erf(void)

ck_gelu_system_tanhf Forward
ck_gelu_math_f32_fn ck_gelu_system_tanhf(void)

ck_gelu_tanh_f32 Forward
float ck_gelu_tanh_f32(float x)

ck_gelu_tanh_ggml_reference_f32 Forward
float ck_gelu_tanh_ggml_reference_f32(float x)

ck_gelu_tanh_parity_f32 Forward
float ck_gelu_tanh_parity_f32(float x)

ck_gelu_try_bind_runtime Forward
void ck_gelu_try_bind_runtime(void * handle)

fast_gelu_scalar Forward
float fast_gelu_scalar(float x)

gelu_backward_exact Backward
void gelu_backward_exact(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

gelu_backward_exact_bf16 Backward
void gelu_backward_exact_bf16(const uint16_t * input, const uint16_t * d_output, uint16_t * d_input, size_t n, float * scratch_input, float * scratch_d_output, float * scratch_d_input)

Backward pass / gradient computation

gelu_backward_fast Backward
void gelu_backward_fast(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

gelu_backward_fast_bf16 Backward
void gelu_backward_fast_bf16(const uint16_t * input, const uint16_t * d_output, uint16_t * d_input, size_t n, float * scratch_input, float * scratch_d_output, float * scratch_d_input)

Backward pass / gradient computation

gelu_backward_scalar Backward
void gelu_backward_scalar(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

gelu_derivative_scalar Forward
float gelu_derivative_scalar(float x)

gelu_erf_bf16_storage Forward
void gelu_erf_bf16_storage(float * data, size_t n)

gelu_erf_fp64_f32_inplace Forward
void gelu_erf_fp64_f32_inplace(float * data, size_t n)

gelu_exact_inplace Forward
void gelu_exact_inplace(float * data, size_t n)

gelu_fast_inplace Forward
void gelu_fast_inplace(float * data, size_t n)

GELU activation forward (fast approximation, in-place) Testtest_gelu.py::TestGELUForward::test_gelu_fast_inplace test_gelu.py::TestGELUForward::test_gelu_vs_exact test_parity.py::test_gelu_parity

gelu_fast_inplace_bf16 Forward
void gelu_fast_inplace_bf16(uint16_t * data, size_t n, float * scratch)

gelu_ggml_inplace Forward
void gelu_ggml_inplace(float * data, size_t n)

gelu_ggml_native_inplace Forward
void gelu_ggml_native_inplace(float * data, size_t n)

gelu_pytorch_erf_f32_inplace Forward
void gelu_pytorch_erf_f32_inplace(float * data, size_t n)

gelu_pytorch_erf_sleef_bf16_storage Forward
void gelu_pytorch_erf_sleef_bf16_storage(float * data, size_t n)

gelu_pytorch_tanh_bf16_storage Forward
void gelu_pytorch_tanh_bf16_storage(float * data, size_t n)

gelu_scalar Forward
float gelu_scalar(float x)

Softmax (22)
backward_causal_softmax_head_major Backward
void backward_causal_softmax_head_major(float * d_scores, const float * weights, int num_heads, int num_tokens, int aligned_context_window)

Backward pass / gradient computation

backward_causal_softmax_head_major_bf16 Backward
void backward_causal_softmax_head_major_bf16(uint16_t * d_scores, const uint16_t * weights, int num_heads, int num_tokens, int aligned_context_window, float * scratch_d_scores, float * scratch_weights)

Backward pass / gradient computation

causal_softmax_head_major Forward
void causal_softmax_head_major(float * scores, int num_heads, int num_tokens, int aligned_context_window)

Causal softmax (in-place, row-wise) Testtest_softmax.py::TestSoftmaxForward::test_causal_softmax test_softmax.py::TestSoftmaxForward::test_causal_vs_softmax test_attention.py::TestAttentionForward::test_softmax_correctness

causal_softmax_head_major_bf16 Forward
void causal_softmax_head_major_bf16(uint16_t * scores, int num_heads, int num_tokens, int aligned_context_window, float * scratch)

causal_softmax_head_major_exact Forward
void causal_softmax_head_major_exact(float * scores, int num_heads, int num_tokens, int aligned_context_window)

Causal softmax (exact version using stdlib expf) Testtest_softmax.py::TestSoftmaxForward::test_causal_softmax_exact test_softmax.py::TestSoftmaxForward::test_exact_vs_fast

ck_moe_llama_softmax_row Forward
double ck_moe_llama_softmax_row(float * probabilities, const float * logits, int n_experts, float max_value)

ck_test_softmax Forward
void ck_test_softmax(const float * input, float * output, int n)

Softmax (simple, non-causal)

deepseek_dsa_topk_softmax_backward_f32 Backward
void deepseek_dsa_topk_softmax_backward_f32(const int * indices, const float * weights, const float * d_weights, float * d_scores, int tokens, int heads, int key_count, int top_k)

Backward pass / gradient computation

deepseek_dsa_topk_softmax_f32 Forward
void deepseek_dsa_topk_softmax_f32(const float * scores, int * indices, float * weights, int tokens, int heads, int key_count, int top_k)

ds_softmax Forward
void ds_softmax(float * x, int n)

moe_softmax_topk_router_llama_f32_workspace Forward
int moe_softmax_topk_router_llama_f32_workspace(const float * logits, int * indices, float * weights, int rows, int n_experts, int top_k, float routed_scaling_factor, void * workspace, size_t workspace_bytes)

moe_softmax_topk_router_pytorch_bf16_workspace Forward
int moe_softmax_topk_router_pytorch_bf16_workspace(const float * logits, int * indices, float * weights, int rows, int n_experts, int top_k, float routed_scaling_factor, void * workspace, size_t workspace_bytes)

moe_softmax_topk_router_workspace_bytes Forward
size_t moe_softmax_topk_router_workspace_bytes(int n_experts)

softmax Forward
void softmax(float * x, int n)

softmax_cross_entropy_loss Forward
void softmax_cross_entropy_loss(const float * logits, const int32_t * targets, int tokens, int vocab_size, float * d_logits, float * loss_out)

softmax_cross_entropy_loss_bf16 Forward
void softmax_cross_entropy_loss_bf16(const uint16_t * logits, const int32_t * targets, int tokens, int vocab_size, uint16_t * d_logits, float * loss_out, float * scratch_logits, float * scratch_d_logits)

softmax_cross_entropy_loss_index_mean_impl Forward
void softmax_cross_entropy_loss_index_mean_impl(const float * logits, const int32_t * targets, int tokens, int vocab_size, float * d_logits, float * loss_out, int force_strict_math)

softmax_cross_entropy_loss_legacy_mean_impl Forward
void softmax_cross_entropy_loss_legacy_mean_impl(const float * logits, const int32_t * targets, int tokens, int vocab_size, float * d_logits, float * loss_out)

softmax_cross_entropy_loss_ptref Forward
void softmax_cross_entropy_loss_ptref(const float * logits, const int32_t * targets, int tokens, int vocab_size, float * d_logits, float * loss_out)

softmax_inplace Forward
void softmax_inplace(float * x, int n)

topk_softmax_backward_f32 Backward
void topk_softmax_backward_f32(const int * indices, const float * weights, const float * d_weights, float * d_scores, int num_tokens, int n_experts_or_keys, int k)

Backward for hard top-k followed by softmax over selected values.

topk_softmax_f32 Forward
void topk_softmax_f32(const float * scores, int n, int k, int * indices, float * weights)

Find top-K indices with softmax-normalized weights.

Attention (171)
attention_backward_causal_head_major Backward
void attention_backward_causal_head_major(const float * d_output, const float * q, const float * k, const float * v, const float * attn_weights, float * d_q, float * d_k, float * d_v, float * d_scores, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

Causal attention backward (non-GQA version) Testtest_attention_backward.py::TestAttentionBackward::test_backward test_attention_backward.py::TestAttentionBackward::test_backward_vs_separate test_parity.py::test_attention_backward_parity

attention_backward_causal_head_major_gqa Backward
void attention_backward_causal_head_major_gqa(const float * d_output, const float * q, const float * k, const float * v, const float * attn_weights, float * d_q, float * d_k, float * d_v, float * d_scores, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

GQA causal attention backward (score-matrix version) Testtest_attention_backward.py::TestAttentionBackwardGQA::test_gqa_backward test_attention_backward.py::TestAttentionBackwardGQA::test_gqa_vs_separate test_parity.py::test_attention_backward_parity

attention_backward_causal_head_major_gqa_bf16 Backward
void attention_backward_causal_head_major_gqa_bf16(const uint16_t * d_output, float * d_x, const uint16_t * q, const uint16_t * k, const uint16_t * v, const float * attn_weights, float * d_q, float * d_k, float * d_v, float * d_scores, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window, float * scratch_d_output, float * scratch_q, float * scratch_k, float * scratch_v)

BF16 attention backward with caller-provided scratch buffers Testbf16/test_attention_bf16.py::TestAttentionBF16::test_bf16_backward

attention_flash_cleanup Forward
void attention_flash_cleanup(void)

Clean up flash attention resources.

attention_flash_decode Forward
void attention_flash_decode(float * out, const float * q, const float * k, const float * v, int T_q, int T_k, int H, int D_h, float scale)

Main flash attention function with SIMD dispatch.

attention_flash_decode_scalar Forward
void attention_flash_decode_scalar(float * out, const float * q, const float * k, const float * v, int T_q, int T_k, int H, int D_h, float scale)

Scalar flash-style attention (online softmax)

attention_flash_init Forward
void attention_flash_init(int max_context, int max_heads, int max_head_dim)

Initialize flash attention buffers.

attention_flash_query_causal Forward
void attention_flash_query_causal(const float * q_vec, const float * k_head, const float * v_head, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec)

attention_flash_query_causal_exact Forward
void attention_flash_query_causal_exact(const float * q_vec, const float * k_head, const float * v_head, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec)

attention_flash_query_causal_exact_f16kv Forward
void attention_flash_query_causal_exact_f16kv(const float * q_vec, const float * k_head, const float * v_head, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec)

attention_flash_query_causal_exact_prerounded_f16kv Forward
void attention_flash_query_causal_exact_prerounded_f16kv(const float * q_vec, const float * k_head, const float * v_head, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec)

attention_flash_query_sliding Forward
void attention_flash_query_sliding(const float * q_vec, const float * k_head, const float * v_head, int query_pos, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec, int sliding_window)

attention_forward_causal_head_major Forward
void attention_forward_causal_head_major(const float * q, const float * k, const float * v, float * scores, float * output, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

Causal attention forward (score-matrix version) Testtest_attention.py::TestAttentionForward::test_causal_forward test_attention.py::TestAttentionForward::test_gqa_broadcast test_attention.py::TestAttentionForward::test_exact_vs_fast test_parity.py::test_attention_parity

attention_forward_causal_head_major_exact Forward
void attention_forward_causal_head_major_exact(const float * q, const float * k, const float * v, float * scores, float * output, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

Causal attention forward (exact version using stdlib expf) Testtest_attention.py::TestAttentionForward::test_exact_single test_attention.py::TestAttentionForward::test_exact_vs_fast

attention_forward_causal_head_major_gqa Forward
void attention_forward_causal_head_major_gqa(const float * q, const float * k, const float * v, float * scores, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

GQA causal attention forward (score-matrix version) Testtest_attention.py::TestAttentionForward::test_gqa_forward test_attention.py::TestAttentionForward::test_gqa_broadcast test_attention_backward.py::TestAttentionBackwardGQA::test_gqa_backward test_parity.py::test_attention_gqa_parity

attention_forward_causal_head_major_gqa_bf16 Forward
void attention_forward_causal_head_major_gqa_bf16(const uint16_t * q, const uint16_t * k, const uint16_t * v, float * scores, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window, float * scratch_q, float * scratch_k, float * scratch_v)

BF16 GQA causal attention forward Testbf16/test_attention_bf16.py::TestAttentionBF16::test_bf16_forward bf16/test_attention_bf16.py::TestAttentionBF16::test_bf16_gqa bf16/test_attention_bf16.py::TestAttentionBF16::test_bf16_flash

attention_forward_causal_head_major_gqa_exact Forward
void attention_forward_causal_head_major_gqa_exact(const float * q, const float * k, const float * v, float * scores, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int aligned_context_window)

GQA causal attention forward (exact version using stdlib expf) Testtest_attention.py::TestAttentionForward::test_gqa_exact bf16/test_attention_bf16.py::TestAttentionBF16::test_bf16_gqa

attention_forward_causal_head_major_gqa_flash Forward
void attention_forward_causal_head_major_gqa_flash(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided Forward
void attention_forward_causal_head_major_gqa_flash_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Flash attention forward with custom KV stride (for KV cache) Testtest_flash_attention.py::TestFlashAttention::test_flash_strided test_kv_cache_attention.py::TestKVCacheAttention::test_flash_attention

attention_forward_causal_head_major_gqa_flash_strided_f16kv Forward
void attention_forward_causal_head_major_gqa_flash_strided_f16kv(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_f16kv_serial Forward
void attention_forward_causal_head_major_gqa_flash_strided_f16kv_serial(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_f16kv_workspace Forward
void attention_forward_causal_head_major_gqa_flash_strided_f16kv_workspace(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, float * rounded_kv, size_t rounded_kv_bytes)

Forward pass computation

attention_forward_causal_head_major_gqa_flash_strided_sliding Forward
void attention_forward_causal_head_major_gqa_flash_strided_sliding(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window)

Flash attention forward with sliding window (prefill) Testtest_attention.py::TestAttentionForward::test_sliding_window_prefill

attention_forward_causal_head_major_gqa_flash_strided_token_output Forward
void attention_forward_causal_head_major_gqa_flash_strided_token_output(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_causal_head_major_gqa_llama_regular_strided_sliding_workspace Forward
void attention_forward_causal_head_major_gqa_llama_regular_strided_sliding_workspace(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window, float * scores, size_t scores_bytes, float * value_columns, size_t value_columns_bytes, float * scaled_scores, size_t scaled_scores_bytes)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_bf16cache_pytorch_contract Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_bf16cache_pytorch_contract(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_bf16cache_pytorch_contract_workspace Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_bf16cache_pytorch_contract_workspace(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction, float * token_workspace, size_t token_workspace_bytes)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_auto_workspace Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_auto_workspace(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction, float * token_workspace, size_t token_workspace_bytes, void * gqa_workspace, size_t gqa_workspace_bytes, int route_num_heads, int route_num_kv_heads, int route_head_dim, int route_query_tokens, int route_min_kv_tokens, int route_workers, int route_query_tile_size, int route_concurrent_query_tiles)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_contract Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_contract(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_contract_workspace Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_contract_workspace(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction, float * token_workspace, size_t token_workspace_bytes)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_gqa_reuse_config Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_gqa_reuse_config(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int query_tile_size, int concurrent_query_tiles, void * workspace, size_t workspace_bytes)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_gqa_reuse_workspace_bytes Forward
size_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_gqa_reuse_workspace_bytes(int num_heads, int num_kv_heads, int head_dim, int workers, int query_tile_size, int concurrent_query_tiles)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_append_f16cache_qtile64_schedule Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_append_f16cache_qtile64_schedule(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_prefill_schedule_t schedule)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_full_bf16cache_pytorch_contract Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_full_bf16cache_pytorch_contract(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction)

Forward pass computation

attention_forward_causal_head_major_gqa_prefill_segmented_f16cache_contract_workspace Forward
ck_attention_status_t attention_forward_causal_head_major_gqa_prefill_segmented_f16cache_contract_workspace(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction, float * token_workspace, size_t token_workspace_bytes, const int * segment_lengths, int num_segments)

Forward pass computation

attention_forward_decode_head_major_gqa_bf16cache_pytorch_contract Forward
ck_attention_status_t attention_forward_decode_head_major_gqa_bf16cache_pytorch_contract(const float * q_token, const uint16_t * k_cache, const uint16_t * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction)

Forward pass computation

attention_forward_decode_head_major_gqa_flash Forward
void attention_forward_decode_head_major_gqa_flash(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Flash attention decode (single token attends to KV cache) Testtest_flash_attention.py::TestFlashAttention::test_flash_decode test_kv_cache_attention.py::TestKVCacheAttention::test_flash_decode test_fused_attention_decode.py::TestFusedAttentionDecode::test_flash_decode test_attention.py::TestAttentionForward::test_flash_decode

attention_forward_decode_head_major_gqa_flash_f16cache Forward
void attention_forward_decode_head_major_gqa_flash_f16cache(const float * q_token, const uint16_t * k_cache, const uint16_t * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_f16cache_contract Forward
ck_attention_status_t attention_forward_decode_head_major_gqa_flash_f16cache_contract(const float * q_token, const uint16_t * k_cache, const uint16_t * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_f16cache_split Forward
void attention_forward_decode_head_major_gqa_flash_f16cache_split(const float * q_token, const uint16_t * k_cache, const uint16_t * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int split_chunks)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_f16cache_split_partitioned Forward
void attention_forward_decode_head_major_gqa_flash_f16cache_split_partitioned(const float * q_token, const uint16_t * k_cache, const uint16_t * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int split_chunks, int partition_tokens)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_f16kv Forward
void attention_forward_decode_head_major_gqa_flash_f16kv(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_decode_head_major_gqa_flash_sliding Forward
void attention_forward_decode_head_major_gqa_flash_sliding(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window)

Forward pass computation

attention_forward_decode_head_major_gqa_llama_regular_sliding_workspace Forward
void attention_forward_decode_head_major_gqa_llama_regular_sliding_workspace(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int live_tokens, int kv_stride_tokens, int head_dim, int aligned_head_dim, int sliding_window, float * scores, size_t scores_bytes, float * value_columns, size_t value_columns_bytes, float * scaled_scores, size_t scaled_scores_bytes)

Forward pass computation

attention_forward_decode_head_major_gqa_regular Forward
void attention_forward_decode_head_major_gqa_regular(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

WARNING: This is NOT true flash attention!

attention_forward_full_head_major_gqa_exact_strided Forward
void attention_forward_full_head_major_gqa_exact_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_flash Forward
void attention_forward_full_head_major_gqa_flash(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim)

Forward pass computation

attention_forward_full_head_major_gqa_flash_strided Forward
void attention_forward_full_head_major_gqa_flash_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_flash_strided_bf16_storage Forward
void attention_forward_full_head_major_gqa_flash_strided_bf16_storage(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_ggml_strided Forward
void attention_forward_full_head_major_gqa_ggml_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_ggml_strided_workspace Forward
void attention_forward_full_head_major_gqa_ggml_strided_workspace(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, float * score_rows, size_t score_rows_bytes, float * v_columns, size_t v_columns_bytes, float * probability_row, size_t probability_row_bytes)

Forward pass computation

attention_forward_full_head_major_gqa_pytorch_cpu_flash_bf16_storage Forward
void attention_forward_full_head_major_gqa_pytorch_cpu_flash_bf16_storage(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_pytorch_cpu_flash_bf16_storage_token_output Forward
void attention_forward_full_head_major_gqa_pytorch_cpu_flash_bf16_storage_token_output(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_sdpa_bf16_storage Forward
void attention_forward_full_head_major_gqa_sdpa_bf16_storage(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_tiled336_f16kv_fp32_strided Forward
void attention_forward_full_head_major_gqa_tiled336_f16kv_fp32_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_tiled64_f16kv_fp32_strided Forward
void attention_forward_full_head_major_gqa_tiled64_f16kv_fp32_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_full_head_major_gqa_tiled_f16kv_fp32_strided Forward
void attention_forward_full_head_major_gqa_tiled_f16kv_fp32_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

Forward pass computation

attention_forward_head_major_gqa_flash_impl Forward
void attention_forward_head_major_gqa_flash_impl(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int causal, int round_full_kv_fp16, int output_token_major, float scale)

Forward pass computation

attention_forward_head_major_gqa_unfused_f16_strict Forward
int attention_forward_head_major_gqa_unfused_f16_strict(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int causal, int output_token_major, float scale, int debug_layer_id)

Forward pass computation

attention_forward_query_key_head_major_f32 Forward
int attention_forward_query_key_head_major_f32(const float * query, const float * key, const float * value, float * output, float * score_scratch, int num_heads, int query_tokens, int key_tokens, int head_dim, float scale)

Forward pass computation

attention_forward_query_key_head_major_f32_packed_k Forward
int attention_forward_query_key_head_major_f32_packed_k(const float * query, const float * key, const float * value, float * output, float * score_scratch, float * key_transpose_scratch, int num_heads, int query_tokens, int key_tokens, int head_dim, float scale)

Forward pass computation

attention_forward_query_key_head_major_tiled_f16kv_fp32 Forward
int attention_forward_query_key_head_major_tiled_f16kv_fp32(const float * query, const float * key, const float * value, float * output, int num_heads, int query_tokens, int key_tokens, int head_dim, float scale)

Forward pass computation

attention_forward_sparse_token_major_gqa_bf16cache_pytorch_cpu_flash_contract Forward
void attention_forward_sparse_token_major_gqa_bf16cache_pytorch_cpu_flash_contract(const float * query, const uint16_t * key_cache, const uint16_t * value_cache, const float * selected_indices, float * output, float * score_scratch, int rows, int query_heads, int kv_heads, int head_dim, int selection_width, int context_length, int position)

Forward pass computation

attention_mlp_fused_fp32 Forward
void attention_mlp_fused_fp32(const float * q, const float * k_cache, const float * v_cache, int seq_len, int num_heads, int num_kv_heads, int head_dim, float attn_scale, const float * wo, const float * residual_1, const float * rms_weight, float eps, const float * w_gate, const float * w_up, const float * w_down, int embed_dim, int intermediate_dim, float * hidden_out)

attention_mlp_fused_q4k Forward
void attention_mlp_fused_q4k(const float * q, const float * k_cache, const float * v_cache, int seq_len, int num_heads, int num_kv_heads, int head_dim, float attn_scale, const void * wo, const float * residual_1, const float * rms_weight, float eps, const void * w_gate, const void * w_up, const void * w_down, int embed_dim, int intermediate_dim, float * hidden_out)

attention_mlp_separate_fp32 Forward
void attention_mlp_separate_fp32(const float * q, const float * k_cache, const float * v_cache, int seq_len, int num_heads, int num_kv_heads, int head_dim, float attn_scale, const float * wo, const float * residual_1, const float * rms_weight, float eps, const float * w_gate, const float * w_up, const float * w_down, int embed_dim, int intermediate_dim, float * attn_out_buf, float * hidden_after_attn_buf, float * normed_buf, float * gate_buf, float * up_buf, float * mlp_out_buf, float * hidden_out)

attention_output_index Forward
size_t attention_output_index(int h, int token, int num_heads, int num_tokens, int aligned_head_dim, int token_major)

attention_query_full_exact_regular Forward
void attention_query_full_exact_regular(const float * q_vec, const float * k_head, const float * v_cols, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * score_row, float * out_vec, int layer_id, int head_id, int query_id)

attention_query_full_ggml_regular Forward
void attention_query_full_ggml_regular(const float * q_vec, const float * k_head, const float * v_cols, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * score_row, float * out_vec, int layer_id, int head_id, int query_id)

ck_attention_align64_size Forward
size_t ck_attention_align64_size(size_t value)

ck_attention_bf16_pytorch_flash_work Forward
void ck_attention_bf16_pytorch_flash_work(int ith, int nth, void * opaque)

ck_attention_bf16_pytorch_gqa_available Forward
int ck_attention_bf16_pytorch_gqa_available(void)

ck_attention_bf16_sdpa_work Forward
void ck_attention_bf16_sdpa_work(int ith, int nth, void * opaque)

ck_attention_causal_f16kv_work Forward
void ck_attention_causal_f16kv_work(int ith, int nth, void * opaque)

ck_attention_dot_f16_llama Forward
float ck_attention_dot_f16_llama(const uint16_t * x, const uint16_t * y, int n)

ck_attention_dot_f16_unfused_llama Forward
float ck_attention_dot_f16_unfused_llama(const uint16_t * x, const uint16_t * y, int n)

ck_attention_f16_prefill_gqa_reuse_work Forward
void ck_attention_f16_prefill_gqa_reuse_work(int ith, int nth, void * opaque)

ck_attention_f16_prefill_gqa_reuse_worker_bytes Forward
size_t ck_attention_f16_prefill_gqa_reuse_worker_bytes(int num_heads, int num_kv_heads, int head_dim, int workers, int query_tile_size, int concurrent_query_tiles)

ck_attention_f16_prefill_qtile64_dispatch Forward
ck_attention_status_t ck_attention_f16_prefill_qtile64_dispatch(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int cache_is_bf16, size_t q_head_stride, size_t output_head_stride, ck_attention_prefill_schedule_t schedule)

ck_attention_f16_prefill_qtile64_work Forward
void ck_attention_f16_prefill_qtile64_work(int ith, int nth, void * opaque)

ck_attention_f16_reduce_expf Forward
float ck_attention_f16_reduce_expf(float value)

ck_attention_f16_split_work Forward
void ck_attention_f16_split_work(int ith, int nth, void * opaque)

ck_attention_flash_decode_wrapper Forward
void ck_attention_flash_decode_wrapper(const float * q_token, const float * k_cache, const float * v_cache, float * out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim)

Wrapper to call TRUE flash attention from orchestration layer.

ck_attention_flash_query_auto Forward
void ck_attention_flash_query_auto(const float * q_vec, const float * k_head, const float * v_head, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float * out_vec)

ck_attention_forward_causal_head_major_gqa_prefill_segmented_f16cache_schedule_workspace Forward
ck_attention_status_t ck_attention_forward_causal_head_major_gqa_prefill_segmented_f16cache_schedule_workspace(const float * q, const uint16_t * k_cache, const uint16_t * v_cache, float * output, int num_heads, int num_kv_heads, int q_tokens, int past_tokens, int cache_capacity, int head_dim, int aligned_head_dim, ck_attention_reduction_t reduction, float * token_workspace, size_t token_workspace_bytes, const int * segment_lengths, int num_segments, ck_attention_prefill_schedule_t schedule)

Forward pass computation

ck_attention_forward_full_head_major_gqa_tiled_f16kv_fp32_strided Forward
void ck_attention_forward_full_head_major_gqa_tiled_f16kv_fp32_strided(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int query_tile_size)

Forward pass computation

ck_attention_forward_query_key_head_major_f32_run Forward
int ck_attention_forward_query_key_head_major_f32_run(const float * query, const float * key, const float * value, float * output, float * score_scratch, float * key_transpose_scratch, int num_heads, int query_tokens, int key_tokens, int head_dim, float scale)

Forward pass computation

ck_attention_full_bf16_pytorch_flash Forward
int ck_attention_full_bf16_pytorch_flash(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int output_token_major)

ck_attention_full_bf16_sdpa_amx_range Forward
int ck_attention_full_bf16_sdpa_amx_range(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int head_begin, int head_step, int output_token_major)

ck_attention_full_bf16_sdpa_tiled Forward
int ck_attention_full_bf16_sdpa_tiled(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens)

ck_attention_full_bf16_sdpa_tiled_range Forward
int ck_attention_full_bf16_sdpa_tiled_range(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int head_begin, int head_step)

ck_attention_full_ggml_graph_oracle_multihead Forward
int ck_attention_full_ggml_graph_oracle_multihead(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, float scale)

ck_attention_full_grid_work Forward
void ck_attention_full_grid_work(int ith, int nth, void * opaque)

ck_attention_full_tiled_f16kv_fp32_range Forward
void ck_attention_full_tiled_f16kv_fp32_range(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int query_tile_size, int ith, int nth)

ck_attention_full_tiled_f16kv_fp32_work Forward
void ck_attention_full_tiled_f16kv_fp32_work(int ith, int nth, void * opaque)

ck_attention_ggml_out_graph_enabled Forward
int ck_attention_ggml_out_graph_enabled(void)

ck_attention_gqa_team_barrier_wait Forward
void ck_attention_gqa_team_barrier_wait(ck_attention_gqa_team_barrier_t * barrier)

ck_attention_head_full_ggml_graph_oracle_regular Forward
int ck_attention_head_full_ggml_graph_oracle_regular(const float * q_head, const float * k_head, const float * v_head, float * out_head, int num_tokens, int head_dim, int aligned_head_dim, float scale)

ck_attention_llama_regular_impl Forward
void ck_attention_llama_regular_impl(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int query_tokens, int live_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window, float * scores, size_t scores_bytes, float * value_columns, size_t value_columns_bytes, float * scaled_scores, size_t scaled_scores_bytes, int batched_prefill)

ck_attention_llama_regular_query Forward
void ck_attention_llama_regular_query(const float * query, const float * key_head, const float * value_columns, float * output, float * scores, float * scaled_scores, int live_tokens, int padded_tokens, int query_position, int head_dim, int aligned_head_dim, int sliding_window, int batched_prefill)

ck_attention_mad_f16_llama Forward
void ck_attention_mad_f16_llama(uint16_t * y, const uint16_t * x, float scale, int n)

ck_attention_matmul_f32_accum Forward
void ck_attention_matmul_f32_accum(float * c, const float * a, const float * b, int m, int k, int n)

ck_attention_oracle_dump_enabled Forward
int ck_attention_oracle_dump_enabled(void)

ck_attention_oracle_dump_layer_id Forward
int ck_attention_oracle_dump_layer_id(void)

ck_attention_oracle_dump_meta Forward
void ck_attention_oracle_dump_meta(const char * name, int layer_id, const struct ggml_tensor * t)

ck_attention_oracle_dump_tensor Forward
void ck_attention_oracle_dump_tensor(const char * name, int layer_id, const struct ggml_tensor * t)

ck_attention_oracle_exact_dump_layer Forward
int ck_attention_oracle_exact_dump_layer(int * layer_id_out)

ck_attention_oracle_meta_dump_enabled Forward
int ck_attention_oracle_meta_dump_enabled(void)

ck_attention_oracle_qkv_index Forward
size_t ck_attention_oracle_qkv_index(int h, int t, int d, int num_tokens, int aligned_head_dim)

ck_attention_oracle_should_dump_layer Forward
int ck_attention_oracle_should_dump_layer(int layer_id)

ck_attention_oracle_tensor_f32_at Forward
float ck_attention_oracle_tensor_f32_at(const struct ggml_tensor * t, size_t i0, size_t i1, size_t i2, size_t i3)

ck_attention_parallel_enabled Forward
int ck_attention_parallel_enabled(int total_queries, int num_tokens, int head_dim)

ck_attention_pick_active_threads Forward
int ck_attention_pick_active_threads(const ck_threadpool_t * pool, int total_queries, int num_tokens)

ck_attention_project_head_major Forward
void ck_attention_project_head_major(const float * attn_out, const float * wo, const float * bo, float * out, float * scratch, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_project_head_major_backward Backward
void ck_attention_project_head_major_backward(const float * d_out, const float * attn_out, const float * wo, float * d_attn_out, float * d_wo, float * d_bo, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

Backward pass / gradient computation

ck_attention_project_head_major_decode_token Forward
void ck_attention_project_head_major_decode_token(const float * attn_token, const float * wo, const float * bo, float * out_token, int embed_dim, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_project_head_major_decode_token_residual Forward
void ck_attention_project_head_major_decode_token_residual(const float * attn_token, const float * wo, const float * bo, const float * residual_in, float * proj_out, float * residual_out, int embed_dim, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_project_head_major_q4_k Forward
void ck_attention_project_head_major_q4_k(const float * attn_out, const void * wo, const float * bo, float * out, float * scratch, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_project_head_major_q4_k_q8_k Forward
void ck_attention_project_head_major_q4_k_q8_k(const float * attn_out, const void * wo, const float * bo, float * out, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_project_head_major_quant Forward
void ck_attention_project_head_major_quant(const float * attn_out, const void * wo, const float * bo, float * out, float * scratch, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim, CKDataType wo_dtype)

ck_attention_project_head_major_ref Forward
void ck_attention_project_head_major_ref(const float * attn_out, const float * wo, const float * bo, float * out, float * scratch, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

ck_attention_pytorch_sdpa_scale_f32 Forward
float ck_attention_pytorch_sdpa_scale_f32(int head_dim)

ck_attention_query_key_f32_transpose_work Forward
void ck_attention_query_key_f32_transpose_work(int ith, int nth, void * opaque)

ck_attention_query_key_f32_work Forward
void ck_attention_query_key_f32_work(int ith, int nth, void * opaque)

ck_attention_reference_expf Forward
float ck_attention_reference_expf(float value)

ck_attention_reverse_out_dot_enabled Forward
int ck_attention_reverse_out_dot_enabled(void)

ck_attention_scale_f16_llama Forward
void ck_attention_scale_f16_llama(uint16_t * y, float scale, int n)

ck_attention_sparse_bf16_pytorch_gqa_available Forward
int ck_attention_sparse_bf16_pytorch_gqa_available(void)

ck_attention_strict_scale_f32 Forward
float ck_attention_strict_scale_f32(int head_dim)

ck_attention_strict_unfused_f16_enabled Forward
int ck_attention_strict_unfused_f16_enabled(void)

ck_attention_trace Forward
void ck_attention_trace(const char * branch, int layer_id, int head_id)

ck_attention_trace_float Forward
void ck_attention_trace_float(const char * tag, int layer_id, int head_id, float value)

ck_attention_trace_query Forward
void ck_attention_trace_query(const char * tag, int layer_id, int head_id, int query_id, int value)

ck_attention_u16_cache_to_f32
float ck_attention_u16_cache_to_f32(uint16_t value, int cache_is_bf16)

ck_attention_vec_dump_enabled Forward
int ck_attention_vec_dump_enabled(void)

ck_attention_vec_dump_exact_query Forward
void ck_attention_vec_dump_exact_query(const float * q_vec, const float * k_head, const float * out_vec, int kv_tokens, int head_dim, int aligned_head_dim, float scale, int layer_id, int head_id, int query_id)

ck_attention_vec_dump_next_layer_id Forward
int ck_attention_vec_dump_next_layer_id(void)

ck_attention_vec_dump_parse_env_int Forward
int ck_attention_vec_dump_parse_env_int(const char * name, int * out)

ck_attention_vec_dump_selected_query Forward
void ck_attention_vec_dump_selected_query(const float * raw_scores, const float * probs, const float * out_vec, const float * v_cols, int kv_tokens, int head_dim, int layer_id, int head_id, int query_id)

ck_attention_vec_dump_should_emit Forward
int ck_attention_vec_dump_should_emit(int layer_id, int head_id, int query_id)

ck_attention_vec_dump_tensor Forward
void ck_attention_vec_dump_tensor(const char * name, int layer_id, int query_id, const float * data, size_t elem_count)

ck_attention_vec_dump_vcols_enabled Forward
int ck_attention_vec_dump_vcols_enabled(void)

ck_sliding_attention_compute_one Forward
void ck_sliding_attention_compute_one(const ck_sliding_attention_args_t * a, int job)

ck_sliding_attention_parallel_disabled Forward
int ck_sliding_attention_parallel_disabled(void)

ck_sliding_attention_pick_threads Forward
int ck_sliding_attention_pick_threads(ck_threadpool_t * pool, int total_jobs, int num_tokens, int head_dim)

ck_sliding_attention_work_fn Forward
void ck_sliding_attention_work_fn(int ith, int nth, void * args)

ck_test_attention_causal Forward
void ck_test_attention_causal(const float * q, const float * k, const float * v, float * out, int num_heads, int num_kv_heads, int tokens, int seq_len, int head_dim)

Multi-head causal attention for prefill (head-major layout)

deepseek_csa_attention_backward_f32 Backward
void deepseek_csa_attention_backward_f32(const float * d_out, const float * q, const float * k, const float * v, const int * indices, const float * attn, float * d_q, float * d_k, float * d_v, int query_tokens, int key_tokens, int heads, int dim, int top_k, float scale)

Backward pass / gradient computation

deepseek_csa_attention_f32 Forward
void deepseek_csa_attention_f32(const float * q, const float * k, const float * v, const int * indices, float * out, float * attn, int query_tokens, int key_tokens, int heads, int dim, int top_k, float scale)

deepseek_hybrid_attention_f32 Forward
void deepseek_hybrid_attention_f32(const float * q, const float * k, const float * v, const int * indices, float * out, float * attn, int query_tokens, int key_tokens, int heads, int dim, int top_k, float scale, int mode)

deepseek_mla_attention_decode_f32 Forward
void deepseek_mla_attention_decode_f32(const float * q, const float * k_cache, const float * v_cache, float * output, int num_heads, int num_kv_heads, int cache_len, int qk_head_dim, int v_head_dim, int max_seq_len, int cache_stride)

deepseek_mla_attention_decode_f32_workspace Forward
void deepseek_mla_attention_decode_f32_workspace(const float * q, const float * k_cache, const float * v_cache, float * output, int num_heads, int num_kv_heads, int cache_len, int qk_head_dim, int v_head_dim, int max_seq_len, int cache_stride, float scale, float * scores, size_t scores_bytes)

deepseek_mla_attention_f32 Forward
void deepseek_mla_attention_f32(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int qk_head_dim, int v_head_dim)

deepseek_mla_attention_f32_parallel_dispatch Forward
void deepseek_mla_attention_f32_parallel_dispatch(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int qk_head_dim, int v_head_dim, float scale, float * scores, size_t scores_bytes)

deepseek_mla_attention_f32_workspace Forward
void deepseek_mla_attention_f32_workspace(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int qk_head_dim, int v_head_dim, float scale, float * scores, size_t scores_bytes)

ds_mla_attention_f32_query_range Forward
void ds_mla_attention_f32_query_range(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int num_tokens, int qk_head_dim, int v_head_dim, float scale, float * scores, int query_begin, int query_end, int query_step)

ds_mla_attention_f32_work Forward
void ds_mla_attention_f32_work(int ith, int nth, void * opaque)

fused_flash_attention_all_heads Forward
void fused_flash_attention_all_heads(float * o_out, const float * q_all, const float * kv_cache_k, const float * kv_cache_v, int num_heads, int num_kv_heads, int head_dim, int seq_len, int kv_tile_size)

Fused Flash Attention for all heads (parallel dispatch)

fused_flash_attention_head Forward
void fused_flash_attention_head(float * o_out, const float * q, const float * kv_cache_k, const float * kv_cache_v, int kv_head_idx, int seq_len, int head_dim, int kv_tile_size)

Fused Flash Attention for single head.

mega_fuse_flash_attention_avx Forward
void mega_fuse_flash_attention_avx(float * o_out, const float * q, const float * kv_cache_k, const float * kv_cache_v, int num_heads, int num_kv_heads, int seq_len, int cache_capacity, int head_dim, int aligned_head_dim)

Flash attention with online softmax (AVX version)

mega_fused_attention Forward
void mega_fused_attention(float * output, const float * input, const float * residual, const float * W_qkv, const float * b_qkv, const float * W_o, const float * b_o, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int pos, int seq_len, int hidden, int num_heads, int num_kv_heads, int head_dim, int max_seq, float eps)

Complete mega-fused attention block.

mega_fused_attention_decode Forward
void mega_fused_attention_decode(float * output, const float * input, const float * residual, const float * ln1_gamma, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, const float * wo, const float * bo, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int pos, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int cache_capacity, float eps)

Mega-fused attention for decode mode (single token)

mega_fused_attention_decode_q5_0 Forward
void mega_fused_attention_decode_q5_0(float * output, const float * input, const float * residual, const void * wq_q5_0, const void * wk_q5_0, const void * wv_q8_0, const void * wo_q5_0, const float * ln_gamma, const float * bq, const float * bk, const float * bv, const float * bo, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int pos, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int cache_capacity, float eps, void * scratch)

Serial mega-fused attention decode kernel.

mega_fused_attention_decode_q5_0_parallel_simd Forward
void mega_fused_attention_decode_q5_0_parallel_simd(float * output, const float * input, const float * residual, const void * wq_q5_0, const void * wk_q5_0, const void * wv_q8_0, const void * wo_q5_0, const float * ln_gamma, const float * bq, const float * bk, const float * bv, const float * bo, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int pos, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int cache_capacity, float eps, void * scratch, int ith, int nth)

Parallel SIMD mega-fused attention decode kernel (threadpool-aware)

mega_fused_attention_decode_scratch_size Forward
int mega_fused_attention_decode_scratch_size(int AE, int H, int KV, int AD)

Calculate scratch buffer size needed for the kernel.

mega_fused_attention_decode_workspace Forward
void mega_fused_attention_decode_workspace(float * output, const float * input, const float * residual, const float * ln1_gamma, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, const float * wo, const float * bo, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int pos, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, int cache_capacity, float eps, float * q_output_workspace, size_t q_output_workspace_bytes, float * kv_workspace, size_t kv_workspace_bytes)

Full mega-fused attention for decode.

mega_fused_attention_prefill Forward
void mega_fused_attention_prefill(float * output, const float * input, const float * residual, const float * ln1_gamma, const void * wq, const float * bq, CKDataType wq_dt, const void * wk, const float * bk, CKDataType wk_dt, const void * wv, const float * bv, CKDataType wv_dt, const void * wo, const float * bo, CKDataType wo_dt, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int start_pos, int tokens, int cache_capacity, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, float eps, void * scratch)

Mega-fused attention for prefill mode (multiple tokens)

mega_fused_attention_prefill_q8_0 Forward
void mega_fused_attention_prefill_q8_0(float * output, const float * input, const float * residual, const float * ln1_gamma, const void * wq, const float * bq, CKDataType wq_dt, const void * wk, const float * bk, CKDataType wk_dt, const void * wv, const float * bv, CKDataType wv_dt, const void * wo, const float * bo, CKDataType wo_dt, float * kv_cache_k, float * kv_cache_v, const float * rope_cos, const float * rope_sin, int start_pos, int tokens, int cache_capacity, int embed_dim, int aligned_embed_dim, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim, float eps, void * scratch)

Mega-fused prefill attention kernel (Q8_0 out-proj)

mega_fused_attention_prefill_q8_0_scratch_size Forward
size_t mega_fused_attention_prefill_q8_0_scratch_size(int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

Get scratch buffer size for mega_fused_attention_prefill_q8_0.

mega_fused_attention_prefill_scratch_size Forward
size_t mega_fused_attention_prefill_scratch_size(int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

Get scratch buffer size for mega_fused_attention_prefill.

simple_attention Forward
void simple_attention(const float * q, const float * k, const float * v, float * output, int num_heads, int num_kv_heads, int seq_len, int head_dim)

MLP / Feed-Forward (31)
ck_mlp_swiglu_forward Forward
void ck_mlp_swiglu_forward(const float * input, const float * w1, const float * b1, const float * w2, const float * b2, float * fc1_out, float * swiglu_out, float * output, int tokens, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_fully_fused_token Forward
void ck_mlp_swiglu_forward_fully_fused_token(const float * input_row, const float * w1, const float * b1, const float * w2, const float * b2, float * output_row, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_fused_token Forward
void ck_mlp_swiglu_forward_fused_token(const float * input_row, const float * w1, const float * b1, const float * w2, const float * b2, float * swiglu_row, float * output_row, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_q4_k Forward
void ck_mlp_swiglu_forward_q4_k(const float * input, const void * w1, const float * b1, const void * w2, const float * b2, float * fc1_out, float * swiglu_out, float * output, int tokens, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_q4_k_q8_k Forward
void ck_mlp_swiglu_forward_q4_k_q8_k(const float * input, const void * w1, const float * b1, const void * w2, const float * b2, float * fc1_out, float * swiglu_out, float * output, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_q4_k_q8_k_prefill Forward
void ck_mlp_swiglu_forward_q4_k_q8_k_prefill(const float * input, const void * w1, const float * b1, const void * w2, const float * b2, float * fc1_out, float * swiglu_out, float * output, int tokens, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_quant Forward
void ck_mlp_swiglu_forward_quant(const float * input, const void * w1, const float * b1, CKDataType w1_dtype, const void * w2, const float * b2, CKDataType w2_dtype, float * fc1_out, float * swiglu_out, float * output, int tokens, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_mlp_swiglu_forward_ref Forward
void ck_mlp_swiglu_forward_ref(const float * input, const float * w1, const float * b1, const float * w2, const float * b2, float * fc1_out, float * swiglu_out, float * output, int tokens, int aligned_embed_dim, int aligned_intermediate_dim)

Forward pass computation

ck_test_outproj_mlp_fused_q5_0 Forward
void ck_test_outproj_mlp_fused_q5_0(const float * attn_out, const float * residual, const float * ln2_gamma, const void * wo, const void * w1, const void * w2, float * output, int tokens, int num_heads, int head_dim, int embed_dim, int intermediate, float eps, int w2_is_q6k)

Test mega-fused OutProj + MLP kernel (Q5_0 weights)

fused_mlp_swiglu_decode Forward
void fused_mlp_swiglu_decode(const float * x, const float * W_gate, const float * W_up, const float * W_down, const float * b_gate, const float * b_up, const float * b_down, float * output, int D, int Hff)

fused_mlp_swiglu_decode_tiled Forward
void fused_mlp_swiglu_decode_tiled(const float * x, const float * W_gate, const float * W_up, const float * W_down, const float * b_gate, const float * b_up, const float * b_down, float * output, int D, int Hff)

fused_mlp_swiglu_decode_v2 Forward
void fused_mlp_swiglu_decode_v2(const float * x, const float * W_gate, const float * W_up, const float * W_down, const float * b_gate, const float * b_up, const float * b_down, float * output, int D, int Hff)

fused_mlp_swiglu_prefill Forward
void fused_mlp_swiglu_prefill(const float * x, const float * W_gate, const float * W_up, const float * W_down, float * output, int seq_len, int hidden, int intermediate, float * scratch)

Fused MLP (Gate + Up + SwiGLU + Down) for prefill.

fused_mlp_swiglu_prefill_bias Forward
void fused_mlp_swiglu_prefill_bias(const float * x, const float * W_gate, const float * W_up, const float * W_down, const float * B_gate, const float * B_up, const float * B_down, float * output, int seq_len, int hidden, int intermediate, float * scratch)

Fused MLP (Gate + Up + SwiGLU + Down) for prefill with biases.

fused_mlp_swiglu_prefill_w1w2_quant Forward
void fused_mlp_swiglu_prefill_w1w2_quant(const float * x, const void * W1, const float * B1, CKDataType w1_dt, const void * W2, const float * B2, CKDataType w2_dt, float * output, int seq_len, int embed_dim, int aligned_embed_dim, int intermediate_dim, int aligned_intermediate_dim, void * scratch)

Quantized fused MLP for prefill (W1=gate+up, W2=down)

fused_mlp_swiglu_prefill_w1w2_quant_scratch_size Forward
size_t fused_mlp_swiglu_prefill_w1w2_quant_scratch_size(int aligned_embed_dim, int aligned_intermediate_dim)

Get scratch buffer size for fused_mlp_swiglu_prefill_w1w2_quant.

fused_mlp_swiglu_scratch_size Forward
size_t fused_mlp_swiglu_scratch_size(int intermediate)

Get scratch buffer size for fused_mlp_swiglu_prefill.

layer_fused_attn_mlp_qkv_q4k Forward
void layer_fused_attn_mlp_qkv_q4k(const float * q, const float * k_cache, const float * v_cache, int seq_len, float attn_scale, const void * wo, const float * rms_weight_mlp, const void * w_gate, const void * w_up, const void * w_down, const float * rms_weight_attn, const void * wq_next, const void * wk_next, const void * wv_next, const float * residual_in, int embed_dim, int intermediate_dim, int num_heads, int num_kv_heads, int head_dim, float eps, float * q_next, float * k_next, float * v_next, float * hidden_out)

mega_fused_outproj_mlp_prefill Forward
void mega_fused_outproj_mlp_prefill(float * output, const float * attn_out, const float * residual, const float * ln2_gamma, const void * wo, const float * bo, int wo_dt, const void * w1, const float * b1, int w1_dt, const void * w2, const float * b2, int w2_dt, int tokens, int embed_dim, int aligned_embed_dim, int num_heads, int aligned_head_dim, int intermediate_dim, int aligned_intermediate_dim, float eps, void * scratch)

mega_fused_outproj_mlp_prefill_scratch_size Forward
size_t mega_fused_outproj_mlp_prefill_scratch_size(int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim, int aligned_intermediate_dim)

Get scratch buffer size for mega_fused_outproj_mlp_prefill.

mlp_fused_fp32_v2 Forward
void mlp_fused_fp32_v2(const float * hidden_in, const float * rms_weight, float eps, const float * w_gate, const float * w_up, const float * w_down, int embed_dim, int intermediate_dim, float * hidden_out)

mlp_fused_fp32_v3 Forward
void mlp_fused_fp32_v3(const float * hidden_in, const float * rms_weight, float eps, const float * w_gate, const float * w_up, const float * w_down, int embed_dim, int intermediate_dim, float * hidden_out)

mlp_parallel Forward
void mlp_parallel(const void * ln2_q8, const void * W_gate, const void * W_up, const void * W_down, float * gate_buf, float * up_buf, float * swiglu_buf, void * down_q8, float * mlp_out, int intermediate, int embed_dim, int num_threads)

Parallel MLP (gate/up + SwiGLU + down projection).

mlp_q8_0_dtype_supported Forward
int mlp_q8_0_dtype_supported(CKDataType dt)

mlp_q8_k_dtype_supported Forward
int mlp_q8_k_dtype_supported(CKDataType dt)

mlp_separate_fp32 Forward
void mlp_separate_fp32(const float * hidden_in, const float * rms_weight, float eps, const float * w_gate, const float * w_up, const float * w_down, float * normed_buf, float * gate_buf, float * up_buf, int embed_dim, int intermediate_dim, float * hidden_out)

mlp_token_parallel Forward
void mlp_token_parallel(const float * input, const float * W_fc1, const float * b_fc1, const float * W_fc2, const float * b_fc2, float * fc1_output, float * output, int T, int aligned_dim, int num_threads)

mlp_token_parallel_bf16 Forward
void mlp_token_parallel_bf16(const uint16_t * input, const uint16_t * W_fc1, const uint16_t * b_fc1, const uint16_t * W_fc2, const uint16_t * b_fc2, float * fc1_output, float * output, int T, int aligned_dim, int num_threads, float * scratch_bias1_f, float * scratch_bias2_f, uint16_t * scratch_fc1_bf16)

Optimized MLP Forward (BF16 weights, FP32 activations)

mlp_token_parallel_bf16_backward_mixed Backward
void mlp_token_parallel_bf16_backward_mixed(const uint16_t * input, const uint16_t * W_fc1, const uint16_t * b_fc1, const uint16_t * W_fc2, const uint16_t * d_output, float * d_input, float * d_W_fc1, float * d_b_fc1, float * d_W_fc2, float * d_b_fc2, int T, int aligned_dim, int num_threads, float * scratch_fc1_pre, uint16_t * scratch_fc1_act_bf16, float * scratch_d_fc1)

BF16 MLP backward with FP32 gradient accumulation.

mlp_token_parallel_bf16_fp32act Forward
void mlp_token_parallel_bf16_fp32act(const uint16_t * input, const uint16_t * W_fc1, const uint16_t * b_fc1, const uint16_t * W_fc2, const uint16_t * b_fc2, float * fc1_output, float * output, int T, int aligned_dim, int num_threads, float * scratch_input_f, float * scratch_bias1_f, float * scratch_bias2_f, uint16_t * scratch_fc1_bf16)

Alternative: Fully FP32 activations throughout

mlp_token_parallel_exact Forward
void mlp_token_parallel_exact(const float * input, const float * W_fc1, const float * b_fc1, const float * W_fc2, const float * b_fc2, float * fc1_output, float * output, int T, int aligned_dim, int num_threads)

Sigmoid Activation (23)
attn_gate_sigmoid_mul_backward Backward
void attn_gate_sigmoid_mul_backward(const float * d_out, const float * x, const float * gate, float * d_x, float * d_gate, int rows, int num_heads, int state_dim)

Backward pass / gradient computation

attn_gate_sigmoid_mul_forward Forward
void attn_gate_sigmoid_mul_forward(const float * x, const float * gate, float * out, int rows, int num_heads, int state_dim)

Forward pass computation

attn_gate_sigmoid_mul_pytorch_bf16_storage Forward
void attn_gate_sigmoid_mul_pytorch_bf16_storage(const float * x, const float * gate, float * out, int rows, int num_heads, int state_dim)

ck_deltanet_llama_sigmoidf Forward
float ck_deltanet_llama_sigmoidf(float x)

ck_deltanet_sigmoidf Forward
float ck_deltanet_sigmoidf(float x)

ck_moe_sigmoid_f32 Forward
float ck_moe_sigmoid_f32(float x)

ck_sigmoid_bf16 Forward
float ck_sigmoid_bf16(float value)

ck_test_attn_gate_sigmoid_mul Forward
void ck_test_attn_gate_sigmoid_mul(const float * x, const float * gate, float * out, int rows, int dim)

Multiply attention output rows by sigmoid(gate) elementwise.

group_limited_topk_router_sigmoid_f32 Forward
void group_limited_topk_router_sigmoid_f32(const float * logits, const float * correction_bias, int * indices, float * weights, int rows, int n_experts, int top_k, int n_group, int topk_group, int norm_topk_prob, float routed_scaling_factor)

hybrid_sigmoid Forward
float hybrid_sigmoid(float x)

mamba2_sigmoid_f32 Forward
float mamba2_sigmoid_f32(float x)

recurrent_norm_sigmoid_gate_llama_avx2_forward Forward
void recurrent_norm_sigmoid_gate_llama_avx2_forward(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

Forward pass computation

recurrent_norm_sigmoid_gate_pytorch_bf16_storage Forward
void recurrent_norm_sigmoid_gate_pytorch_bf16_storage(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

recurrent_sigmoid Forward
float recurrent_sigmoid(float x)

recurrent_sigmoid_forward_ggml Forward
void recurrent_sigmoid_forward_ggml(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_sigmoid_forward_pytorch_bf16_input_fp32_output Forward
void recurrent_sigmoid_forward_pytorch_bf16_input_fp32_output(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_sigmoid_local Forward
float recurrent_sigmoid_local(float x)

sigmoid_backward Backward
void sigmoid_backward(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

sigmoid_backward_bf16 Backward
void sigmoid_backward_bf16(const uint16_t * input, const uint16_t * d_output, uint16_t * d_input, size_t n, float * scratch_input, float * scratch_d_output, float * scratch_d_input)

Backward pass / gradient computation

sigmoid_forward Forward
void sigmoid_forward(const float * input, float * output, size_t n)

Forward pass computation

sigmoid_forward_bf16 Forward
void sigmoid_forward_bf16(const uint16_t * input, uint16_t * output, size_t n, float * scratch_input, float * scratch_output)

Forward pass computation

sigmoid_scalar Forward
float sigmoid_scalar(float x)

sigmoid_scalar_parity Forward
float sigmoid_scalar_parity(float x)

SwiGLU Activation (62)
ck_moe_swiglu_expert_forward_q4k_q5k_bucketed_impl Forward
int ck_moe_swiglu_expert_forward_q4k_q5k_bucketed_impl(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, const void * expert_gate_packed, const void * expert_up_packed, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

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

ck_test_swiglu Forward
void ck_test_swiglu(const float * gate_up, float * output, int n_tokens, int intermediate_dim)

SwiGLU activation.

farskip_swiglu_shared_combine_bf16 Forward
void farskip_swiglu_shared_combine_bf16(const float * hidden, const float * routed, const float * post_attn_residual, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * main_output, float * routed_free_output, int rows, int hidden_dim, int intermediate_dim)

farskip_swiglu_shared_combine_bf16_parallel_dispatch Forward
void farskip_swiglu_shared_combine_bf16_parallel_dispatch(const float * hidden, const float * routed, const float * post_attn_residual, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * main_output, float * routed_free_output, int rows, int hidden_dim, int intermediate_dim)

farskip_swiglu_shared_combine_bf16_row_range Forward
void farskip_swiglu_shared_combine_bf16_row_range(const float * hidden, const float * routed, const float * post_attn_residual, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * main_output, float * routed_free_output, int rows, int hidden_dim, int intermediate_dim, int row_begin, int row_end)

moe_swiglu_expert_backward_f32 Backward
void moe_swiglu_expert_backward_f32(const float * d_output, const float * hidden, const int * indices, const float * routing_weights, const float * expert_gate, const float * expert_up, const float * expert_down, float * d_hidden, float * d_routing_weights, float * d_expert_gate, float * d_expert_up, float * d_expert_down, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Backward pass / gradient computation

moe_swiglu_expert_forward_bf16 Forward
void moe_swiglu_expert_forward_bf16(const float * hidden, const int * indices, const float * routing_weights, const uint16_t * expert_gate, const uint16_t * expert_up, const uint16_t * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_swiglu_expert_forward_bf16_parallel_dispatch Forward
void moe_swiglu_expert_forward_bf16_parallel_dispatch(const float * hidden, const int * indices, const float * routing_weights, const uint16_t * expert_gate, const uint16_t * expert_up, const uint16_t * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_swiglu_expert_forward_bf16_row_range Forward
void moe_swiglu_expert_forward_bf16_row_range(const float * hidden, const int * indices, const float * routing_weights, const uint16_t * expert_gate, const uint16_t * expert_up, const uint16_t * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, int row_begin, int row_end)

Forward pass computation

moe_swiglu_expert_forward_f32 Forward
void moe_swiglu_expert_forward_f32(const float * hidden, const int * indices, const float * routing_weights, const float * expert_gate, const float * expert_up, const float * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

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

Forward pass computation

moe_swiglu_expert_forward_q4k_q4k_parallel_workspace Forward
int moe_swiglu_expert_forward_q4k_q4k_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q4k_workspace Forward
int moe_swiglu_expert_forward_q4k_q4k_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5_0_parallel_workspace Forward
int moe_swiglu_expert_forward_q4k_q5_0_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5_0_workspace Forward
int moe_swiglu_expert_forward_q4k_q5_0_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_auto_prepared_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_auto_prepared_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_auto_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_auto_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_bucketed_prepared_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_bucketed_prepared_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, const void * expert_gate_packed, const void * expert_up_packed, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_bucketed_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_bucketed_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_parallel_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q5k_workspace Forward
int moe_swiglu_expert_forward_q4k_q5k_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q6k_parallel_workspace Forward
int moe_swiglu_expert_forward_q4k_q6k_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q6k_workspace Forward
int moe_swiglu_expert_forward_q4k_q6k_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q8_0_parallel_workspace Forward
int moe_swiglu_expert_forward_q4k_q8_0_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_forward_q4k_q8_0_workspace Forward
int moe_swiglu_expert_forward_q4k_q8_0_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_expert_q4k_q5k_bucketed_workspace_bytes Forward
size_t moe_swiglu_expert_q4k_q5k_bucketed_workspace_bytes(int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

moe_swiglu_expert_q4k_q5k_workspace_bytes Forward
size_t moe_swiglu_expert_q4k_q5k_workspace_bytes(int hidden_dim, int intermediate_dim)

moe_swiglu_expert_q4k_q8_0_workspace_bytes Forward
size_t moe_swiglu_expert_q4k_q8_0_workspace_bytes(int hidden_dim, int intermediate_dim)

moe_swiglu_nvfp4_workspace_bytes Forward
size_t moe_swiglu_nvfp4_workspace_bytes(int hidden_dim, int intermediate_dim)

moe_swiglu_packed_expert_forward_bf16 Forward
void moe_swiglu_packed_expert_forward_bf16(const float * hidden, const int * indices, const float * routing_weights, const uint16_t * expert_gate_up, const uint16_t * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_swiglu_shared_backward_f32 Backward
void moe_swiglu_shared_backward_f32(const float * d_output, const float * hidden, const float * shared_gate, const float * shared_up, const float * shared_down, float * d_hidden, float * d_routed, float * d_shared_gate, float * d_shared_up, float * d_shared_down, int rows, int hidden_dim, int intermediate_dim)

Backward pass / gradient computation

moe_swiglu_shared_forward_bf16 Forward
void moe_swiglu_shared_forward_bf16(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

moe_swiglu_shared_forward_bf16_gated Forward
void moe_swiglu_shared_forward_bf16_gated(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, const uint16_t * shared_router, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

moe_swiglu_shared_forward_bf16_gated_parallel_dispatch Forward
void moe_swiglu_shared_forward_bf16_gated_parallel_dispatch(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, const uint16_t * shared_router, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

moe_swiglu_shared_forward_bf16_gated_row_range Forward
void moe_swiglu_shared_forward_bf16_gated_row_range(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, const uint16_t * shared_router, float * output, int rows, int hidden_dim, int intermediate_dim, int row_begin, int row_end)

Forward pass computation

moe_swiglu_shared_forward_bf16_parallel_dispatch Forward
void moe_swiglu_shared_forward_bf16_parallel_dispatch(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

moe_swiglu_shared_forward_bf16_row_range Forward
void moe_swiglu_shared_forward_bf16_row_range(const float * hidden, const float * routed, const uint16_t * shared_gate, const uint16_t * shared_up, const uint16_t * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, int row_begin, int row_end)

Forward pass computation

moe_swiglu_shared_forward_f32 Forward
void moe_swiglu_shared_forward_f32(const float * hidden, const float * routed, const float * shared_gate, const float * shared_up, const float * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

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

Forward pass computation

moe_swiglu_shared_forward_q4k_q4k_parallel_workspace Forward
int moe_swiglu_shared_forward_q4k_q4k_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q4k_workspace Forward
int moe_swiglu_shared_forward_q4k_q4k_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q5_0_gated_parallel_workspace Forward
int moe_swiglu_shared_forward_q4k_q5_0_gated_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q5_0_gated_workspace Forward
int moe_swiglu_shared_forward_q4k_q5_0_gated_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q6k_parallel_workspace Forward
int moe_swiglu_shared_forward_q4k_q6k_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q6k_workspace Forward
int moe_swiglu_shared_forward_q4k_q6k_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q8_0_gated_parallel_workspace Forward
int moe_swiglu_shared_forward_q4k_q8_0_gated_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q4k_q8_0_gated_workspace Forward
int moe_swiglu_shared_forward_q4k_q8_0_gated_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q8_0_gated_parallel_workspace Forward
int moe_swiglu_shared_forward_q8_0_gated_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_forward_q8_0_gated_workspace Forward
int moe_swiglu_shared_forward_q8_0_gated_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes)

Forward pass computation

moe_swiglu_shared_q4k_q8_0_gated_workspace_bytes Forward
size_t moe_swiglu_shared_q4k_q8_0_gated_workspace_bytes(int hidden_dim, int intermediate_dim)

moe_swiglu_shared_q8_0_gated_workspace_bytes Forward
size_t moe_swiglu_shared_q8_0_gated_workspace_bytes(int hidden_dim, int intermediate_dim)

swiglu_backward Backward
void swiglu_backward(const float * input, const float * d_output, float * d_input, int tokens, int dim)

SwiGLU backward pass Testtest_swiglu.py::TestSwiGLUBackward::test_backward_tokens test_swiglu.py::TestSwiGLUBackward::test_backward_single test_parity.py::test_swiglu_backward_parity

swiglu_backward_bf16 Backward
void swiglu_backward_bf16(const uint16_t * input, const uint16_t * d_output, uint16_t * d_input, int tokens, int dim)

Backward pass / gradient computation

swiglu_backward_exact Backward
void swiglu_backward_exact(const float * input, const float * d_output, float * d_input, int tokens, int dim)

SwiGLU backward pass (exact version using stdlib sigmoid) Testtest_swiglu.py::TestSwiGLUBackward::test_exact_vs_fast test_swiglu.py::TestSwiGLUBackward::test_exact_single

swiglu_forward Forward
void swiglu_forward(const float * input, float * output, int tokens, int dim)

SwiGLU forward pass Testtest_swiglu.py::TestSwiGLUForward::test_forward_tokens test_swiglu.py::TestSwiGLUForward::test_forward_single test_mlp.py::TestMLPForward::test_swiglu_mlp test_fused_swiglu_decode.py::TestFusedSwiGLUDecode::test_fused_swiglu_decode test_parity.py::test_swiglu_parity

swiglu_forward_bf16 Forward
void swiglu_forward_bf16(const uint16_t * input, uint16_t * output, int tokens, int dim)

Forward pass computation

swiglu_forward_exact Forward
void swiglu_forward_exact(const float * input, float * output, int tokens, int dim)

SwiGLU forward pass (exact version using stdlib sigmoid) Testtest_swiglu.py::TestSwiGLUForward::test_exact_vs_fast test_swiglu.py::TestSwiGLUForward::test_exact_single

swiglu_forward_ggml Forward
void swiglu_forward_ggml(const float * input, float * output, int tokens, int dim)

Forward pass computation

swiglu_forward_ggml_split Forward
void swiglu_forward_ggml_split(const float * gate, const float * up, float * output, int tokens, int dim)

Forward pass computation

swiglu_forward_pytorch_bf16_storage Forward
void swiglu_forward_pytorch_bf16_storage(const float * input, float * output, int tokens, int dim)

Forward pass computation

swiglu_forward_q8_k Forward
void swiglu_forward_q8_k(const float * input, void * output_q8, int tokens, int dim)

Forward pass computation

RoPE (Rotary Position Embedding) (81)
apply_rope Forward
void apply_rope(float * x, int seq_len, int head_dim)

apply_rope_inline Forward
void apply_rope_inline(float * q, float * k, const float * rope_cos, const float * rope_sin, int pos, int H, int KV, int AD)

ck_model_precompute_rope Forward
void ck_model_precompute_rope(void * model)

Precompute RoPE cos/sin caches. Call once after allocation, before inference.

ck_mrope_round_storage Forward
void ck_mrope_round_storage(float * data, size_t count, int storage_kind)

ck_multimodal_mrope_positions_2d Forward
int ck_multimodal_mrope_positions_2d(int32_t * positions, int total_tokens, int prefix_start, int position_base, int prefix_tokens, int grid_x, int grid_y, int text_pos)

ck_resolve_ggml_rope_multi_inplace Forward
ck_ggml_rope_multi_inplace_fn ck_resolve_ggml_rope_multi_inplace(void)

ck_rope_ensure_ggml_loaded Forward
void ck_rope_ensure_ggml_loaded(void)

ck_rope_reference_cosf Forward
float ck_rope_reference_cosf(float value)

ck_rope_reference_powf Forward
float ck_rope_reference_powf(float base, float exponent)

ck_rope_reference_sinf Forward
float ck_rope_reference_sinf(float value)

ck_rope_resolve_ggml_symbol Forward
void * ck_rope_resolve_ggml_symbol(const char * name)

ck_rope_resolve_system_math_f32 Forward
ck_rope_math_f32_fn ck_rope_resolve_system_math_f32(const char * name)

ck_test_rope Forward
void ck_test_rope(float * q, float * k, int n_tokens, int n_heads, int n_heads_kv, int head_dim, int pos_offset, float theta)

RoPE (Rotary Position Embedding)

ck_test_rope_interleaved Forward
void ck_test_rope_interleaved(float * q, float * k, int n_tokens, int n_heads, int n_heads_kv, int head_dim, int pos_offset, float theta)

RoPE with interleaved format (for llama.cpp compatibility)

deepseek_mla_partial_rope_concat_f32 Forward
void deepseek_mla_partial_rope_concat_f32(const float * q_nope, const float * q_pe, const float * k_nope, const float * k_pe, const float * cos, const float * sin, float * query, float * key, int tokens, int heads, int qk_nope_dim, int qk_rope_dim)

deepseek_mla_partial_rope_concat_packed_bf16_storage Forward
void deepseek_mla_partial_rope_concat_packed_bf16_storage(const float * q_packed, const float * k_nope, const float * kv_a_packed, const float * cos, const float * sin, float * query, float * key, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int qk_rope_dim)

deepseek_mla_partial_rope_concat_packed_f32 Forward
void deepseek_mla_partial_rope_concat_packed_f32(const float * q_packed, const float * k_nope, const float * kv_a_packed, const float * cos, const float * sin, float * query, float * key, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int qk_rope_dim)

ds_mla_apply_kimi_rope Forward
void ds_mla_apply_kimi_rope(const float * src, float * dst, const float * cos_row, const float * sin_row, int dim)

explicit_mrope_apply_ggml_exact Forward
int explicit_mrope_apply_ggml_exact(float * x, const int32_t * positions, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, const int sections, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, int rope_type)

explicit_mrope_apply_head Forward
void explicit_mrope_apply_head(float * x, const int32_t * positions, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, const int sections, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, int is_imrope)

fused_rope_inplace Forward
void fused_rope_inplace(float * q, float * k, const float * rope_cos, const float * rope_sin, int pos, int num_heads, int num_kv_heads, int head_dim, int max_seq)

Fused RoPE application (in-place on pre-allocated buffers)

mega_fuse_rope_inplace_avx Forward
void mega_fuse_rope_inplace_avx(float * q, float * k, const float * rope_cos, const float * rope_sin, int pos, int num_heads, int num_kv_heads, int head_dim, int aligned_head_dim)

Apply RoPE to Q and K (in-place, from L1)

model_precompute_rope Forward
void model_precompute_rope(MODELModel * model)

mrope_qk_imrope_positions Forward
void mrope_qk_imrope_positions(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_text Forward
void mrope_qk_text(float * q, float * k, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_text_imrope Forward
void mrope_qk_text_imrope(float * q, float * k, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_text_imrope_bf16_pytorch_storage Forward
void mrope_qk_text_imrope_bf16_pytorch_storage(float * q, float * k, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_text_imrope_positions_bf16_pytorch_storage Forward
void mrope_qk_text_imrope_positions_bf16_pytorch_storage(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_vision Forward
void mrope_qk_vision(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_vision_bf16_pytorch_storage Forward
void mrope_qk_vision_bf16_pytorch_storage(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_vision_bf16_storage Forward
void mrope_qk_vision_bf16_storage(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_qk_vision_fp16_storage Forward
void mrope_qk_vision_fp16_storage(float * q, float * k, const int32_t * positions, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, int section_0, int section_1, int section_2, int section_3, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

mrope_rotate_pair Forward
void mrope_rotate_pair(float x0, float x1, float cos_theta, float sin_theta, float * out0, float * out1)

qwen2_0_5b_decode_precompute_rope Forward
void qwen2_0_5b_decode_precompute_rope(QWEN2_0_5B_DECODEModel * model)

qwen4_rope_split_inplace Forward
void qwen4_rope_split_inplace(float * vector, int rotary_dim, int position, float theta)

rope_apply_decode_pairwise_llama_cpu Forward
void rope_apply_decode_pairwise_llama_cpu(float * rows, const float * cos_row, const float * sin_row, int num_heads, int aligned_head_dim, int rotary_dim)

rope_apply_head Forward
void rope_apply_head(float * x, const float * cos_cache, const float * sin_cache, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

rope_apply_head_pairwise Forward
void rope_apply_head_pairwise(float * x, const float * cos_cache, const float * sin_cache, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

rope_backward Backward
void rope_backward(const float * d_out, float * d_x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)

RoPE backward (inverse rotation) Testtest_rope.py::TestRoPEBackward::test_rope_backward test_rope.py::TestRoPEBackward::test_rope_backward_vs_separate

rope_backward_apply_head_pairwise Backward
void rope_backward_apply_head_pairwise(const float * d_out, float * d_x, const float * cos_cache, const float * sin_cache, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Backward pass / gradient computation

rope_backward_bf16 Backward
void rope_backward_bf16(const uint16_t * d_out, uint16_t * d_x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float * scratch_d_out, float * scratch_d_x)

Backward pass / gradient computation

rope_backward_inplace Backward
void rope_backward_inplace(float * d_x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)

RoPE backward in-place (overwrite with inverse rotation) Testtest_rope.py::TestRoPEBackward::test_rope_backward_inplace

rope_backward_qk Backward
void rope_backward_qk(const float * d_q_out, const float * d_k_out, float * d_q, float * d_k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)

RoPE backward for both dQ and dK Testtest_rope.py::TestRoPEBackward::test_rope_backward_qk

rope_backward_qk_bf16 Backward
void rope_backward_qk_bf16(const uint16_t * d_q_out, const uint16_t * d_k_out, uint16_t * d_q, uint16_t * d_k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float * scratch_dq_out, float * scratch_dq, float * scratch_dk_out, float * scratch_dk)

Backward pass / gradient computation

rope_backward_qk_pairwise_with_rotary_dim Backward
void rope_backward_qk_pairwise_with_rotary_dim(const float * d_q_out, const float * d_k_out, float * d_q, float * d_k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Backward pass / gradient computation

rope_forward Forward
void rope_forward(float * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)

RoPE forward (head-major layout, in-place) Testtest_rope.py::TestRoPEForward::test_rope_forward test_rope.py::TestRoPEForward::test_rope_vs_separate test_parity.py::test_rope_parity

rope_forward_bf16 Forward
void rope_forward_bf16(uint16_t * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float * scratch)

Forward pass computation

rope_forward_bf16_with_rotary_dim Forward
void rope_forward_bf16_with_rotary_dim(uint16_t * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float * scratch)

Forward pass computation

rope_forward_q_split_direct_f32 Forward
void rope_forward_q_split_direct_f32(float * q, const float * freq_factors, int use_freq_factors, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base)

Forward pass computation

rope_forward_qk Forward
void rope_forward_qk(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)

RoPE forward for both Q and K (common inference pattern) Testtest_rope.py::TestRoPEForward::test_rope_forward_qk test_fused_attention_decode.py::TestFusedAttentionDecode::test_qk_rope test_parity.py::test_rope_qk_parity

rope_forward_qk_bf16 Forward
void rope_forward_qk_bf16(uint16_t * q, uint16_t * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float * scratch_q, float * scratch_k)

Forward pass computation

rope_forward_qk_bf16_with_rotary_dim Forward
void rope_forward_qk_bf16_with_rotary_dim(uint16_t * q, uint16_t * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float * scratch_q, float * scratch_k)

Forward pass computation

rope_forward_qk_pairwise_llama_cpu Forward
void rope_forward_qk_pairwise_llama_cpu(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Forward pass computation

rope_forward_qk_pairwise_with_rotary_dim Forward
void rope_forward_qk_pairwise_with_rotary_dim(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Forward pass computation

rope_forward_qk_split_direct_f32 Forward
void rope_forward_qk_split_direct_f32(float * q, float * k, const float * freq_factors, int use_freq_factors, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base)

Forward pass computation

rope_forward_qk_split_direct_token_range_f32 Forward
void rope_forward_qk_split_direct_token_range_f32(float * q, float * k, const float * freq_factors, int use_freq_factors, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base, int token_begin, int token_end)

Forward pass computation

rope_forward_qk_split_llama_token_range_f32 Forward
void rope_forward_qk_split_llama_token_range_f32(float * q, float * k, const float * freq_factors, int use_freq_factors, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base, int token_begin, int token_end)

Forward pass computation

rope_forward_qk_strided Forward
void rope_forward_qk_strided(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int q_stride_tokens, int k_stride_tokens)

RoPE forward for both Q and K with custom strides (KV cache layouts) Testtest_rope.py::TestRoPEForward::test_rope_forward_qk_strided test_kv_cache_attention.py::TestKVCacheAttention::test_qk_rope_strided

rope_forward_qk_strided_with_rotary_dim Forward
void rope_forward_qk_strided_with_rotary_dim(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int q_stride_tokens, int k_stride_tokens, int rotary_dim)

Forward pass computation

rope_forward_qk_with_rotary_dim Forward
void rope_forward_qk_with_rotary_dim(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Forward pass computation

rope_forward_qk_with_rotary_dim_cache_stride Forward
void rope_forward_qk_with_rotary_dim_cache_stride(float * q, float * k, const float * cos_cache, const float * sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, int cache_rotary_dim)

Forward pass computation

rope_forward_split_direct_one Forward
void rope_forward_split_direct_one(float * x, const float * freq_factors, int use_freq_factors, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float freq_base)

Forward pass computation

rope_forward_strided Forward
void rope_forward_strided(float * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int head_stride_tokens)

RoPE forward with custom head stride (for KV cache layouts) Testtest_rope.py::TestRoPEForward::test_rope_strided test_kv_cache_attention.py::TestKVCacheAttention::test_rope_decode

rope_forward_strided_with_rotary_dim Forward
void rope_forward_strided_with_rotary_dim(float * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int head_stride_tokens, int rotary_dim)

Forward pass computation

rope_forward_with_rotary_dim Forward
void rope_forward_with_rotary_dim(float * x, const float * cos_cache, const float * sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)

Forward pass computation

rope_precompute_cache
void rope_precompute_cache(float * cos_cache, float * sin_cache, int max_seq_len, int head_dim, float base, int rotary_dim, const char * scaling_type, float scaling_factor)

Precompute RoPE cos/sin cache with rotary_dim and scaling support Testtest_rope.py::TestRoPECache::test_cache_computation test_rope.py::TestRoPECache::test_cache_values

rope_precompute_cache_llama_cpu
void rope_precompute_cache_llama_cpu(float * cos_cache, float * sin_cache, int max_seq_len, int head_dim, float base, int rotary_dim, const char * scaling_type, float scaling_factor)

rope_precompute_cache_split
void rope_precompute_cache_split(float * cos_cache, float * sin_cache, int max_seq_len, int head_dim, float base)

Precompute RoPE cos/sin cache (split layout: head_dim/2) Legacy layout used before rotary_dim/scaling support.

text_mrope_apply_head Forward
void text_mrope_apply_head(float * x, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int n_dims, const int sections, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, int is_imrope)

text_mrope_apply_positions_pytorch_bf16_storage Forward
void text_mrope_apply_positions_pytorch_bf16_storage(float * x, const int32_t * positions, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, const int sections, float freq_base, float freq_scale)

text_mrope_apply_pytorch_bf16_storage Forward
void text_mrope_apply_pytorch_bf16_storage(float * x, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int n_dims, float freq_base, float freq_scale)

text_mrope_yarn Forward
void text_mrope_yarn(float theta_extrap, float freq_scale, const float corr_dims, int chan, float ext_factor, float attn_factor, float * cos_theta, float * sin_theta)

vision_mrope_apply_head Forward
void vision_mrope_apply_head(float * x, const int32_t * positions, int num_tokens, int head_dim, int aligned_head_dim, int n_dims, const int sections, int n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow)

vision_mrope_yarn Forward
void vision_mrope_yarn(float theta_extrap, float freq_scale, const float corr_dims, int chan, float ext_factor, float attn_factor, float * cos_theta, float * sin_theta)

vision_mrope_yarn_corr_dim Forward
float vision_mrope_yarn_corr_dim(int n_dims, int n_ctx_orig, float n_rot, float base)

vision_mrope_yarn_corr_dims Forward
void vision_mrope_yarn_corr_dims(int n_dims, int n_ctx_orig, float freq_base, float beta_fast, float beta_slow, float dims)

vision_mrope_yarn_ramp Forward
float vision_mrope_yarn_ramp(float low, float high, int chan)

yarn_rope_cache_contiguous_positions_f32
void yarn_rope_cache_contiguous_positions_f32(float * cos_cache, float * sin_cache, int num_tokens, int rotary_dim, float freq_base, float factor, int original_context, float beta_fast, float beta_slow, float mscale, float mscale_all_dim)

yarn_rope_cache_explicit_positions_bf16
void yarn_rope_cache_explicit_positions_bf16(uint16_t * cos_cache, uint16_t * sin_cache, const int32_t * positions, int num_tokens, int rotary_dim, float freq_base, float factor, int original_context, float beta_fast, float beta_slow, float mscale, float mscale_all_dim)

yarn_rope_cache_explicit_positions_f32
void yarn_rope_cache_explicit_positions_f32(float * cos_cache, float * sin_cache, const int32_t * positions, int num_tokens, int rotary_dim, float freq_base, float factor, int original_context, float beta_fast, float beta_slow, float mscale, float mscale_all_dim)

yarn_rope_cache_explicit_positions_impl
void yarn_rope_cache_explicit_positions_impl(float * cos_f32, float * sin_f32, uint16_t * cos_bf16, uint16_t * sin_bf16, const int32_t * positions, int num_tokens, int rotary_dim, float freq_base, float factor, int original_context, float beta_fast, float beta_slow, float mscale, float mscale_all_dim)

Fully Connected Layers (2)
fc1_backward_kernel Backward
void fc1_backward_kernel(const float * d_output, const float * fc1_input, const float * W_fc1, float * d_input, float * d_W_fc1, float * d_b_fc1, int T, int aligned_in, int aligned_out, int num_threads)

Backward pass / gradient computation

fc2_backward_kernel Backward
void fc2_backward_kernel(const float * d_output, const float * fc2_input, const float * W_fc2, float * d_input, float * d_W_fc2, float * d_b_fc2, int T, int aligned_in, int aligned_out, int num_threads)

Backward pass / gradient computation

Other Functions (1368)
_Static_assert Forward
 _Static_assert(sizeof(MagicHeader), "MagicHeader must be 64 bytes")

__attribute__ Forward
 __attribute__((unused))

accum_q4_k_packed_meta_x16_q8_k_block Forward
void accum_q4_k_packed_meta_x16_q8_k_block(float acc, const block_q4_K_packed_meta_x16 * w, int active, const block_q8_K * x)

accum_q4_k_packed_meta_x16_q8_k_block_mreuse Forward
void accum_q4_k_packed_meta_x16_q8_k_block_mreuse(float acc, const block_q4_K_packed_meta_x16 * w, int active, const block_q8_K * A, int blocks_per_vec, int block_index, int m0, int m_count)

accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4 Forward
void accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(float acc, const block_q4_K_packed_meta_x16 * w, int active, const block_q8_K * A, int blocks_per_vec, int block_index, int m0, int m_count)

accum_q4_k_packed_meta_x8_q8_k_block Forward
void accum_q4_k_packed_meta_x8_q8_k_block(float acc, const block_q4_K_packed_meta_x8 * w, int active, const block_q8_K * x)

accum_q4_k_packed_meta_x8_q8_k_block_mreuse Forward
void accum_q4_k_packed_meta_x8_q8_k_block_mreuse(float acc, const block_q4_K_packed_meta_x8 * w, int active, const block_q8_K * A, int blocks_per_vec, int block_index, int m0, int m_count)

accum_q4_k_packed_meta_x8_q8_k_gemv_block Forward
void accum_q4_k_packed_meta_x8_q8_k_gemv_block(float acc, float acc_min, const block_q4_K_packed_meta_x8 * w, int active, const block_q8_K * x)

accum_q4_k_packed_meta_x8_q8_k_superblock Forward
void accum_q4_k_packed_meta_x8_q8_k_superblock(float acc, float acc_min, const block_q4_K_packed_meta_x8 * w, int active, const block_q8_K * x)

accum_q4_k_packed_meta_x8_q8_k_superblock_rows Forward
void accum_q4_k_packed_meta_x8_q8_k_superblock_rows(float acc, float acc_min, const block_q4_K_packed_meta_x8 * w, int active, const block_q8_K * x, int rows)

accum_q4_k_packed_u8_x16_q8_k_block Forward
void accum_q4_k_packed_u8_x16_q8_k_block(float acc, const block_q4_K_packed_u8_x16 * w, int active, const block_q8_K * x)

accum_q4_k_packed_vnni_x16_q8_k_16m_superblock Forward
void accum_q4_k_packed_vnni_x16_q8_k_16m_superblock(float acc, float acc_min, const block_q4_K_packed_vnni_x16 * w, const block_q8_K * x, int rows)

accum_q4_k_packed_vnni_x16_q8_k_gemv_block Forward
void accum_q4_k_packed_vnni_x16_q8_k_gemv_block(float acc, float acc_min, const block_q4_K_packed_vnni_x16 * w, const block_q8_K * x)

accum_q4_k_packed_vnni_x8_q8_k_4m_superblock Forward
void accum_q4_k_packed_vnni_x8_q8_k_4m_superblock(float acc, float acc_min, const block_q4_K_packed_vnni_x8 * w, const block_q8_K * x, int rows)

adamw_clip_update_multi_f32 Forward
void adamw_clip_update_multi_f32(float *const * grads, float *const * weights, float *const * m_states, float *const * v_states, const size_t * numels, int tensor_count, float lr, float beta1, float beta2, float eps, float weight_decay, float max_grad_norm, int step)

adamw_update_bf16 Forward
void adamw_update_bf16(const uint16_t * grad, uint16_t * weight, float * m, float * v, size_t numel, float lr, float beta1, float beta2, float eps, float weight_decay, int step)

AdamW optimizer update (bf16 weights/gradients, fp32 optimizer state)

adamw_update_f32 Forward
void adamw_update_f32(const float * grad, float * weight, float * m, float * v, size_t numel, float lr, float beta1, float beta2, float eps, float weight_decay, int step)

adamw_update_f32_impl Forward
void adamw_update_f32_impl(const float * grad, float * weight, float * m, float * v, size_t numel, float lr, float beta1, float beta2, float eps, float weight_decay, int step)

AdamW optimizer update (fp32 version)

add_backward_bf16 Backward
void add_backward_bf16(const uint16_t * d_y, uint16_t * d_a, uint16_t * d_b, size_t n)

Backward pass / gradient computation

add_bias_tile Forward
void add_bias_tile(float * out, const float * bias, int tile_m, int out_dim)

add_forward_2d_bf16 Forward
void add_forward_2d_bf16(const uint16_t * a, const uint16_t * b, uint16_t * y, int tokens, int dim, int aligned_dim)

Forward pass computation

add_forward_bf16 Forward
void add_forward_bf16(const uint16_t * a, const uint16_t * b, uint16_t * y, size_t n)

Forward pass computation

add_forward_f32 Forward
void add_forward_f32(const float * a, const float * b, float * y, size_t n)

Element-wise add: y = a + b Testtest_add.py::TestAddForward::test_add_forward_f32 test_add.py::TestAddForward::test_add_inplace_f32 test_multi_layer_parity.py::TestMultiLayerParity::test_residual_add

add_inplace_bf16 Forward
void add_inplace_bf16(uint16_t * a, const uint16_t * b, size_t n)

add_inplace_f32 Forward
void add_inplace_f32(float * a, const float * b, size_t n)

add_scaled_forward_bf16 Forward
void add_scaled_forward_bf16(const uint16_t * a, const uint16_t * b, uint16_t * y, float alpha, size_t n)

Forward pass computation

add_scaled_inplace_bf16 Forward
void add_scaled_inplace_bf16(uint16_t * a, const uint16_t * b, float alpha, size_t n)

add_stream_inplace Forward
void add_stream_inplace(float * a, const float * b, size_t n)

add_stream_reorder_2d Forward
void add_stream_reorder_2d(float * main_inout, float * aux_scratch, int grid_h, int grid_w, int embed_dim, int merge_size)

align_up Forward
size_t align_up(size_t n, size_t align)

align_up_bytes Forward
size_t align_up_bytes(size_t n, size_t align)

align_up_elems Forward
size_t align_up_elems(size_t elems, size_t elem_bytes, size_t align_bytes)

align_up_size Forward
size_t align_up_size(size_t value, size_t align)

amx_available Forward
bool amx_available(void)

apply_bpe_merges Forward
int apply_bpe_merges(CKTrueBPE * bpe, CKBPETokenList * list)

apply_chat_template Forward
char * apply_chat_template(const ChatTemplate * tmpl, const char * system, const char * user)

arena_for_role Forward
CKMemArenaKind arena_for_role(CKBufferRole role)

argmax_f32 Forward
int argmax_f32(const float * scores, int n)

Find index of maximum value.

assistant_layer_scale_forward Forward
void assistant_layer_scale_forward(float * hidden, const float * scale, int tokens, int embed_dim)

Forward pass computation

attn_gate_softplus_mul_forward Forward
void attn_gate_softplus_mul_forward(const float * x, const float * gate, float * out, int rows, int num_heads, int state_dim)

Forward pass computation

audio_conv1d_channel_major_f32 Forward
int audio_conv1d_channel_major_f32(const float * input, const float * weight, const float * bias, float * output, int input_channels, int output_channels, int input_frames, int kernel_size, int stride, int padding, int output_frames)

audio_conv2d_whc_grouped_f32 Forward
int audio_conv2d_whc_grouped_f32(const float * input, const float * weight, const float * bias, float * output, int input_width, int input_height, int input_channels, int output_channels, int kernel_width, int kernel_height, int stride_width, int stride_height, int padding_width, int padding_height, int groups, int output_width, int output_height)

audio_feature_normalize_per_feature_f32 Forward
int audio_feature_normalize_per_feature_f32(const float * input, float * output, int channels, int frames, float epsilon)

audio_glu_split_channel_major_f32 Forward
int audio_glu_split_channel_major_f32(const float * input, float * output, int channels, int frames)

audio_hz_to_mel_slaney Forward
double audio_hz_to_mel_slaney(double hz)

audio_log_mel_time_major_f32 Forward
int audio_log_mel_time_major_f32(const float * power, const float * mel_filters, float * log_mel, int frames, int bins, int channels, float epsilon)

audio_mel_to_hz_slaney Forward
double audio_mel_to_hz_slaney(double mel)

audio_pad_or_truncate_f32 Forward
int audio_pad_or_truncate_f32(const float * input, int input_frames, float * output, int output_frames)

audio_pcm_s16_to_mono_f32 Forward
int audio_pcm_s16_to_mono_f32(const int16_t * interleaved, int n_frames, int n_channels, float * mono)

audio_preemphasis_f32 Forward
int audio_preemphasis_f32(const float * input, float * output, int frames, float coefficient)

audio_relative_shift_f32 Forward
int audio_relative_shift_f32(const float * raw_scores, float * scores, int heads, int query_frames)

audio_resample_linear_f32 Forward
int audio_resample_linear_f32(const float * input, int input_frames, int input_rate, float * output, int output_frames, int output_rate)

audio_resample_windowed_sinc_f32 Forward
int audio_resample_windowed_sinc_f32(const float * input, int input_frames, int input_rate, float * output, int output_frames, int output_rate, int radius)

audio_resampled_frame_count Forward
int audio_resampled_frame_count(int input_frames, int input_rate, int output_rate)

audio_stft_power_centered_window_f32 Forward
int audio_stft_power_centered_window_f32(const float * samples, int n_samples, const float * window, int window_length, const float * cos_table, const float * sin_table, int n_fft, int hop_length, int reflect_padding, float * power, int n_frames)

audio_stft_power_fft400_f32 Forward
int audio_stft_power_fft400_f32(const float * samples, int n_samples, const float * window, const float * cos_table, const float * sin_table, int hop_length, float * power, int n_frames, float * fft_scratch)

audio_stft_power_fft400_frame_f32 Forward
void audio_stft_power_fft400_frame_f32(const float * samples, int n_samples, int frame, const float * window, const float * cos_table, const float * sin_table, float * power, float * fft_scratch)

audio_stft_power_precomputed_f32 Forward
int audio_stft_power_precomputed_f32(const float * samples, int n_samples, const float * window, const float * cos_table, const float * sin_table, int n_fft, int hop_length, float * power, int n_frames)

audio_stft_precompute_tables_f32 Forward
int audio_stft_precompute_tables_f32(int n_fft, float * window, float * cos_table, float * sin_table)

audio_transpose_channel_to_token_f32 Forward
int audio_transpose_channel_to_token_f32(const float * input, float * output, int channels, int frames)

audio_wav_decode_memory_pcm16_mono_f32 Forward
int audio_wav_decode_memory_pcm16_mono_f32(const uint8_t * bytes, size_t byte_count, float * mono, int mono_capacity, CKAudioWavInfo * info)

audio_wav_decode_memory_pcm16_mono_window_f32 Forward
int audio_wav_decode_memory_pcm16_mono_window_f32(const uint8_t * bytes, size_t byte_count, int start_frame, float * mono, int mono_capacity, CKAudioWavInfo * info)

audio_wav_decode_pcm16_mono_f32 Forward
int audio_wav_decode_pcm16_mono_f32(const uint8_t * bytes, size_t byte_count, const CKAudioWavInfo * info, float * mono, int mono_capacity)

audio_wav_parse_memory Forward
int audio_wav_parse_memory(const uint8_t * bytes, size_t byte_count, CKAudioWavInfo * info)

audio_whisper_log_mel_from_power_reference_f32 Forward
int audio_whisper_log_mel_from_power_reference_f32(const float * power, const float * mel_filters, int n_mels, int n_frames, float * log_mel)

audio_whisper_log_mel_reference_f32 Forward
int audio_whisper_log_mel_reference_f32(const float * samples, int n_samples, const float * mel_filters, int n_mels, float * power_scratch, float * log_mel, int n_frames)

audio_whisper_log_mel_window_wav_pcm16_f32 Forward
int audio_whisper_log_mel_window_wav_pcm16_f32(const uint8_t * bytes, size_t byte_count, int start_frame, int target_sample_rate, const float * window, const float * cos_table, const float * sin_table, const float * mel_filters, int n_mels, int output_frames, float * log_mel)

audio_whisper_mel_filters_slaney_f32 Forward
int audio_whisper_mel_filters_slaney_f32(int sample_rate, int n_fft, int n_mels, float * mel_filters)

audio_whisper_stft_power_reference_f32 Forward
int audio_whisper_stft_power_reference_f32(const float * samples, int n_samples, float * power, int n_frames)

axpy_2d_f32 Forward
void axpy_2d_f32(float * Y, const float * X, float alpha, int num_tokens, int dim, int y_stride, int x_stride)

Batched AXPY for 2D tensors: Y[t,:] += alpha * X[t,:].

axpy_f32 Forward
void axpy_f32(float * y, const float * x, float alpha, int n)

In-place AXPY: y += alpha * x.

axpy_zero_f32 Forward
void axpy_zero_f32(float * y, const float * x, float alpha, int n)

Zero output then accumulate: y = 0; y += alpha * x.

barrier_init Forward
void barrier_init(ck_barrier_t * b, int n_threads)

barrier_wait Forward
void barrier_wait(ck_barrier_t * b)

Spin-wait barrier. All threads must call this. Uses phase counter to allow re-use without reset.

bf16_tensor_to_float Forward
void bf16_tensor_to_float(const uint16_t * src, float * dst, size_t count)

bf16_to_float Forward
float bf16_to_float(uint16_t v)

buffer_bytes Forward
size_t buffer_bytes(const CKIRV2Buffer * buf, const CKModelConfig * cfg, const CKV2AlignInfo * align)

buffer_enabled Forward
int buffer_enabled(const CKIRV2Graph * graph, const CKIRV2Buffer * buf, int training_enabled)

build_plan Forward
int build_plan(const CKIRV2Graph * graph, CKMemPlan * plan, size_t alignment_bytes, int training_enabled, int tokens_override)

bump_bytes Forward
size_t bump_bytes(size_t * off, size_t bytes, size_t align)

byte_to_gpt2 Forward
int byte_to_gpt2(unsigned char byte, char * out)

cache_compare
int cache_compare(const void * a, const void * b)

ce_legacy_mode_enabled Forward
int ce_legacy_mode_enabled(void)

ce_targets_all_valid_no_ignore Forward
int ce_targets_all_valid_no_ignore(const int32_t * targets, int tokens, int vocab_size)

ck_accum_multi_parallel_work Forward
void ck_accum_multi_parallel_work(int ith, int nth, void * argp)

ck_accum_parallel_work Forward
void ck_accum_parallel_work(int ith, int nth, void * argp)

ck_adamw_multi_parallel_work Forward
void ck_adamw_multi_parallel_work(int ith, int nth, void * argp)

ck_adamw_parallel_work Forward
void ck_adamw_parallel_work(int ith, int nth, void * argp)

ck_add_inplace Forward
void ck_add_inplace(float * dst, const float * src, int tokens, int aligned_embed_dim)

ck_audio_conv1d_channel_major_f32_work Forward
void ck_audio_conv1d_channel_major_f32_work(int ith, int nth, void * opaque)

ck_audio_conv2d_whc_grouped_f32_range Forward
void ck_audio_conv2d_whc_grouped_f32_range(int begin, int end, void * opaque)

ck_audio_glu_split_f32_range Forward
void ck_audio_glu_split_f32_range(int begin, int end, void * opaque)

ck_audio_relative_shift_f32_range Forward
void ck_audio_relative_shift_f32_range(int begin, int end, void * opaque)

ck_available_logical_cpus Forward
int ck_available_logical_cpus(void)

ck_bf16_convert_work Forward
void ck_bf16_convert_work(int ith, int nth, void * opaque)

ck_bf16_dot_contract Forward
float ck_bf16_dot_contract(const float * a, const float * b, int count)

ck_bf16_round Forward
float ck_bf16_round(float value)

ck_bf16_round_work Forward
void ck_bf16_round_work(int ith, int nth, void * opaque)

ck_bf16_to_f32 Forward
float ck_bf16_to_f32(uint16_t v)

ck_bind_deltanet_llama_libm Forward
void ck_bind_deltanet_llama_libm(void)

ck_bind_deltanet_pytorch_primitives Forward
void ck_bind_deltanet_pytorch_primitives(void)

ck_bind_hybrid_llama_libm Forward
void ck_bind_hybrid_llama_libm(void)

ck_bind_recurrent_llama_libm Forward
void ck_bind_recurrent_llama_libm(void)

ck_buffer_should_alloc Forward
int ck_buffer_should_alloc(const CKBufferSpec * spec)

ck_buffer_uses_weight_dtype Forward
int ck_buffer_uses_weight_dtype(const CKBufferSpec * spec)

ck_build_decoder_backward_ir Backward
int ck_build_decoder_backward_ir(const CKIRGraph * forward, CKIRGraph * backward)

Build a naive backward IR graph from a forward decoder IR.

ck_build_decoder_ir Forward
int ck_build_decoder_ir(const CKModelConfig * cfg, CKIRGraph * graph)

Build a simple decoder-only IR graph for the given config.

ck_bump_alloc_free Forward
void ck_bump_alloc_free(ck_bump_alloc_t * alloc)

ck_bump_alloc_init Forward
int ck_bump_alloc_init(ck_bump_alloc_t * alloc, const char * weights_path, size_t total_size, size_t weights_base, size_t activations_base)

ck_bump_alloc_needs_weight_materialization Forward
int ck_bump_alloc_needs_weight_materialization(const ck_bump_alloc_t * alloc)

ck_bump_alloc_reset Forward
void ck_bump_alloc_reset(ck_bump_alloc_t * alloc)

ck_codegen_c_skeleton Forward
void ck_codegen_c_skeleton(const CKIRGraph * forward, const CKIRGraph * backward, FILE * out)

Emit a C skeleton for forward + backward execution based on the IR.

ck_codegen_emit_runtime Forward
int ck_codegen_emit_runtime(const CKIRGraph * forward, const char * path, CKEmitMode mode)

Emit a C runtime file that stitches kernels for the given forward IR.

ck_codegen_v2_dtype_name Forward
const char * ck_codegen_v2_dtype_name(CKDataType dtype)

ck_codegen_v2_emit_dispatch Forward
void ck_codegen_v2_emit_dispatch(FILE * out, const CKIRV2Graph * graph)

ck_codegen_v2_emit_preamble Forward
int ck_codegen_v2_emit_preamble(FILE * out)

ck_codegen_v2_emit_runtime Forward
int ck_codegen_v2_emit_runtime(const CKIRV2Graph * graph, const char * path, CKEmitMode mode)

Emit a C runtime file from a CKIRV2Graph.

ck_codegen_v2_emit_schedule Forward
void ck_codegen_v2_emit_schedule(FILE * out, const CKIRV2Graph * graph, const char * prefill_runtime, const char * decode_runtime, const char * backward_runtime)

ck_codegen_v2_emit_sections Forward
void ck_codegen_v2_emit_sections(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * prefill_plan, const CKMemPlan * decode_plan, const CKMemPlan * backward_plan)

ck_codegen_v2_emit_struct Forward
void ck_codegen_v2_emit_struct(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * plan, const char * tag)

ck_debug_check_buffer Forward
void ck_debug_check_buffer(const char * stage, const float * buf, int size)

ck_debug_check_q4k_weights Forward
void ck_debug_check_q4k_weights(const char * stage, const void * q4_buf, int num_blocks)

ck_debug_check_q8k Forward
void ck_debug_check_q8k(const char * stage, const void * q8_buf, int num_blocks)

ck_deltanet_ceil_log2 Forward
int ck_deltanet_ceil_log2(int value)

ck_deltanet_force_ref Forward
int ck_deltanet_force_ref(void)

ck_deltanet_pytorch_gate_values Forward
void ck_deltanet_pytorch_gate_values(const float * g, const float * beta, float * gate_values, float * beta_values, int num_heads)

ck_deltanet_pytorch_outer_sum Forward
void ck_deltanet_pytorch_outer_sum(const float * matrix, const float * row_weights, float * output, int state_dim)

ck_dot_f16_f16_local Forward
float ck_dot_f16_f16_local(const uint16_t * w, const uint16_t * x, int k)

ck_dot_f32 Forward
float ck_dot_f32(const float * a, const float * b, int len)

ck_dot_q6_k_q8_k_fast_or_ref Forward
float ck_dot_q6_k_q8_k_fast_or_ref(const block_q6_K * w, const block_q8_K * x, int K)

ck_dtype_block_bytes Forward
size_t ck_dtype_block_bytes(CKDataType dt)

Get bytes per block for quantized types.

ck_dtype_block_size Forward
size_t ck_dtype_block_size(CKDataType dt)

Get the number of elements per quantization block.

ck_dtype_bytes Forward
size_t ck_dtype_bytes(CKDataType dt)

Get bytes per element for non-quantized types.

ck_dtype_is_quantized Forward
int ck_dtype_is_quantized(CKDataType dt)

Check if a data type is block-quantized (GGML-style)

ck_dtype_row_bytes Forward
size_t ck_dtype_row_bytes(CKDataType dt, size_t n_elements)

Calculate total bytes for n_elements of given dtype.

ck_dtype_supported Forward
int ck_dtype_supported(CKDataTypeMask mask, CKDataType dt)

ck_env_int_default Forward
int ck_env_int_default(const char * name, int fallback)

ck_env_truthy_or_qwen3vl_ocr_profile Forward
int ck_env_truthy_or_qwen3vl_ocr_profile(const char * name)

ck_env_value_truthy Forward
int ck_env_value_truthy(const char * v)

ck_expf Forward
float ck_expf(float x)

ck_f16_resolve_ggml_build_forward_expand Forward
ck_f16_ggml_build_forward_expand_fn ck_f16_resolve_ggml_build_forward_expand(void)

Forward pass computation

ck_f16_resolve_ggml_cpu_init Forward
ck_f16_ggml_cpu_init_fn ck_f16_resolve_ggml_cpu_init(void)

ck_f16_resolve_ggml_free Forward
ck_f16_ggml_free_fn ck_f16_resolve_ggml_free(void)

ck_f16_resolve_ggml_get_data Forward
ck_f16_ggml_get_data_fn ck_f16_resolve_ggml_get_data(void)

ck_f16_resolve_ggml_get_data_f32 Forward
ck_f16_ggml_get_data_f32_fn ck_f16_resolve_ggml_get_data_f32(void)

ck_f16_resolve_ggml_graph_compute_with_ctx Forward
ck_f16_ggml_graph_compute_with_ctx_fn ck_f16_resolve_ggml_graph_compute_with_ctx(void)

ck_f16_resolve_ggml_init Forward
ck_f16_ggml_init_fn ck_f16_resolve_ggml_init(void)

ck_f16_resolve_ggml_mul_mat Forward
ck_f16_ggml_mul_mat_fn ck_f16_resolve_ggml_mul_mat(void)

ck_f16_resolve_ggml_new_graph Forward
ck_f16_ggml_new_graph_fn ck_f16_resolve_ggml_new_graph(void)

ck_f16_resolve_ggml_new_tensor_2d Forward
ck_f16_ggml_new_tensor_2d_fn ck_f16_resolve_ggml_new_tensor_2d(void)

ck_f32_to_f16_row_local Forward
void ck_f32_to_f16_row_local(uint16_t * dst, const float * src, int n)

ck_fast_expf Forward
float ck_fast_expf(float x)

ck_find_buffer_spec Forward
const CKBufferSpec * ck_find_buffer_spec(const char * name)

ck_find_kernel_spec Forward
const CKKernelSpec * ck_find_kernel_spec(const char * name)

ck_first_layer_buffer_name Forward
const char * ck_first_layer_buffer_name(void)

ck_flash_attn_choose_tile_k Forward
int ck_flash_attn_choose_tile_k(int D_h)

ck_flash_attn_fast_exp_kind Forward
int ck_flash_attn_fast_exp_kind(void)

ck_flash_attn_tile_k Forward
int ck_flash_attn_tile_k(int D_h)

ck_fma_f32_to_f16 Forward
void ck_fma_f32_to_f16(const float * a, const float * b, const float * c, uint16_t * dst, int n)

FMA in FP32, store result as FP16: dst = a * b + c.

ck_fp16_to_fp32 Forward
float ck_fp16_to_fp32(ck_half h)

ck_fp16_to_fp32_2d Forward
void ck_fp16_to_fp32_2d(const uint16_t * src, float * dst, int rows, int cols, int src_stride, int dst_stride)

Convert 2D FP16 matrix to FP32 with strided access.

ck_fp16_to_fp32_row Forward
void ck_fp16_to_fp32_row(const uint16_t * src, float * dst, int n)

Convert FP16 row to FP32 (auto-select best implementation)

ck_fp16_to_fp32_scalar Forward
float ck_fp16_to_fp32_scalar(uint16_t h)

ck_fp16_to_fp32_soft Forward
float ck_fp16_to_fp32_soft(ck_half h)

Convert FP16 (ck_half) to FP32 — software implementation.

ck_fp32_from_bits Forward
float ck_fp32_from_bits(uint32_t u32)

ck_fp32_to_bits Forward
uint32_t ck_fp32_to_bits(float f)

ck_fp32_to_fp16 Forward
ck_half ck_fp32_to_fp16(float f)

ck_fp32_to_fp16_2d Forward
void ck_fp32_to_fp16_2d(const float * src, uint16_t * dst, int rows, int cols, int src_stride, int dst_stride)

Convert 2D FP32 matrix to FP16 with strided access.

ck_fp32_to_fp16_inplace Forward
void ck_fp32_to_fp16_inplace(float * data, void * scratch, int n)

Convert FP32 to FP16 in-place using scratch buffer.

ck_fp32_to_fp16_row Forward
void ck_fp32_to_fp16_row(const float * src, uint16_t * dst, int n)

Convert FP32 row to FP16 (auto-select best implementation)

ck_fp32_to_fp16_scalar Forward
uint16_t ck_fp32_to_fp16_scalar(float f)

ck_fp32_to_fp16_soft Forward
ck_half ck_fp32_to_fp16_soft(float f)

Convert FP32 to FP16 (ck_half) — software implementation.

ck_gemv_bf16_rows Forward
void ck_gemv_bf16_rows(int begin, int end, void * opaque)

ck_gemv_bf16_storage_rows Forward
void ck_gemv_bf16_storage_rows(int begin, int end, void * opaque)

ck_get_block_q4_k_size Forward
int ck_get_block_q4_k_size(void)

Get Q4_K block size in bytes.

ck_get_block_q6_k_size Forward
int ck_get_block_q6_k_size(void)

Get Q6_K block size in bytes.

ck_get_block_q8_k_size Forward
int ck_get_block_q8_k_size(void)

Get Q8_K block size in bytes.

ck_get_capabilities Forward
ck_capability_t ck_get_capabilities(void)

Get current platform capabilities.

ck_get_num_threads Forward
int ck_get_num_threads(void)

ck_get_physical_cores Forward
int ck_get_physical_cores(void)

ck_get_qk_k Forward
int ck_get_qk_k(void)

Get QK_K (elements per super-block)

ck_get_threadpool Forward
ck_threadpool_t * ck_get_threadpool(void)

Get the global thread pool handle for dispatch. Convenience wrapper — initializes on first call.

ck_ggml_vec_dot_f32_contig Forward
float ck_ggml_vec_dot_f32_contig(const float * x, const float * y, int n)

ck_ggml_vec_soft_max_row Forward
double ck_ggml_vec_soft_max_row(int n, float * y, const float * x, float max)

ck_huge_alloc Forward
void * ck_huge_alloc(size_t bytes)

Allocate a large, contiguous memory region for model weights/activations.

ck_huge_free Forward
void ck_huge_free(void * ptr, size_t bytes)

Free memory allocated by ck_huge_alloc.

ck_ir_dump Forward
void ck_ir_dump(const CKIRGraph * graph, FILE * out)

Dump a human-readable view of the IR to the given stream.

ck_ir_free Forward
void ck_ir_free(CKIRGraph * graph)

Free any heap-allocated memory owned by the graph.

ck_ir_parse_json Forward
int ck_ir_parse_json(const char * path, CKIRGraph * graph)

Parse a JSON IR map file (as produced by ck_ir_serialize_json) back into a CKIRGraph. This enables a two-stage pipeline:

ck_ir_serialize_json Forward
int ck_ir_serialize_json(const CKIRGraph * graph, const char * path)

Serialize a CKIRGraph to a simple JSON IR map file.

ck_ir_v2_align_up_bytes Forward
size_t ck_ir_v2_align_up_bytes(size_t n, size_t align)

ck_ir_v2_align_up_elems Forward
size_t ck_ir_v2_align_up_elems(size_t elems, size_t elem_bytes, size_t align_bytes)

ck_ir_v2_apply_meta Forward
int ck_ir_v2_apply_meta(const char * path, CKIRV2Graph * graph)

ck_ir_v2_apply_weight_dtypes Forward
int ck_ir_v2_apply_weight_dtypes(const char * json, const char * end, CKIRV2Graph * graph)

ck_ir_v2_build_decoder Forward
int ck_ir_v2_build_decoder(const CKModelConfig * cfg, CKIRV2Graph * graph)

ck_ir_v2_build_decoder_backward Backward
int ck_ir_v2_build_decoder_backward(const CKIRV2Graph * forward, CKIRV2Graph * backward)

Backward pass / gradient computation

ck_ir_v2_copy_buffer_spec Forward
int ck_ir_v2_copy_buffer_spec(const CKBufferSpec * spec, CKIRV2Buffer * out)

ck_ir_v2_copy_shape Forward
void ck_ir_v2_copy_shape(CKDimToken * dst, const CKDimToken * src)

ck_ir_v2_dim_kind_from_name Forward
CKDimKind ck_ir_v2_dim_kind_from_name(const char * name)

ck_ir_v2_dim_name Forward
const char * ck_ir_v2_dim_name(CKDimKind dim)

ck_ir_v2_dtype_name Forward
const char * ck_ir_v2_dtype_name(CKDataType dtype)

ck_ir_v2_emit_dimensions Forward
void ck_ir_v2_emit_dimensions(FILE * out, const CKModelConfig * cfg, const CKIRV2AlignInfo * align, int tokens_override)

ck_ir_v2_emit_memory_plan Forward
void ck_ir_v2_emit_memory_plan(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * plan)

ck_ir_v2_emit_resolved_shape Forward
void ck_ir_v2_emit_resolved_shape(FILE * out, const CKModelConfig * cfg, const CKIRV2AlignInfo * align, const CKDimToken * shape, int tokens_override)

ck_ir_v2_emit_shape Forward
int ck_ir_v2_emit_shape(FILE * out, const CKDimToken * shape)

ck_ir_v2_find_array_end Forward
const char * ck_ir_v2_find_array_end(const char * open, const char * end)

ck_ir_v2_find_buffer_index Forward
int ck_ir_v2_find_buffer_index(const CKIRV2Graph * graph, const char * name)

ck_ir_v2_find_buffer_spec Forward
const CKBufferSpec * ck_ir_v2_find_buffer_spec(const char * name)

ck_ir_v2_find_kernel_spec Forward
const CKKernelSpec * ck_ir_v2_find_kernel_spec(const char * name)

ck_ir_v2_find_key Forward
const char * ck_ir_v2_find_key(const char * json, const char * key, const char * end)

ck_ir_v2_free Forward
void ck_ir_v2_free(CKIRV2Graph * graph)

ck_ir_v2_free_buffer Forward
void ck_ir_v2_free_buffer(CKIRV2Buffer * buf)

ck_ir_v2_free_node Forward
void ck_ir_v2_free_node(CKIRV2Node * node)

ck_ir_v2_lower_copy_buffers Forward
int ck_ir_v2_lower_copy_buffers(const CKIRV2Graph * input, CKIRV2Graph * output)

ck_ir_v2_lower_copy_nodes Forward
int ck_ir_v2_lower_copy_nodes(const CKIRV2Graph * input, CKIRV2LowerMode mode, CKIRV2Graph * output)

ck_ir_v2_lower_emit_json Forward
int ck_ir_v2_lower_emit_json(const CKIRV2Graph * input, CKIRV2LowerMode mode, const char * path)

ck_ir_v2_lower_graph Forward
int ck_ir_v2_lower_graph(const CKIRV2Graph * input, CKIRV2LowerMode mode, CKIRV2Graph * output, CKMemPlan * plan)

ck_ir_v2_lower_mode_from_string Forward
int ck_ir_v2_lower_mode_from_string(const char * name, CKIRV2LowerMode * out_mode)

ck_ir_v2_lower_mode_name Forward
const char * ck_ir_v2_lower_mode_name(CKIRV2LowerMode mode)

ck_ir_v2_lower_node_enabled Forward
int ck_ir_v2_lower_node_enabled(const CKIRV2Node * node, CKIRV2LowerMode mode)

ck_ir_v2_lower_strdup Forward
char * ck_ir_v2_lower_strdup(const char * s)

ck_ir_v2_mem_arena_name Forward
const char * ck_ir_v2_mem_arena_name(CKMemArenaKind arena)

ck_ir_v2_next_object Forward
const char * ck_ir_v2_next_object(const char * cur, const char * end, const char ** obj_start, const char ** obj_end)

ck_ir_v2_parse_bindings Forward
int ck_ir_v2_parse_bindings(const char * obj_start, const char * obj_end, CKIRV2Graph * graph, CKIRV2Node * node)

ck_ir_v2_parse_bool Forward
int ck_ir_v2_parse_bool(const char * json, const char * key, const char * end, int * out_val)

ck_ir_v2_parse_buffers Forward
int ck_ir_v2_parse_buffers(const char * json, const char * end, CKIRV2Graph * graph)

ck_ir_v2_parse_dim_kind Forward
CKDimKind ck_ir_v2_parse_dim_kind(const char * obj_start, const char * obj_end)

ck_ir_v2_parse_dtype Forward
CKDataType ck_ir_v2_parse_dtype(const char * s)

ck_ir_v2_parse_float Forward
int ck_ir_v2_parse_float(const char * json, const char * key, const char * end, float * out_val)

ck_ir_v2_parse_int Forward
int ck_ir_v2_parse_int(const char * json, const char * key, const char * end, int * out_val)

ck_ir_v2_parse_json Forward
int ck_ir_v2_parse_json(const char * path, CKIRV2Graph * graph)

ck_ir_v2_parse_nodes Forward
int ck_ir_v2_parse_nodes(const char * json, const char * end, CKIRV2Graph * graph)

ck_ir_v2_parse_role Forward
CKBufferRole ck_ir_v2_parse_role(const char * s)

ck_ir_v2_parse_scope Forward
CKBufferScope ck_ir_v2_parse_scope(const char * s)

ck_ir_v2_parse_shape Forward
int ck_ir_v2_parse_shape(const char * obj_start, const char * obj_end, CKDimToken * shape_out)

ck_ir_v2_parse_string Forward
int ck_ir_v2_parse_string(const char * start, const char * end, char ** out_str)

ck_ir_v2_parse_string_field Forward
int ck_ir_v2_parse_string_field(const char * json, const char * key, const char * end, char ** out_str)

ck_ir_v2_resolve_align Forward
void ck_ir_v2_resolve_align(const CKModelConfig * cfg, size_t alignment_bytes, CKIRV2AlignInfo * align)

ck_ir_v2_resolve_dim_value Forward
size_t ck_ir_v2_resolve_dim_value(const CKModelConfig * cfg, const CKIRV2AlignInfo * align, CKDimKind dim, int tokens_override)

ck_ir_v2_role_name Forward
const char * ck_ir_v2_role_name(CKBufferRole role)

ck_ir_v2_scope_name Forward
const char * ck_ir_v2_scope_name(CKBufferScope scope)

ck_ir_v2_select_kernel Forward
const char * ck_ir_v2_select_kernel(const CKKernelSpec * spec, CKDataType dtype, int backward)

ck_ir_v2_serialize_json Forward
int ck_ir_v2_serialize_json(const CKIRV2Graph * graph, const char * path)

ck_ir_v2_serialize_json_internal Forward
int ck_ir_v2_serialize_json_internal(const CKIRV2Graph * graph, const CKMemPlan * plan, const char * mode, int tokens_override, int base_context_window, const char * path)

ck_ir_v2_serialize_json_with_plan Forward
int ck_ir_v2_serialize_json_with_plan(const CKIRV2Graph * graph, const struct CKMemPlan * plan, const char * mode, int tokens_override, int base_context_window, const char * path)

ck_ir_v2_skip_string Forward
const char * ck_ir_v2_skip_string(const char * cur, const char * end)

ck_ir_v2_skip_ws Forward
const char * ck_ir_v2_skip_ws(const char * cur, const char * end)

ck_ir_v2_strdup Forward
char * ck_ir_v2_strdup(const char * s)

ck_ir_validate_supported Forward
int ck_ir_validate_supported(const CKIRGraph * graph)

ck_layer_debug_enabled Forward
int ck_layer_debug_enabled(void)

ck_layout_head_to_token_f32 Forward
void ck_layout_head_to_token_f32(const float * src, float * dst, int heads, int tokens, int head_dim)

ck_layout_token_to_head_f32 Forward
void ck_layout_token_to_head_f32(const float * src, float * dst, int tokens, int heads, int head_dim)

ck_llama_kv_pad_256 Forward
int ck_llama_kv_pad_256(int live_tokens, int capacity)

ck_llama_regular_dot_f16 Forward
float ck_llama_regular_dot_f16(const float * a, const float * b, int count)

ck_load_weights_manifest_v4 Forward
int ck_load_weights_manifest_v4(void * base, const char * weights_path, const char * manifest_path)

Load BUMPWGT4 weights into a v4 model buffer using a manifest map.

ck_local_fp16_to_fp32_2d Forward
void ck_local_fp16_to_fp32_2d(const uint16_t * src, float * dst, int rows, int cols, int src_stride, int dst_stride)

ck_local_fp16_to_fp32_row Forward
void ck_local_fp16_to_fp32_row(const uint16_t * src, float * dst, int n)

ck_local_fp32_to_bf16_row Forward
void ck_local_fp32_to_bf16_row(const float * src, uint16_t * dst, int n)

ck_local_fp32_to_fp16_row Forward
void ck_local_fp32_to_fp16_row(const float * src, uint16_t * dst, int n)

ck_mamba_debug_enabled Forward
int ck_mamba_debug_enabled(void)

ck_mamba_debug_finite Forward
void ck_mamba_debug_finite(const char * name, const float * x, size_t n)

ck_mem_plan_build_inference Forward
int ck_mem_plan_build_inference(const CKIRV2Graph * graph, CKMemPlan * plan, size_t alignment_bytes)

ck_mem_plan_build_inference_with_tokens Forward
int ck_mem_plan_build_inference_with_tokens(const CKIRV2Graph * graph, CKMemPlan * plan, size_t alignment_bytes, int tokens_override)

ck_mem_plan_build_training Forward
int ck_mem_plan_build_training(const CKIRV2Graph * graph, CKMemPlan * plan, size_t alignment_bytes)

ck_mem_plan_build_training_with_tokens Forward
int ck_mem_plan_build_training_with_tokens(const CKIRV2Graph * graph, CKMemPlan * plan, size_t alignment_bytes, int tokens_override)

ck_mem_plan_free Forward
void ck_mem_plan_free(CKMemPlan * plan)

ck_memcpy_parallel_dispatch Forward
void * ck_memcpy_parallel_dispatch(void * dst, const void * src, size_t size)

ck_memory_allocate Forward
int ck_memory_allocate(CKModel * model, int use_hugepages)

Allocate the planned memory.

ck_memory_free Forward
void ck_memory_free(CKModel * model)

Free the model memory.

ck_memory_plan Forward
size_t ck_memory_plan(const CKSectionConfig * sections, int num_sections, int mode, uint32_t fusion_flags, CKModel * out_model)

Plan memory layout for a model.

ck_metrics_cleanup Forward
void ck_metrics_cleanup(void)

Cleanup and free resources

ck_metrics_create_context Forward
CKMetricsContext * ck_metrics_create_context(void)

ck_metrics_ctx_end Forward
void ck_metrics_ctx_end(CKMetricsContext * ctx, const char * status)

ck_metrics_ctx_init Forward
bool ck_metrics_ctx_init(CKMetricsContext * ctx, const char * run_id, const char * endpoint, CKMetricsMode mode)

ck_metrics_ctx_log_f Forward
void ck_metrics_ctx_log_f(CKMetricsContext * ctx, const char * name, double value)

ck_metrics_ctx_log_i Forward
void ck_metrics_ctx_log_i(CKMetricsContext * ctx, const char * name, int64_t value)

ck_metrics_ctx_step Forward
void ck_metrics_ctx_step(CKMetricsContext * ctx, int64_t step)

ck_metrics_destroy_context Forward
void ck_metrics_destroy_context(CKMetricsContext * ctx)

ck_metrics_end Forward
void ck_metrics_end(const char * status)

End the training run

ck_metrics_generate_run_id Forward
void ck_metrics_generate_run_id(char * buffer, size_t size)

Generate a unique run ID based on current timestamp

ck_metrics_get_memory_mb Forward
int64_t ck_metrics_get_memory_mb(void)

Get current memory usage in MB (platform-specific)

ck_metrics_init Forward
bool ck_metrics_init(const char * run_id, const char * endpoint, CKMetricsMode mode)

Initialize metrics logging

ck_metrics_init_full Forward
bool ck_metrics_init_full(const char * run_id, const char * endpoint, CKMetricsMode mode, const char * model, const char * dataset, int batch_size, double lr, int max_steps)

Initialize with full configuration

ck_metrics_log_f Forward
void ck_metrics_log_f(const char * name, double value)

Log a float metric (e.g., loss, learning rate)

ck_metrics_log_i Forward
void ck_metrics_log_i(const char * name, int64_t value)

Log an integer metric (e.g., step, tokens_per_sec)

ck_metrics_log_s Forward
void ck_metrics_log_s(const char * name, const char * value)

Log a string metric (e.g., phase, status)

ck_metrics_step Forward
void ck_metrics_step(int64_t step)

Flush metrics for the current step and advance to next step Call this at the end of each training step

ck_metrics_timestamp Forward
double ck_metrics_timestamp(void)

Get current timestamp in seconds with microsecond precision

ck_min Forward
int ck_min(int a, int b)

ck_min_i Forward
int ck_min_i(int a, int b)

ck_model_allocate Forward
int ck_model_allocate(CKModel * model, int hugepage_mode)

Allocate the planned memory. hugepage_mode 0=normal, 1=2MB hugepages, 2=1GB hugepages 0 on success, -1 on failure

ck_model_config_from_hf_json Forward
int ck_model_config_from_hf_json(const char * path, CKModelConfig * cfg)

Parse a HuggingFace-style config.json into CKModelConfig.

ck_model_create Forward
void * ck_model_create(void)

Create and allocate model memory. Returns opaque model pointer, or NULL on failure.

ck_model_decode Forward
void ck_model_decode(void * model, const int * token, int token_index)

Decode single token at position token_index. Used for autoregressive generation.

ck_model_forward Forward
void ck_model_forward(void * model, const int * tokens, int num_tokens)

Forward pass (prefill) - process multiple tokens. Used for initial prompt processing.

ck_model_free Forward
void ck_model_free(void * model)

Free model memory.

ck_model_get_base Forward
void * ck_model_get_base(void * model)

Get model base pointer (for weight loading).

ck_model_get_config Forward
const CKModelConfig * ck_model_get_config(void)

Get model configuration (dimensions, sizes, etc.) This is available before allocation.

ck_model_get_logits Forward
float * ck_model_get_logits(void * model)

Get pointer to output logits buffer. Size is vocab_size floats.

ck_model_get_total_bytes Forward
size_t ck_model_get_total_bytes(void * model)

Get total model size in bytes.

ck_model_load_weights Forward
int ck_model_load_weights(void * model, const char * bump_path)

Load weights from BUMP file into model. Returns 0 on success, -1 on failure.

ck_model_load_weights_flat Forward
int ck_model_load_weights_flat(TransformerModel * m, const char * path)

Load weights from a single flat binary file into model->memory_base.

ck_model_plan Forward
size_t ck_model_plan(CKModel * model, const CKSectionConfig * configs, int num_sections, int training_enabled, uint32_t fusion_flags)

Plan memory layout for complete model. Returns total bytes needed.

ck_model_verify_canaries Forward
int ck_model_verify_canaries(void * model)

Verify memory canaries (debug). Returns number of corrupted canaries (0 = OK).

ck_moe_align64 Forward
size_t ck_moe_align64(size_t value)

ck_moe_bf16_round Forward
float ck_moe_bf16_round(float x)

ck_moe_bucket_expert_for_position Forward
int ck_moe_bucket_expert_for_position(const int * offsets, int n_experts, int position)

ck_moe_bucket_layout Forward
int ck_moe_bucket_layout(int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, ck_moe_q4k_q5k_bucket_layout_t * layout)

ck_moe_debug_enabled Forward
int ck_moe_debug_enabled(void)

ck_moe_debug_finite Forward
void ck_moe_debug_finite(const char * name, const float * x, size_t n)

ck_moe_down_idx Forward
size_t ck_moe_down_idx(int e, int h, int i, int hidden_dim, int intermediate_dim)

ck_moe_dsilu_f32 Forward
float ck_moe_dsilu_f32(float x)

ck_moe_llama_weighted_accumulate Forward
void ck_moe_llama_weighted_accumulate(float * output, const float * expert_output, float route_weight, int n)

ck_moe_q4k_llama_projection Forward
void ck_moe_q4k_llama_projection(float * output, const void * weights, const void * input_q8, int output_dim, int input_dim, void * scratch)

ck_moe_q4k_llama_projection_scratch_bytes Forward
size_t ck_moe_q4k_llama_projection_scratch_bytes(int output_dim, int input_dim)

ck_moe_q4k_mixed_parallel_work Forward
void ck_moe_q4k_mixed_parallel_work(int ith, int nth, void * opaque)

ck_moe_q4k_mixed_parallel_workspace Forward
int ck_moe_q4k_mixed_parallel_workspace(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes, ck_moe_expert_workspace_fn serial_fn, ck_moe_down_kind_t down_kind)

ck_moe_q4k_mixed_route_parallel Forward
int ck_moe_q4k_mixed_route_parallel(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes, size_t workspace_stride, ck_threadpool_t * pool, ck_moe_down_kind_t down_kind)

ck_moe_q4k_mixed_route_work Forward
void ck_moe_q4k_mixed_route_work(int ith, int nth, void * opaque)

ck_moe_q4k_q5k_bucket_work Forward
void ck_moe_q4k_q5k_bucket_work(int ith, int nth, void * opaque)

ck_moe_q4k_q5k_parallel_work Forward
void ck_moe_q4k_q5k_parallel_work(int ith, int nth, void * opaque)

ck_moe_q4k_q5k_quantize_work Forward
void ck_moe_q4k_q5k_quantize_work(int ith, int nth, void * opaque)

ck_moe_q4k_q5k_route_parallel Forward
int ck_moe_q4k_q5k_route_parallel(const float * hidden, const int * indices, const float * routing_weights, const void * expert_gate, const void * expert_up, const void * expert_down, float * output, int hidden_dim, int intermediate_dim, int n_experts, int top_k, void * workspace, size_t workspace_bytes, size_t workspace_stride, ck_threadpool_t * pool)

ck_moe_q4k_q5k_route_work Forward
void ck_moe_q4k_q5k_route_work(int ith, int nth, void * opaque)

ck_moe_shared_gated_parallel_workspace Forward
int ck_moe_shared_gated_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes, size_t stride, ck_moe_shared_gated_workspace_fn serial_fn)

ck_moe_shared_q4k_gated_workspace Forward
int ck_moe_shared_q4k_gated_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, const float * shared_gate_input, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes, void(*)(float *, const void *, const void *, int, int) down_projection)

ck_moe_shared_q4k_parallel_work Forward
void ck_moe_shared_q4k_parallel_work(int ith, int nth, void * opaque)

ck_moe_shared_q4k_parallel_workspace Forward
int ck_moe_shared_q4k_parallel_workspace(const float * hidden, const float * routed, const void * shared_gate, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim, void * workspace, size_t workspace_bytes, ck_moe_shared_workspace_fn serial_fn)

ck_moe_shared_q8_0_parallel_work Forward
void ck_moe_shared_q8_0_parallel_work(int ith, int nth, void * opaque)

ck_moe_silu_f32 Forward
float ck_moe_silu_f32(float x)

ck_moe_size_add Forward
int ck_moe_size_add(size_t a, size_t b, size_t * result)

ck_moe_size_mul Forward
int ck_moe_size_mul(size_t a, size_t b, size_t * result)

ck_moe_up_idx Forward
size_t ck_moe_up_idx(int e, int i, int h, int intermediate_dim, int hidden_dim)

ck_multimodal_prefix_insert_f32 Forward
int ck_multimodal_prefix_insert_f32(const float * source_rows, int32_t * token_ids, float * decoder_rows, int row_count, int source_row_stride, int decoder_row_stride, int copy_dim, int start_row, int decoder_capacity)

ck_murmurhash3 Forward
uint32_t ck_murmurhash3(const char * key, uint32_t len, uint32_t seed)

MurmurHash3-32bit hash function (original HPC_Embeddings version).

ck_murmurhash3_128 Forward
void ck_murmurhash3_128(const void * key, size_t len, uint32_t seed, uint64_t * out1, uint64_t * out2)

MurmurHash3-128bit hash function (produces two 64-bit values).

ck_murmurhash3_32 Forward
uint32_t ck_murmurhash3_32(const void * key, size_t len, uint32_t seed)

MurmurHash3-32bit hash function (alternative name).

ck_murmurhash3_str Forward
uint32_t ck_murmurhash3_str(const char * key, uint32_t seed)

ck_murmurhash3_strn Forward
uint32_t ck_murmurhash3_strn(const char * str, size_t len, uint32_t seed)

MurmurHash3-32bit for string with length.

ck_nearest_int Forward
int ck_nearest_int(float fval)

ck_nearest_int_fused Forward
int ck_nearest_int_fused(float fval)

ck_nearest_int_q8_0 Forward
int ck_nearest_int_q8_0(float fval)

ck_nearest_int_q8_0_ref Forward
int ck_nearest_int_q8_0_ref(float fval)

ck_nvfp4_align64 Forward
size_t ck_nvfp4_align64(size_t value)

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

ck_op_name Forward
const char * ck_op_name(CKOpType op)

ck_op_supported Forward
int ck_op_supported(CKOpType op)

ck_opt_pick_active_threads Forward
int ck_opt_pick_active_threads(int nth, size_t work_items, size_t min_chunk)

ck_parallel_for_worker Forward
void ck_parallel_for_worker(int ith, int nth, void * opaque)

ck_parse_env_int Forward
int ck_parse_env_int(const char * name)

ck_patch_projection_bf16_native_work Forward
void ck_patch_projection_bf16_native_work(int ith, int nth, void * opaque)

ck_plan_step_enabled Forward
int ck_plan_step_enabled(const CKPlanStep * step, const CKIRGraph * cfg)

ck_pool_alloc Forward
void * ck_pool_alloc(CKMemPool * pool, size_t size)

ck_pool_free Forward
void ck_pool_free(CKMemPool * pool)

ck_pool_init Forward
void ck_pool_init(CKMemPool * pool)

ck_pool_strdup Forward
char * ck_pool_strdup(CKMemPool * pool, const char * s, int len)

ck_prefill_forward Forward
int ck_prefill_forward(const void * weights, const int32_t * tokens, int n_tokens, float * hidden_out, void * kv_cache, int kv_pos)

Forward pass computation

ck_q4k_debug_q8_contract Forward
int ck_q4k_debug_q8_contract(void)

ck_q4k_packed_vnni_x16_available Forward
int ck_q4k_packed_vnni_x16_available(void)

ck_q4k_packed_vnni_x8_available Forward
int ck_q4k_packed_vnni_x8_available(void)

ck_q4k_packed_vnni_x8_compact_order_available Forward
int ck_q4k_packed_vnni_x8_compact_order_available(void)

ck_q4k_q8k_force_ref Forward
int ck_q4k_q8k_force_ref(void)

ck_q4k_silu_f32 Forward
float ck_q4k_silu_f32(float x)

ck_q4k_x16_chunk4_enabled Forward
int ck_q4k_x16_chunk4_enabled(void)

ck_q5_k_prepare_weight Forward
void ck_q5_k_prepare_weight(const void * src, void * dst, int N, int K)

ck_q5_k_prepared_block_size Forward
size_t ck_q5_k_prepared_block_size(void)

ck_q5k_debug_fp32_fallback Forward
int ck_q5k_debug_fp32_fallback(void)

ck_q5k_debug_generic_dot Forward
int ck_q5k_debug_generic_dot(void)

ck_q6_k_prepare_weight Forward
void ck_q6_k_prepare_weight(const void * src, void * dst, int N, int K)

ck_q6_k_prepared_block_size Forward
size_t ck_q6_k_prepared_block_size(void)

ck_q6_k_prepared_provider_name Forward
const char * ck_q6_k_prepared_provider_name(void)

ck_q6_k_q8_k_provider_name Forward
const char * ck_q6_k_q8_k_provider_name(void)

ck_q6k_q8k_force_ref Forward
int ck_q6k_q8k_force_ref(void)

ck_q80_contract_cached_input_enabled
int ck_q80_contract_cached_input_enabled(void)

ck_q80_contract_dump_enabled Forward
int ck_q80_contract_dump_enabled(void)

ck_q80_contract_dump_tensor Forward
void ck_q80_contract_dump_tensor(const char * name, int layer_id, const float * data, size_t elem_count)

ck_q80_resolve_ggml_build_forward_expand Forward
ck_q80_ggml_build_forward_expand_fn ck_q80_resolve_ggml_build_forward_expand(void)

Forward pass computation

ck_q80_resolve_ggml_cpu_init Forward
ck_q80_ggml_cpu_init_fn ck_q80_resolve_ggml_cpu_init(void)

ck_q80_resolve_ggml_free Forward
ck_q80_ggml_free_fn ck_q80_resolve_ggml_free(void)

ck_q80_resolve_ggml_get_data Forward
ck_q80_ggml_get_data_fn ck_q80_resolve_ggml_get_data(void)

ck_q80_resolve_ggml_get_data_f32 Forward
ck_q80_ggml_get_data_f32_fn ck_q80_resolve_ggml_get_data_f32(void)

ck_q80_resolve_ggml_graph_compute_with_ctx Forward
ck_q80_ggml_graph_compute_with_ctx_fn ck_q80_resolve_ggml_graph_compute_with_ctx(void)

ck_q80_resolve_ggml_init Forward
ck_q80_ggml_init_fn ck_q80_resolve_ggml_init(void)

ck_q80_resolve_ggml_mul_mat Forward
ck_q80_ggml_mul_mat_fn ck_q80_resolve_ggml_mul_mat(void)

ck_q80_resolve_ggml_nbytes Forward
ck_q80_ggml_nbytes_fn ck_q80_resolve_ggml_nbytes(void)

ck_q80_resolve_ggml_new_graph Forward
ck_q80_ggml_new_graph_fn ck_q80_resolve_ggml_new_graph(void)

ck_q80_resolve_ggml_new_tensor_2d Forward
ck_q80_ggml_new_tensor_2d_fn ck_q80_resolve_ggml_new_tensor_2d(void)

ck_q8_0_debug_ref Forward
int ck_q8_0_debug_ref(void)

ck_q8_0_fp32_m4n4_enabled Forward
int ck_q8_0_fp32_m4n4_enabled(void)

ck_q8_0_outproj_enabled Forward
int ck_q8_0_outproj_enabled(void)

ck_q8_0_q8_0_debug_ref Forward
int ck_q8_0_q8_0_debug_ref(void)

ck_q8k_activations_enabled Forward
int ck_q8k_activations_enabled(void)

ck_qkv_project_head_major Forward
void ck_qkv_project_head_major(const float * input, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, float * q, float * k, float * v, int tokens, int kv_stride_tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_backward Backward
void ck_qkv_project_head_major_backward(const float * d_q, const float * d_k, const float * d_v, const float * input, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, float * d_input, float * d_wq, float * d_bq, float * d_wk, float * d_bk, float * d_wv, float * d_bv, float * scratch, int tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim, int num_threads)

Backward pass / gradient computation

ck_qkv_project_head_major_q4_k Forward
void ck_qkv_project_head_major_q4_k(const float * input, const void * wq, const float * bq, const void * wk, const float * bk, const void * wv, const float * bv, float * q, float * k, float * v, int tokens, int kv_stride_tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_q4_k_q8_k Forward
void ck_qkv_project_head_major_q4_k_q8_k(const float * input, const void * wq, const float * bq, const void * wk, const float * bk, const void * wv, const float * bv, float * q, float * k, float * v, int tokens, int kv_stride_tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_quant Forward
void ck_qkv_project_head_major_quant(const float * input, const void * wq, const float * bq, CKDataType wq_dtype, const void * wk, const float * bk, CKDataType wk_dtype, const void * wv, const float * bv, CKDataType wv_dtype, float * q, float * k, float * v, int tokens, int kv_stride_tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_ref Forward
void ck_qkv_project_head_major_ref(const float * input, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, float * q, float * k, float * v, int tokens, int kv_stride_tokens, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_token Forward
void ck_qkv_project_head_major_token(const float * input_row, const float * wq, const float * bq, const float * wk, const float * bk, const float * wv, const float * bv, float * q_token, float * k_token, float * v_token, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_token_q4_k Forward
void ck_qkv_project_head_major_token_q4_k(const float * input_row, const void * wq, const float * bq, const void * wk, const float * bk, const void * wv, const float * bv, float * q_token, float * k_token, float * v_token, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_token_q4_k_q8_k Forward
void ck_qkv_project_head_major_token_q4_k_q8_k(const block_q8_K * input_q8, const void * wq, const float * bq, const void * wk, const float * bk, const void * wv, const float * bv, float * q_token, float * k_token, float * v_token, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_qkv_project_head_major_token_quant Forward
void ck_qkv_project_head_major_token_quant(const float * input_row, const void * wq, const float * bq, CKDataType wq_dtype, const void * wk, const float * bk, CKDataType wk_dtype, const void * wv, const float * bv, CKDataType wv_dtype, float * q_token, float * k_token, float * v_token, int aligned_embed_dim, int num_heads, int num_kv_heads, int aligned_head_dim)

ck_quant_block_size Forward
size_t ck_quant_block_size(int type)

Get the block size (number of weights per block) for a quant type.

ck_quant_row_size Forward
size_t ck_quant_row_size(int type, int64_t n_elements)

Calculate total bytes needed for n_elements with given quant type.

ck_quant_type_size Forward
size_t ck_quant_type_size(int type)

Get the byte size per block for a quant type.

ck_residual_add_backward Backward
void ck_residual_add_backward(const float * d_out, float * d_a, float * d_b, int tokens, int aligned_embed_dim)

Backward pass / gradient computation

ck_residual_add_token_major Forward
void ck_residual_add_token_major(const float * a, const float * b, float * out, int tokens, int aligned_embed_dim)

ck_residual_add_token_major_bf16_storage Forward
void ck_residual_add_token_major_bf16_storage(const float * a, const float * b, float * out, int tokens, int aligned_embed_dim)

ck_residual_add_token_major_parallel_dispatch Forward
void ck_residual_add_token_major_parallel_dispatch(const float * a, const float * b, float * out, int tokens, int aligned_embed_dim)

ck_resolve_ggml_backend_cpu_set_n_threads
ck_ggml_backend_cpu_set_n_threads_fn ck_resolve_ggml_backend_cpu_set_n_threads(void)

ck_resolve_ggml_backend_free
ck_ggml_backend_free_fn ck_resolve_ggml_backend_free(void)

ck_resolve_ggml_backend_get_default_buffer_type
ck_ggml_backend_get_default_buffer_type_fn ck_resolve_ggml_backend_get_default_buffer_type(void)

ck_resolve_ggml_backend_init_by_type
ck_ggml_backend_init_by_type_fn ck_resolve_ggml_backend_init_by_type(void)

ck_resolve_ggml_backend_sched_alloc_graph
ck_ggml_backend_sched_alloc_graph_fn ck_resolve_ggml_backend_sched_alloc_graph(void)

ck_resolve_ggml_backend_sched_free
ck_ggml_backend_sched_free_fn ck_resolve_ggml_backend_sched_free(void)

ck_resolve_ggml_backend_sched_graph_compute
ck_ggml_backend_sched_graph_compute_fn ck_resolve_ggml_backend_sched_graph_compute(void)

ck_resolve_ggml_backend_sched_new
ck_ggml_backend_sched_new_fn ck_resolve_ggml_backend_sched_new(void)

ck_resolve_ggml_backend_sched_reset
ck_ggml_backend_sched_reset_fn ck_resolve_ggml_backend_sched_reset(void)

ck_resolve_ggml_backend_tensor_get
ck_ggml_backend_tensor_get_fn ck_resolve_ggml_backend_tensor_get(void)

ck_resolve_ggml_backend_tensor_set
ck_ggml_backend_tensor_set_fn ck_resolve_ggml_backend_tensor_set(void)

ck_resolve_ggml_build_forward_expand Forward
ck_ggml_build_forward_expand_fn ck_resolve_ggml_build_forward_expand(void)

Forward pass computation

ck_resolve_ggml_cont Forward
ck_ggml_cont_fn ck_resolve_ggml_cont(void)

ck_resolve_ggml_cont_2d Forward
ck_ggml_cont_2d_fn ck_resolve_ggml_cont_2d(void)

ck_resolve_ggml_cpu_init Forward
ck_ggml_cpu_init_fn ck_resolve_ggml_cpu_init(void)

ck_resolve_ggml_free Forward
ck_ggml_free_fn ck_resolve_ggml_free(void)

ck_resolve_ggml_get_data Forward
ck_ggml_get_data_fn ck_resolve_ggml_get_data(void)

ck_resolve_ggml_graph_compute_with_ctx Forward
ck_ggml_graph_compute_with_ctx_fn ck_resolve_ggml_graph_compute_with_ctx(void)

ck_resolve_ggml_init Forward
ck_ggml_init_fn ck_resolve_ggml_init(void)

ck_resolve_ggml_mul_mat_graph Forward
ck_ggml_mul_mat_graph_fn ck_resolve_ggml_mul_mat_graph(void)

ck_resolve_ggml_new_graph Forward
ck_ggml_new_graph_fn ck_resolve_ggml_new_graph(void)

ck_resolve_ggml_new_tensor_1d Forward
ck_ggml_new_tensor_1d_fn ck_resolve_ggml_new_tensor_1d(void)

ck_resolve_ggml_new_tensor_2d Forward
ck_ggml_new_tensor_2d_fn ck_resolve_ggml_new_tensor_2d(void)

ck_resolve_ggml_permute Forward
ck_ggml_permute_fn ck_resolve_ggml_permute(void)

ck_resolve_ggml_set_input Forward
ck_ggml_set_input_fn ck_resolve_ggml_set_input(void)

ck_resolve_ggml_soft_max_ext Forward
ck_ggml_soft_max_ext_fn ck_resolve_ggml_soft_max_ext(void)

ck_resolve_ggml_view_3d Forward
ck_ggml_view_3d_fn ck_resolve_ggml_view_3d(void)

ck_round_fp16_buffer Forward
void ck_round_fp16_buffer(const float * src, float * dst, size_t count)

ck_round_fp16_scalar Forward
float ck_round_fp16_scalar(float x)

ck_round_nearest Forward
int ck_round_nearest(float v)

Round to nearest int, half away from zero (matches quantize_row_q8_0)

ck_sample_top_p_v8 Forward
int ck_sample_top_p_v8(float * logits, int vocab_size, float temperature, float top_p, float random_value)

ck_scale_f32_to_f16 Forward
void ck_scale_f32_to_f16(const float * src, float scale, uint16_t * dst, int n)

Scale FP32 array and store as FP16: dst = scale * src.

ck_scale_parallel_work Forward
void ck_scale_parallel_work(int ith, int nth, void * argp)

ck_section_config_init Forward
void ck_section_config_init(CKSectionConfig * config, size_t simd_align)

Initialize section config with computed alignments.

ck_section_plan Forward
size_t ck_section_plan(CKSection * section, const CKSectionConfig * config, int training_enabled, size_t base_offset)

Plan memory layout for a single section. Returns bytes needed for this section.

ck_session_v8_cancel Forward
void ck_session_v8_cancel(CKSessionV8 * session)

ck_session_v8_close Forward
void ck_session_v8_close(CKSessionV8 * session)

ck_session_v8_decode Forward
int ck_session_v8_decode(CKSessionV8 * session, const int32_t * tokens, int32_t token_count, char * output, int32_t capacity)

ck_session_v8_encode Forward
int ck_session_v8_encode(CKSessionV8 * session, const char * text, int32_t * output, int32_t capacity)

ck_session_v8_format_chat Forward
int ck_session_v8_format_chat(CKSessionV8 * session, const char * system_text, const char * user_text, char * output, int32_t capacity)

ck_session_v8_generate Forward
int ck_session_v8_generate(CKSessionV8 * session, const CKSessionGenerateRequestV8 * request, ck_session_token_callback_v8 callback, void * user_data, CKSessionGenerateResultV8 * result)

ck_session_v8_get_abi_version Forward
uint32_t ck_session_v8_get_abi_version(void)

ck_session_v8_get_model_descriptor Forward
int ck_session_v8_get_model_descriptor(const CKSessionV8 * session, CKModelRuntimeDescriptorV8 * descriptor, size_t descriptor_size)

ck_session_v8_last_error Forward
const char * ck_session_v8_last_error(const CKSessionV8 * session)

ck_session_v8_open Forward
int ck_session_v8_open(const CKSessionConfigV8 * config, CKSessionV8 ** session_out)

ck_session_v8_reset Forward
int ck_session_v8_reset(CKSessionV8 * session)

ck_set_num_threads Forward
void ck_set_num_threads(int num_threads)

ck_set_strict_parity Forward
void ck_set_strict_parity(int enabled)

ck_speed_profile_qwen3vl_ocr_fast Forward
int ck_speed_profile_qwen3vl_ocr_fast(void)

ck_ssm_conv1d_llama_channel_range Forward
void ck_ssm_conv1d_llama_channel_range(int begin, int end, void * opaque)

ck_ssm_conv1d_llama_fma_channel_range Forward
void ck_ssm_conv1d_llama_fma_channel_range(int begin, int end, void * opaque)

ck_strict_mtmd_clip_encode_planar_f32 Forward
int ck_strict_mtmd_clip_encode_planar_f32(const float * planar, int channels, int height, int width, float * out, size_t out_elems)

ck_strict_parity_enabled Forward
int ck_strict_parity_enabled(void)

ck_sum_sq_multi_parallel_work Forward
void ck_sum_sq_multi_parallel_work(int ith, int nth, void * argp)

ck_sum_sq_parallel_work Forward
void ck_sum_sq_parallel_work(int ith, int nth, void * argp)

ck_test_dequant_q4_0 Forward
void ck_test_dequant_q4_0(const void * src, float * dst, int n)

Dequantize Q4_0 data to FP32.

ck_test_dequant_q4_k Forward
void ck_test_dequant_q4_k(const void * src, float * dst, int n)

Dequantize Q4_K data to FP32.

ck_test_dequant_q6_k Forward
void ck_test_dequant_q6_k(const void * src, float * dst, int n)

Dequantize Q6_K data to FP32.

ck_test_gated_deltanet_autoregressive Forward
void ck_test_gated_deltanet_autoregressive(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)

Gated DeltaNet autoregressive update.

ck_test_gemv_q4_k Forward
void ck_test_gemv_q4_k(const void * weight_q4k, const float * input_f32, float * output, int cols)

Q4_K GEMV - dot product of quantized weights and FP32 input.

ck_test_gemv_q5_0 Forward
void ck_test_gemv_q5_0(const void * weight_q5_0, const float * input_f32, float * output, int rows, int cols)

Q5_0 GEMV - matrix-vector multiply with Q5_0 weights.

ck_test_gemv_q5_0_q8_0 Forward
void ck_test_gemv_q5_0_q8_0(const void * weight_q5_0, const float * input_f32, float * output, int rows, int cols)

Q5_0 x Q8_0 quantized GEMV - matches llama.cpp's approach.

ck_test_gemv_q6_k Forward
void ck_test_gemv_q6_k(const void * weight_q6k, const float * input_f32, float * output, int cols)

Q6_K GEMV.

ck_test_gemv_q8_0 Forward
void ck_test_gemv_q8_0(const void * weight_q8_0, const float * input_f32, float * output, int rows, int cols)

Q8_0 GEMV - matrix-vector multiply with Q8_0 weights.

ck_test_gemv_q8_0_q8_0 Forward
void ck_test_gemv_q8_0_q8_0(const void * weight_q8_0, const float * input_f32, float * output, int rows, int cols)

Q8_0 x Q8_0 quantized GEMV - matches llama.cpp's approach.

ck_test_quantize_q8_k Forward
void ck_test_quantize_q8_k(const float * src, void * dst, int n)

Quantize FP32 to Q8_K (for activations)

ck_test_recurrent_conv_state_update Forward
void ck_test_recurrent_conv_state_update(const float * state_in, const float * q, const float * k, const float * v, float * conv_x, float * state_out, int history_len, int num_seqs, int num_tokens, int q_dim, int k_dim, int v_dim)

Build the recurrent convolution input history window.

ck_test_recurrent_dt_gate Forward
void ck_test_recurrent_dt_gate(const float * alpha, const float * dt_bias, const float * a, float * gate, int rows, int dim)

Transform recurrent alpha rows into the DeltaNet gate.

ck_test_recurrent_norm_gate Forward
void ck_test_recurrent_norm_gate(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

Per-head RMSNorm followed by SiLU(z) gating for recurrent outputs.

ck_test_recurrent_qk_l2_norm Forward
void ck_test_recurrent_qk_l2_norm(float * q, float * k, int rows, int q_dim, int k_dim, int head_dim, float eps)

Apply per-head L2 normalization to recurrent Q/K rows in-place.

ck_test_recurrent_silu Forward
void ck_test_recurrent_silu(const float * x, float * out, int rows, int dim)

Apply SiLU elementwise to recurrent rows.

ck_test_recurrent_split_conv_qkv Forward
void ck_test_recurrent_split_conv_qkv(const float * packed_qkv, float * q, float * k, float * v, int rows, int q_dim, int k_dim, int v_dim)

Split the post-convolution recurrent packed QKV rows.

ck_test_recurrent_split_qkv Forward
void ck_test_recurrent_split_qkv(const float * packed_qkv, float * q, float * k, float * v, int rows, int q_dim, int k_dim, int v_dim)

Split a packed recurrent QKV matrix into explicit Q, K, and V outputs.

ck_test_split_q_gate Forward
void ck_test_split_q_gate(const float * packed_qg, float * q, float * gate, int rows, int q_dim, int gate_dim, int group_dim)

Split a packed full-attention Q+gate matrix into Q rows and gate rows.

ck_test_ssm_conv1d Forward
void ck_test_ssm_conv1d(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

qwen3next/Qwen3.5 SSM causal depthwise convolution.

ck_test_vec_dot_q4_k_q8_k Forward
void ck_test_vec_dot_q4_k_q8_k(const void * weight_q4_k, const void * input_q8_k, float * output, int cols)

Direct Q4_K x Q8_K dot product using identical pre-quantized bytes.

ck_test_vec_dot_q5_0_q8_0 Forward
void ck_test_vec_dot_q5_0_q8_0(const void * weight_q5_0, const void * input_q8_0, float * output, int cols)

Direct Q5_0 x Q8_0 dot product test (takes pre-quantized Q8_0 input)

ck_test_vec_dot_q6_k_q8_k Forward
void ck_test_vec_dot_q6_k_q8_k(const void * weight_q6_k, const void * input_q8_k, float * output, int cols)

Direct Q6_K x Q8_K dot product using identical pre-quantized bytes.

ck_test_vec_dot_q8_0_q8_0 Forward
void ck_test_vec_dot_q8_0_q8_0(const void * weight_q8_0, const void * input_q8_0, float * output, int cols)

Direct Q8_0 x Q8_0 dot product test (takes pre-quantized Q8_0 input)

ck_threadpool_barrier Forward
void ck_threadpool_barrier(ck_threadpool_t * pool)

Barrier synchronization within a dispatched work function.

ck_threadpool_bounded_capacity Forward
int ck_threadpool_bounded_capacity(int default_threads, int logical_threads)

Compute the bounded capacity for an SMT-safe provider. The default width is preserved and at most half of the additional logical CPUs are reserved.

ck_threadpool_capacity Forward
int ck_threadpool_capacity(const ck_threadpool_t * pool)

Get the maximum worker capacity available to explicit dispatch_n calls.

ck_threadpool_create Forward
ck_threadpool_t * ck_threadpool_create(int n_threads)

Create a thread pool with n_threads total threads. Thread 0 is the calling (main) thread; n_threads-1 workers are spawned.

ck_threadpool_create_capacity Forward
ck_threadpool_t * ck_threadpool_create_capacity(int default_threads, int capacity_threads)

Create a pool whose ordinary dispatch width is smaller than its worker capacity. Exact providers may opt into the additional workers with ck_threadpool_dispatch_n(); ordinary dispatch remains at default_threads.

ck_threadpool_destroy Forward
void ck_threadpool_destroy(ck_threadpool_t * pool)

Destroy the thread pool. Signals all workers to exit and joins them. Safe to call with NULL.

ck_threadpool_dispatch Forward
void ck_threadpool_dispatch(ck_threadpool_t * pool, ck_work_fn_t fn, void * args)

Dispatch work to all threads and wait for completion.

ck_threadpool_dispatch_n Forward
void ck_threadpool_dispatch_n(ck_threadpool_t * pool, int active_threads, ck_work_fn_t fn, void * args)

Dispatch work to a subset of the pool and wait for completion.

ck_threadpool_global Forward
ck_threadpool_t * ck_threadpool_global(void)

Get or create the global thread pool. Thread-safe (uses pthread_once internally). Uses ck_get_num_threads() for the default width. In automatic mode the pool may reserve a bounded subset of SMT siblings as explicit provider capacity; ordinary dispatch remains at the default width.

ck_threadpool_global_destroy Forward
void ck_threadpool_global_destroy(void)

Destroy the global thread pool. Called during engine shutdown.

ck_threadpool_init Forward
void ck_threadpool_init(void)

Initialize the global thread pool. Called once during engine startup (e.g., from ck_model_init). Uses ck_get_num_threads() for thread count (respects CK_NUM_THREADS env).

ck_threadpool_n_threads Forward
int ck_threadpool_n_threads(const ck_threadpool_t * pool)

Get the ordinary/default dispatch width (including the main thread).

ck_threadpool_parallel_for_n Forward
void ck_threadpool_parallel_for_n(ck_threadpool_t * pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void * args)

Dynamically distribute independent ranges through the persistent pool.

ck_threadpool_pause Forward
void ck_threadpool_pause(ck_threadpool_t * pool)

Pause workers — they sleep on condvar (0% CPU). Call between batches or during interactive waiting. Workers wake on next dispatch or resume.

ck_threadpool_profile_reset Forward
void ck_threadpool_profile_reset(ck_threadpool_t * pool)

Enable profiling and reset cumulative dispatch timing counters.

ck_threadpool_profile_snapshot Forward
void ck_threadpool_profile_snapshot(const ck_threadpool_t * pool, ck_threadpool_profile_t * profile)

Snapshot cumulative dispatch timing counters without stopping workers.

ck_threadpool_resume Forward
void ck_threadpool_resume(ck_threadpool_t * pool)

Resume workers — transition from sleep to spin-wait. Call before starting a new batch of work.

ck_threadpool_shutdown Forward
void ck_threadpool_shutdown(void)

Shut down the global thread pool. Called during engine teardown. Workers are joined and freed.

ck_threadpool_thread_id Forward
int ck_threadpool_thread_id(const ck_threadpool_t * pool)

Get thread index for current thread (0 = main, -1 if not in pool)

ck_tokenizer_add_merge Forward
int ck_tokenizer_add_merge(CKTokenizer * tok, int32_t left, int32_t right, int32_t merged)

ck_tokenizer_add_special_token Forward
int ck_tokenizer_add_special_token(CKTokenizer * tok, const char * name, int32_t id)

Add special token (UNK, BOS, EOS, PAD, MASK).

ck_tokenizer_add_token Forward
int32_t ck_tokenizer_add_token(CKTokenizer * tok, const char * token, int len)

ck_tokenizer_create Forward
CKTokenizer * ck_tokenizer_create(CKTokenizerType type)

ck_tokenizer_create_bpe Forward
CKTokenizer * ck_tokenizer_create_bpe(void)

Create tokenizer with default BPE config.

ck_tokenizer_create_spm Forward
CKTokenizer * ck_tokenizer_create_spm(void)

Create tokenizer with default SPM config.

ck_tokenizer_create_wordpiece Forward
CKTokenizer * ck_tokenizer_create_wordpiece(void)

Create tokenizer with default WordPiece config.

ck_tokenizer_decode Forward
int ck_tokenizer_decode(const CKTokenizer * tok, const int32_t * ids, int num_ids, char * text, int max_len)

Decode token IDs to text.

ck_tokenizer_detect_space_prefix_style Forward
CKSpacePrefixStyle ck_tokenizer_detect_space_prefix_style(CKTokenizer * tok)

ck_tokenizer_encode Forward
int ck_tokenizer_encode(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

Encode text to token IDs using greedy longest-match.

ck_tokenizer_encode_spm_dispatch Forward
int ck_tokenizer_encode_spm_dispatch(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

ck_tokenizer_encode_spm_impl Forward
int ck_tokenizer_encode_spm_impl(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

ck_tokenizer_encode_spm_llama_impl Forward
int ck_tokenizer_encode_spm_llama_impl(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

ck_tokenizer_encode_spm_plain_segment Forward
int ck_tokenizer_encode_spm_plain_segment(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

ck_tokenizer_encode_tokens Forward
int ck_tokenizer_encode_tokens(const CKTokenizer * tok, const char * text, int text_len, const char ** out_tokens, int max_tokens)

Encode and return tokens as array of strings.

ck_tokenizer_encode_with_special Forward
int ck_tokenizer_encode_with_special(CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids, bool add_special)

Encode with special token handling.

ck_tokenizer_free Forward
void ck_tokenizer_free(CKTokenizer * tok)

ck_tokenizer_hash Forward
uint32_t ck_tokenizer_hash(const char * key, size_t len)

ck_tokenizer_hash_str Forward
uint32_t ck_tokenizer_hash_str(const char * key)

ck_tokenizer_hash_table_clear Forward
void ck_tokenizer_hash_table_clear(CKTokenizerHashTable * table, bool free_values)

Clear all entries (but keep bucket array).

ck_tokenizer_hash_table_contains Forward
bool ck_tokenizer_hash_table_contains(CKTokenizerHashTable * table, const char * key)

Check if key exists.

ck_tokenizer_hash_table_count Forward
size_t ck_tokenizer_hash_table_count(CKTokenizerHashTable * table)

Get the number of entries.

ck_tokenizer_hash_table_create Forward
CKTokenizerHashTable * ck_tokenizer_hash_table_create(size_t bucket_count)

Create a hash table.

ck_tokenizer_hash_table_delete Forward
int ck_tokenizer_hash_table_delete(CKTokenizerHashTable * table, const char * key, bool free_value)

Delete a key.

ck_tokenizer_hash_table_free Forward
void ck_tokenizer_hash_table_free(CKTokenizerHashTable * table, bool free_values)

Free a hash table.

ck_tokenizer_hash_table_insert Forward
int ck_tokenizer_hash_table_insert(CKTokenizerHashTable * table, const char * key, void * value)

Insert a key-value pair.

ck_tokenizer_hash_table_iterate Forward
int ck_tokenizer_hash_table_iterate(CKTokenizerHashTable * table, CKTokenizerHashCallback callback, void * user_data)

ck_tokenizer_hash_table_keys Forward
size_t ck_tokenizer_hash_table_keys(CKTokenizerHashTable * table, const char ** out_keys, size_t max_keys)

Get all keys as an array.

ck_tokenizer_hash_table_lookup Forward
void * ck_tokenizer_hash_table_lookup(CKTokenizerHashTable * table, const char * key)

Look up a key.

ck_tokenizer_hash_table_lookup_avx Forward
void * ck_tokenizer_hash_table_lookup_avx(CKTokenizerHashTable * table, const char * key)

ck_tokenizer_id_to_token Forward
const char * ck_tokenizer_id_to_token(const CKTokenizer * tok, int32_t id)

ck_tokenizer_init Forward
int ck_tokenizer_init(CKTokenizer * tok)

ck_tokenizer_load Forward
int ck_tokenizer_load(CKTokenizer * tok, const char * path)

ck_tokenizer_load_binary Forward
int ck_tokenizer_load_binary(CKTokenizer * tok, int vocab_size, const int32_t * offsets, const char * strings, int num_merges, const int32_t * merges)

Load vocabulary from memory-mapped binary data.

ck_tokenizer_load_binary_with_scores Forward
int ck_tokenizer_load_binary_with_scores(CKTokenizer * tok, int vocab_size, const int32_t * offsets, const char * strings, const float * scores, const uint8_t * types, int num_merges, const int32_t * merges)

Load vocabulary from memory-mapped binary data with scores and types.

ck_tokenizer_load_gguf Forward
int ck_tokenizer_load_gguf(CKTokenizer * tok, const char * path)

Load vocabulary from GGUF file.

ck_tokenizer_load_json Forward
int ck_tokenizer_load_json(CKTokenizer * tok, const char * path)

Load vocabulary from JSON file (HuggingFace format).

ck_tokenizer_load_merges Forward
int ck_tokenizer_load_merges(CKTokenizer * tok, const char * path)

Load BPE merges from text file.

ck_tokenizer_load_text Forward
int ck_tokenizer_load_text(CKTokenizer * tok, const char * path)

Load vocabulary from text file (one token per line).

ck_tokenizer_lookup Forward
int32_t ck_tokenizer_lookup(const CKTokenizer * tok, const char * token, int len)

ck_tokenizer_lookup_exact Forward
int32_t ck_tokenizer_lookup_exact(const CKTokenizer * tok, const char * token)

ck_tokenizer_lookup_exact_n Forward
int32_t ck_tokenizer_lookup_exact_n(const CKTokenizer * tok, const char * text, int text_len)

ck_tokenizer_lookup_merge Forward
int ck_tokenizer_lookup_merge(const CKTokenizer * tok, int32_t left, int32_t right)

ck_tokenizer_mempool_alloc Forward
void * ck_tokenizer_mempool_alloc(CKTokenizerMemPool * pool, size_t size)

Allocate from pool.

ck_tokenizer_mempool_alloc_aligned Forward
void * ck_tokenizer_mempool_alloc_aligned(CKTokenizerMemPool * pool, size_t size, size_t align)

Allocate aligned memory from pool.

ck_tokenizer_mempool_alloc_count Forward
size_t ck_tokenizer_mempool_alloc_count(CKTokenizerMemPool * pool)

Get allocation count.

ck_tokenizer_mempool_available Forward
size_t ck_tokenizer_mempool_available(CKTokenizerMemPool * pool)

Get available bytes in pool.

ck_tokenizer_mempool_free Forward
void ck_tokenizer_mempool_free(CKTokenizerMemPool * pool)

Free a memory pool.

ck_tokenizer_mempool_init Forward
int ck_tokenizer_mempool_init(CKTokenizerMemPool * pool, size_t size)

Initialize a memory pool.

ck_tokenizer_mempool_reset Forward
void ck_tokenizer_mempool_reset(CKTokenizerMemPool * pool)

Reset pool (mark all memory as free).

ck_tokenizer_mempool_strdup Forward
char * ck_tokenizer_mempool_strdup(CKTokenizerMemPool * pool, const char * str)

Allocate and copy string (strdup equivalent).

ck_tokenizer_mempool_strndup Forward
char * ck_tokenizer_mempool_strndup(CKTokenizerMemPool * pool, const char * str, int len)

Allocate and copy string with length.

ck_tokenizer_mempool_used Forward
size_t ck_tokenizer_mempool_used(CKTokenizerMemPool * pool)

Get used bytes in pool.

ck_tokenizer_reset Forward
void ck_tokenizer_reset(CKTokenizer * tok)

ck_tokenizer_set_add_bos_eos Forward
void ck_tokenizer_set_add_bos_eos(CKTokenizer * tok, bool add_bos, bool add_eos)

ck_tokenizer_set_add_space_prefix Forward
void ck_tokenizer_set_add_space_prefix(CKTokenizer * tok, bool add_space_prefix)

ck_tokenizer_set_space_prefix_style Forward
void ck_tokenizer_set_space_prefix_style(CKTokenizer * tok, CKSpacePrefixStyle style)

ck_tokenizer_set_special_ids Forward
void ck_tokenizer_set_special_ids(CKTokenizer * tok, int32_t unk, int32_t bos, int32_t eos, int32_t pad, int32_t mask)

ck_tokenizer_set_spm_mode Forward
void ck_tokenizer_set_spm_mode(CKTokenizer * tok, CKSpmMode spm_mode)

ck_tokenizer_set_use_trie Forward
void ck_tokenizer_set_use_trie(CKTokenizer * tok, bool use_trie)

ck_tokenizer_utf8_normalize_nfc Forward
size_t ck_tokenizer_utf8_normalize_nfc(const char * src, size_t src_len, char * dst, size_t dst_size)

ck_tokenizer_vocab_size Forward
size_t ck_tokenizer_vocab_size(const CKTokenizer * tok)

Get vocabulary size.

ck_topk_insert_desc Forward
void ck_topk_insert_desc(int idx, float val, int * indices, float * values, int k)

ck_train_bias_reduce_compute_range Forward
void ck_train_bias_reduce_compute_range(const float * d_output, float * d_bias, int T, int out_start, int out_end, int aligned_out)

ck_train_outer_t1_compute_range Forward
void ck_train_outer_t1_compute_range(const float * d_output, const float * input, float * d_W, float * d_b, int out_start, int out_end, int aligned_in)

ck_train_outer_t1_work Forward
void ck_train_outer_t1_work(int ith, int nth, void * argp)

ck_train_pick_active_threads Forward
int ck_train_pick_active_threads(int nth, size_t work_items, size_t min_chunk)

ck_trie_clear Forward
void ck_trie_clear(CKTrie * trie)

ck_trie_create Forward
CKTrie * ck_trie_create(size_t max_nodes)

ck_trie_find_longest Forward
int32_t ck_trie_find_longest(const CKTrie * trie, const char * text, size_t text_len, size_t start_pos, size_t * match_len)

ck_trie_free Forward
void ck_trie_free(CKTrie * trie)

ck_trie_has_prefix Forward
bool ck_trie_has_prefix(const CKTrie * trie, const char * text, size_t text_len, size_t pos)

ck_trie_insert Forward
int ck_trie_insert(CKTrie * trie, const char * token, int32_t token_id, bool is_special, int32_t priority)

ck_trie_node_count Forward
size_t ck_trie_node_count(const CKTrie * trie)

ck_true_bpe_add_merge Forward
int ck_true_bpe_add_merge(CKTrueBPE * bpe, int32_t left_id, int32_t right_id, int32_t merged_id, int32_t priority)

ck_true_bpe_add_merge_by_tokens Forward
int ck_true_bpe_add_merge_by_tokens(CKTrueBPE * bpe, const char * left, const char * right, int32_t priority)

ck_true_bpe_add_special_token Forward
int ck_true_bpe_add_special_token(CKTrueBPE * bpe, const char * token, int32_t id)

ck_true_bpe_add_token Forward
int ck_true_bpe_add_token(CKTrueBPE * bpe, const char * token, int32_t id, float score)

ck_true_bpe_create Forward
CKTrueBPE * ck_true_bpe_create(void)

ck_true_bpe_decode Forward
int ck_true_bpe_decode(const CKTrueBPE * bpe, const int32_t * ids, int num_ids, char * text, int max_len)

ck_true_bpe_detect_space_style Forward
CKSpacePrefixStyle ck_true_bpe_detect_space_style(CKTrueBPE * bpe)

ck_true_bpe_encode Forward
int ck_true_bpe_encode(CKTrueBPE * bpe, const char * text, int text_len, int32_t * ids, int max_ids)

ck_true_bpe_free Forward
void ck_true_bpe_free(CKTrueBPE * bpe)

ck_true_bpe_id_to_token Forward
const char * ck_true_bpe_id_to_token(const CKTrueBPE * bpe, int32_t id)

ck_true_bpe_load_binary Forward
int ck_true_bpe_load_binary(CKTrueBPE * bpe, int vocab_size, const int32_t * offsets, const char * strings, int num_merges, const int32_t * merges)

ck_true_bpe_lookup Forward
int32_t ck_true_bpe_lookup(const CKTrueBPE * bpe, const char * token)

ck_true_bpe_num_merges Forward
int32_t ck_true_bpe_num_merges(const CKTrueBPE * bpe)

ck_true_bpe_set_config Forward
void ck_true_bpe_set_config(CKTrueBPE * bpe, const CKBPEConfig * config)

ck_true_bpe_set_special_ids Forward
void ck_true_bpe_set_special_ids(CKTrueBPE * bpe, int32_t unk, int32_t bos, int32_t eos, int32_t pad)

ck_true_bpe_vocab_size Forward
size_t ck_true_bpe_vocab_size(const CKTrueBPE * bpe)

ck_ue4m3_to_fp32 Forward
float ck_ue4m3_to_fp32(uint8_t value)

ck_ue4m3_to_fp32_inline Forward
float ck_ue4m3_to_fp32_inline(uint8_t value)

ck_utf8_byte_to_offset Forward
size_t ck_utf8_byte_to_offset(const char * str, size_t len, size_t byte_offset)

Get the character index from byte offset.

ck_utf8_char_length Forward
int ck_utf8_char_length(unsigned char c)

Get the length of a UTF-8 character from its first byte.

ck_utf8_count_chars Forward
size_t ck_utf8_count_chars(const char * str, size_t len)

Count UTF-8 characters in a string.

ck_utf8_decode_2 Forward
uint32_t ck_utf8_decode_2(const char * s)

Get 2-byte UTF-8 sequence value.

ck_utf8_decode_3 Forward
uint32_t ck_utf8_decode_3(const char * s)

Get 3-byte UTF-8 sequence value.

ck_utf8_decode_4 Forward
uint32_t ck_utf8_decode_4(const char * s)

Get 4-byte UTF-8 sequence value.

ck_utf8_first_byte Forward
unsigned char ck_utf8_first_byte(const char * s)

Get the first byte of a UTF-8 character.

ck_utf8_from_cp Forward
int ck_utf8_from_cp(uint32_t cp, char * out)

Write a Unicode code point as UTF-8.

ck_utf8_is_continuation Forward
int ck_utf8_is_continuation(unsigned char c)

Get the continuation byte mask and value.

ck_utf8_is_valid Forward
bool ck_utf8_is_valid(const char * str, size_t len)

Check if a byte sequence is valid UTF-8.

ck_utf8_is_whitespace Forward
bool ck_utf8_is_whitespace(uint32_t cp)

Check if character is whitespace (Unicode White_Space property).

ck_utf8_next_char Forward
int32_t ck_utf8_next_char(const char ** str, int * out_len)

Get next UTF-8 character, return its code point.

ck_utf8_normalize_nfc Forward
size_t ck_utf8_normalize_nfc(const char * src, size_t src_len, char * dst, size_t dst_size)

Normalize UTF-8 string (Unicode normalization form NFC).

ck_utf8_offset_to_byte Forward
size_t ck_utf8_offset_to_byte(const char * str, size_t len, size_t n)

Get the byte offset of the N-th character.

ck_utf8_validate Forward
size_t ck_utf8_validate(const char * str, size_t len)

Validate a UTF-8 string.

ck_vec_dot_f32_reverse_strict Forward
float ck_vec_dot_f32_reverse_strict(const float * x, const float * y, int n)

ck_vec_dot_f32_strict Forward
float ck_vec_dot_f32_strict(const float * x, const float * y, int n)

ck_vec_dot_f32x_f32_to_f32_via_f64 Forward
float ck_vec_dot_f32x_f32_to_f32_via_f64(const float * x, const float * y, int n)

ck_vec_max_f32_contig Forward
float ck_vec_max_f32_contig(const float * x, int n)

ck_vec_scale_f32_inplace Forward
void ck_vec_scale_f32_inplace(float * x, int n, float scale)

ck_weight_dtype_expr Forward
const char * ck_weight_dtype_expr(const CKBufferSpec * spec)

ckernel_backend_native
CKMathBackend ckernel_backend_native(void)

Obtain the built-in native backend (single-node CPU, C + intrinsics).

clamp_int8 Forward
int8_t clamp_int8(float value)

compute_align Forward
CKV2AlignInfo compute_align(const CKModelConfig * cfg)

compute_rms_scale Forward
float compute_rms_scale(const float * x, int n, float eps)

compute_rms_scale_internal Forward
float compute_rms_scale_internal(const float * x, int n, float eps)

convert_bf16_tensor_to_buf Forward
void convert_bf16_tensor_to_buf(const uint16_t * src, float * dst, size_t count)

convert_f16_to_f32 Forward
void convert_f16_to_f32(float * dst, const uint16_t * src, size_t count)

Convert FP16 tensor to FP32.

convert_f32_to_f16 Forward
void convert_f32_to_f16(uint16_t * dst, const float * src, size_t count)

Convert FP32 tensor to FP16.

convert_float_to_int4 Forward
void convert_float_to_int4(const float * src, uint8_t * dst, size_t count)

convert_float_to_int8 Forward
void convert_float_to_int8(const float * src, int8_t * dst, size_t count)

convert_int4_to_float Forward
void convert_int4_to_float(const uint8_t * src, float * dst, size_t count)

convert_int8_to_float Forward
void convert_int8_to_float(const int8_t * src, float * dst, size_t count)

count_set_bits Forward
int count_set_bits(const char * hex_mask)

cpu_features_init Forward
void cpu_features_init(void)

create_entry Forward
CKTokenizerHashEntry * create_entry(const char * key, const void * value, size_t value_size)

create_node Forward
CKTrieNode * create_node(void)

decode_bpe_token Forward
int decode_bpe_token(const char * token, char * out, int max)

Decode GPT-2 byte-level BPE representation back to actual bytes.

decode_int4 Forward
int8_t decode_int4(uint8_t packed, int index)

decode_layer_parallel Forward
void decode_layer_parallel(float * hidden, const void * ln1_weight, const void * ln2_weight, const void * WQ, const void * WK, const void * WV, const void * WO, const void * W_gate, const void * W_up, const void * W_down, float * k_cache, float * v_cache, int token_index, float * scratch, int embed_dim, int intermediate, int H, int H_kv, int head_dim, int max_seq, float eps, int num_threads)

Process one transformer layer in parallel.

decode_utf8_scalar Forward
int decode_utf8_scalar(const unsigned char * s, int len, int * out_cp, int * out_used)

deepseek_mhc_mix_backward_f32 Backward
void deepseek_mhc_mix_backward_f32(const float * d_out, const float * streams, const float * mix, float * d_streams, float * d_mix, int tokens, int n_streams, int dim)

Backward pass / gradient computation

deepseek_mhc_mix_f32 Forward
void deepseek_mhc_mix_f32(const float * streams, const float * mix, float * out, int tokens, int n_streams, int dim)

deepseek_mla_kv_cache_batch_store_f32
void deepseek_mla_kv_cache_batch_store_f32(float * k_cache, float * v_cache, const float * k, const float * v, int num_tokens, int num_kv_heads, int qk_head_dim, int v_head_dim, int max_seq_len, int cache_stride)

deepseek_mla_kv_cache_store_f32
void deepseek_mla_kv_cache_store_f32(float * k_cache, float * v_cache, const float * k, const float * v, int pos, int num_kv_heads, int qk_head_dim, int v_head_dim, int max_seq_len, int cache_stride)

deepseek_mla_kv_decompress_bf16 Forward
void deepseek_mla_kv_decompress_bf16(const float * compressed_kv, const uint16_t * kv_b_proj, float * k_nope, float * value, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int v_dim)

deepseek_mla_kv_decompress_bf16_parallel_dispatch Forward
void deepseek_mla_kv_decompress_bf16_parallel_dispatch(const float * compressed_kv, const uint16_t * kv_b_proj, float * k_nope, float * value, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int v_dim)

deepseek_mla_kv_decompress_bf16_token_range Forward
void deepseek_mla_kv_decompress_bf16_token_range(const float * compressed_kv, const uint16_t * kv_b_proj, float * k_nope, float * value, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int v_dim, int token_begin, int token_end)

deepseek_mla_kv_decompress_f32 Forward
void deepseek_mla_kv_decompress_f32(const float * compressed_kv, const float * kv_b_proj, float * k_nope, float * value, int tokens, int heads, int kv_lora_rank, int qk_nope_dim, int v_dim)

dequant_q4_0_block Forward
void dequant_q4_0_block(const block_q4_0 * block, float * output)

Dequantize a single Q4_0 block to FP32.

dequant_q4_0_row Forward
void dequant_q4_0_row(const void * src, float * dst, size_t n_elements)

Dequantize Q4_0 row (multiple blocks)

dequant_q4_1_block Forward
void dequant_q4_1_block(const block_q4_1 * block, float * output)

Dequantize a single Q4_1 block to FP32.

dequant_q4_1_row Forward
void dequant_q4_1_row(const void * src, float * dst, size_t n_elements)

Dequantize Q4_1 row (multiple blocks)

dequant_q4_k_block Forward
void dequant_q4_k_block(const block_q4_K * block, float * output)

Dequantize a single Q4_K block to FP32.

dequant_q4_k_row Forward
void dequant_q4_k_row(const void * src, float * dst, size_t n_elements)

Dequantize Q4_K row (multiple blocks)

dequant_q5_0_block Forward
void dequant_q5_0_block(const block_q5_0 * block, float * output)

Dequantize a single Q5_0 block to FP32.

dequant_q5_0_row Forward
void dequant_q5_0_row(const void * src, float * dst, size_t n_elements)

Dequantize Q5_0 row (multiple blocks)

dequant_q5_1_block Forward
void dequant_q5_1_block(const block_q5_1 * block, float * output)

Dequantize a single Q5_1 block to FP32.

dequant_q5_1_row Forward
void dequant_q5_1_row(const void * src, float * dst, size_t n_elements)

Dequantize Q5_1 row (multiple blocks)

dequant_q6_k_block Forward
void dequant_q6_k_block(const block_q6_K * block, float * output)

Dequantize a single Q6_K block to FP32.

dequant_q6_k_row Forward
void dequant_q6_k_row(const void * src, float * dst, size_t n_elements)

Dequantize Q6_K row (multiple blocks)

dequant_q8_0_block Forward
void dequant_q8_0_block(const block_q8_0 * block, float * output)

Dequantize a single Q8_0 block to FP32.

dequant_q8_0_row Forward
void dequant_q8_0_row(const void * src, float * dst, size_t n_elements)

Dequantize Q8_0 row (multiple blocks)

dequant_row Forward
void dequant_row(CKDataType dtype, const void * src, float * dst, size_t n_elements)

Dequantize a row of quantized data to FP32.

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

detach_allocation Forward
ck_huge_alloc_entry_t * detach_allocation(void * ptr)

detect_chat_template Forward
ChatTemplateType detect_chat_template(const char * model_name)

detect_gpt2_byte_fallback Forward
bool detect_gpt2_byte_fallback(CKTrueBPE * tokenizer, int vocab_size)

detect_physical_cores Forward
int detect_physical_cores(void)

dot_f16 Forward
float dot_f16(const uint16_t * w_f16, const float * x, int K)

dot_fp32_q5_0_block Forward
float dot_fp32_q5_0_block(const float * x, const block_q5_0 * block)

Compute dot product of FP32 input with Q5_0 weight block, with online Q8 quantization.

dot_fp32_q8_0_block Forward
float dot_fp32_q8_0_block(const float * x, const block_q8_0 * block)

Compute dot product of FP32 input with Q8_0 weight block, with online Q8 quantization.

dot_q4_0 Forward
float dot_q4_0(const void * w_q4_0, const float * x, int K)

dot_q4_1 Forward
float dot_q4_1(const void * w_q4_1, const float * x, int K)

dot_q4_k Forward
float dot_q4_k(const void * w_q4k, const float * x, int K)

Compute dot product of Q4_K row with FP32 vector.

dot_q4_k_packed_meta_q8_k_block Forward
float dot_q4_k_packed_meta_q8_k_block(const block_q4_K_packed_meta * w, const block_q8_K * x)

dot_q4_k_packed_u8_q8_k_block Forward
float dot_q4_k_packed_u8_q8_k_block(const block_q4_K_packed_u8 * w, const block_q8_K * x)

dot_q4_k_packed_vnni_x8_q8_k_compact_order Forward
void dot_q4_k_packed_vnni_x8_q8_k_compact_order(float block_sums, const block_q4_K_packed_vnni_x8 * w, const block_q8_K * x, int rows)

dot_q4_k_q8_k_ref Forward
float dot_q4_k_q8_k_ref(const block_q4_K * w, const block_q8_K * x, int k)

dot_q4_packed_u8_q8_32_ref Forward
int32_t dot_q4_packed_u8_q8_32_ref(const uint8_t * q4_32, const int8_t * q8_32)

dot_q5_0 Forward
float dot_q5_0(const void * w_q5_0, const float * x, int K)

dot_q5_0_q8_k_32_sse Forward
float dot_q5_0_q8_k_32_sse(const block_q5_0 * bw, const block_q8_K * ba, int q8_offset)

dot_q5_1 Forward
float dot_q5_1(const void * w_q5_1, const float * x, int K)

dot_q5_1_q8_1_block Forward
float dot_q5_1_q8_1_block(const block_q5_1 * w, const block_q8_1 * x)

dot_q5_k_q8_k_row Forward
float dot_q5_k_q8_k_row(const block_q5_K * w, const block_q8_K * x, int nb)

dot_q6_k_q8_k_256_sse Forward
float dot_q6_k_q8_k_256_sse(const block_q6_K * bw, const block_q8_K * ba)

SSE Optimized dot product for Q6_K x Q8_K Q6_K layout: ql: 128 bytes (low 4 bits) qh: 64 bytes (high 2 bits) scales: 16 bytes (int8 scales) d: fp16 super-scale

dot_q6_k_q8_k_ref Forward
float dot_q6_k_q8_k_ref(const block_q6_K * w, const block_q8_K * x, int K)

Scalar dot product for Q6_K x Q8_K.

dot_q6_k_ref Forward
float dot_q6_k_ref(const block_q6_K * w, const float * x, int K)

dot_q8_0 Forward
float dot_q8_0(const void * w_q8_0, const float * x, int K)

ds_mhc_idx Forward
size_t ds_mhc_idx(int t, int s, int d, int n_streams, int dim)

ds_mix_idx Forward
size_t ds_mix_idx(int t, int out_s, int in_s, int n_streams)

ds_mla_bf16_round Forward
float ds_mla_bf16_round(float value)

ds_mla_kv_decompress_bf16_rows Forward
void ds_mla_kv_decompress_bf16_rows(int begin, int end, void * opaque)

ds_mla_thd_idx Forward
size_t ds_mla_thd_idx(int t, int h, int d, int heads, int dim)

ds_mla_tok_idx Forward
size_t ds_mla_tok_idx(int t, int d, int dim)

ds_qkv_idx Forward
size_t ds_qkv_idx(int token, int head, int d, int heads, int dim)

embedding_backward Backward
void embedding_backward(const int32_t * token_ids, int token_count, const float * d_output, float * d_token_embeddings, float * d_pos_embeddings, int vocab_size, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Backward pass / gradient computation

embedding_backward_bf16 Backward
void embedding_backward_bf16(const int32_t * token_ids, int token_count, const uint16_t * d_output, uint16_t * d_token_embeddings, uint16_t * d_pos_embeddings, int vocab_size, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Backward pass / gradient computation

embedding_backward_bf16_mixed Backward
void embedding_backward_bf16_mixed(const int32_t * token_ids, int token_count, const uint16_t * d_output, float * d_token_embeddings, float * d_pos_embeddings, int vocab_size, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Backward pass / gradient computation

embedding_forward Forward
void embedding_forward(const int32_t * token_ids, int token_count, int vocab_size, const float * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_bf16 Forward
void embedding_forward_bf16(const int32_t * token_ids, int token_count, int vocab_size, const uint16_t * token_embeddings, const uint16_t * pos_embeddings, uint16_t * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_bf16_fp32 Forward
void embedding_forward_bf16_fp32(const int32_t * token_ids, int token_count, int vocab_size, const uint16_t * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_q4_k Forward
void embedding_forward_q4_k(const int32_t * token_ids, int token_count, int vocab_size, const void * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_q5_0 Forward
void embedding_forward_q5_0(const int32_t * token_ids, int token_count, int vocab_size, const void * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_q6_k Forward
void embedding_forward_q6_k(const int32_t * token_ids, int token_count, int vocab_size, const void * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

embedding_forward_q8_0 Forward
void embedding_forward_q8_0(const int32_t * token_ids, int token_count, int vocab_size, const void * token_embeddings, const float * pos_embeddings, float * output, int embed_dim, int aligned_embed_dim, int context_window, int add_pos)

Forward pass computation

emit_body_fields Forward
void emit_body_fields(FILE * out, const CKIRV2Graph * graph, CKBufferRole role_filter, int activation_group)

emit_body_values Forward
size_t emit_body_values(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * plan, CKBufferRole role_filter, int activation_group)

emit_bump_bytes_assignment Forward
void emit_bump_bytes_assignment(FILE * out, const char * indent, const char * struct_prefix, const char * name, const CKDimToken * shape)

emit_bump_bytes_assignment_weight_dtype Forward
void emit_bump_bytes_assignment_weight_dtype(FILE * out, const char * indent, const char * struct_prefix, const char * name, const CKDimToken * shape, const char * dtype_expr)

emit_dim_expr Forward
void emit_dim_expr(FILE * out, CKDimKind dim)

emit_footer_fields Forward
void emit_footer_fields(FILE * out, const CKIRV2Graph * graph, CKBufferRole role_filter, int activation_group)

emit_footer_values Forward
void emit_footer_values(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * plan, CKBufferRole role_filter, int activation_group, size_t * offset)

emit_global_aliases_to_layer Forward
void emit_global_aliases_to_layer(FILE * out)

emit_global_allocations Forward
void emit_global_allocations(FILE * out)

emit_global_offset_fields Forward
void emit_global_offset_fields(FILE * out)

emit_header_fields Forward
void emit_header_fields(FILE * out, const CKIRV2Graph * graph, CKBufferRole role_filter, int activation_group)

emit_header_values Forward
void emit_header_values(FILE * out, const CKIRV2Graph * graph, const CKMemPlan * plan, CKBufferRole role_filter, int activation_group, size_t * offset)

emit_kernel_manifest Forward
int emit_kernel_manifest(const CKIRGraph * forward, const char * runtime_path)

emit_layer_allocations Forward
void emit_layer_allocations(FILE * out)

emit_layer_offsets_struct Forward
void emit_layer_offsets_struct(FILE * out)

emit_library_api Forward
void emit_library_api(FILE * out, const CKIRGraph * forward)

emit_model_struct Forward
void emit_model_struct(FILE * out)

emit_offset_field Forward
void emit_offset_field(FILE * out, const char * name)

emit_plan_sources Forward
int emit_plan_sources(FILE * f, const CKPlanStep * plan, size_t plan_count, const CKIRGraph * cfg, const char ** seen, size_t * seen_count, size_t seen_cap)

emit_runtime_preamble Forward
int emit_runtime_preamble(FILE * out)

emit_schedule_block Forward
void emit_schedule_block(FILE * out, const CKIRV2Graph * graph, const char * func_name, const char * label, const char * runtime_sym)

emit_sgd_update Forward
void emit_sgd_update(FILE * out)

emit_shape_expr Forward
void emit_shape_expr(FILE * out, const CKDimToken * shape)

emit_span_field Forward
void emit_span_field(FILE * out, const char * label)

emit_span_value Forward
void emit_span_value(FILE * out, const char * label, size_t offset, size_t size, int comma)

emit_training_conditional_assignment Forward
void emit_training_conditional_assignment(FILE * out, const char * indent, const char * struct_prefix, const char * name, const CKDimToken * shape)

emit_unique_source Forward
int emit_unique_source(FILE * f, const char * path, const char ** seen, size_t * seen_count, size_t seen_cap)

emit_zero_grad Forward
void emit_zero_grad(FILE * out)

encode_chunk Forward
int encode_chunk(CKTrueBPE * bpe, const char * chunk, int chunk_len, int32_t * ids, int max_ids, CKBPETokenList * list)

encode_int4_nibble Forward
uint8_t encode_int4_nibble(int8_t value)

encode_text_segment Forward
int encode_text_segment(CKTrueBPE * bpe, const char * text, int text_len, int32_t * ids, int max_ids)

engine_thread_func Forward
void * engine_thread_func(void * arg)

env_flag_enabled Forward
int env_flag_enabled(const char * name)

eos_is_potential_prefix Forward
bool eos_is_potential_prefix(const char * token)

Check if token might be start of EOS pattern.

eos_pattern_init Forward
void eos_pattern_init(ChatTemplateType tmpl)

eos_pattern_process Forward
bool eos_pattern_process(const char * token_text, char * out_buf, size_t * out_len, void(*)(char *, size_t *, const char *) output_fn, ChatTemplateType tmpl)

Process a token for EOS pattern detection.

eos_pattern_reset Forward
void eos_pattern_reset(void)

feature_concat Forward
void feature_concat(const float * main_input, const float * branch_input, float * output, int rows, int main_dim, int branch_slice_dim, int num_branch_slices)

feature_concat_2way Forward
void feature_concat_2way(const float * main_input, const float * branch_input, float * output, int rows, int main_dim, int branch_slice_dim, int num_branch_slices)

feature_slice_copy Forward
void feature_slice_copy(const float * src, float * dst, int rows, int src_dim, int dst_dim, int dst_feature_offset)

final_logit_scale_f32 Forward
void final_logit_scale_f32(float * logits, int tokens, int vocab_size, float scale)

find_best_merge Forward
int find_best_merge(const CKTrueBPE * bpe, const CKBPETokenList * list, size_t * best_pos, const CKBPEMerge ** best_merge)

find_buffer_by_name Forward
int find_buffer_by_name(const CKIRV2Graph * graph, const char * name)

find_longest_match Forward
int32_t find_longest_match(const CKTokenizer * tok, const char * text, size_t text_len, size_t pos, size_t * match_len)

find_longest_match_hash Forward
int32_t find_longest_match_hash(const CKTokenizer * tok, const char * text, size_t text_len, size_t pos, size_t * match_len)

find_longest_match_trie Forward
int32_t find_longest_match_trie(const CKTokenizer * tok, const char * text, size_t text_len, size_t pos, size_t * match_len)

find_model_in_cache
bool find_model_in_cache(const char * model_name, char * lib_out, char * weights_out, size_t out_size)

find_object_range Forward
int find_object_range(const char * json, const char * key, const char ** out_start, size_t * out_len)

flatten_head_major Forward
void flatten_head_major(const float * attn_out, float * dst, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

float_tensor_to_bf16 Forward
void float_tensor_to_bf16(const float * src, uint16_t * dst, size_t count)

float_to_bf16 Forward
uint16_t float_to_bf16(float f)

format_bandwidth Forward
const char * format_bandwidth(float bw_gbs, char * buf, size_t buf_size)

format_size Forward
const char * format_size(uint64_t size_mb, char * buf, size_t buf_size)

free_entry Forward
void free_entry(CKTokenizerHashEntry * entry, bool free_value)

fused_kernels_compute_kv_tile Forward
int fused_kernels_compute_kv_tile(int l1_size, int head_dim, int bytes_per_elem)

Compute optimal KV tile size for flash attention.

fused_kernels_report_stats Forward
void fused_kernels_report_stats(int hidden, int num_layers, int seq_len)

Report memory savings from mega-fusion.

fused_kernels_validate_constraints Forward
int fused_kernels_validate_constraints(int l1_size, int head_dim, int kv_tile_size, int bytes_per_elem)

Validate cache constraints for fusion.

fused_output_projection_residual Forward
void fused_output_projection_residual(float * output, const float * o_all, const float * W_o, const float * b_o, const float * residual, int hidden, int num_heads, int head_dim)

Fused output projection with residual add.

gated_deltanet_autoregressive_backward Backward
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)

Backward pass / gradient computation

gated_deltanet_autoregressive_backward_ref Backward
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)

Backward pass / gradient computation

gated_deltanet_autoregressive_forward Forward
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)

Forward pass computation

gated_deltanet_autoregressive_forward_ref 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)

Forward pass computation

gated_deltanet_impl_name Forward
const char * gated_deltanet_impl_name(void)

gated_deltanet_llama_avx2_forward Forward
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)

Forward pass computation

gated_deltanet_llama_avx2_forward_head_range 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)

Forward pass computation

gated_deltanet_llama_avx2_grouped_forward_impl Forward
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)

Forward pass computation

gated_deltanet_llama_avx2_grouped_forward_transposed_impl Forward
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)

Forward pass computation

gated_deltanet_llama_avx2_prefill_forward Forward
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)

Forward pass computation

gated_deltanet_llama_chunk64_head_forward 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)

Forward pass computation

gated_deltanet_llama_chunk64_prefill_forward Forward
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)

Forward pass computation

gated_deltanet_prefill_forward 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)

Forward pass computation

gated_deltanet_pytorch_gate_values_debug Forward
void gated_deltanet_pytorch_gate_values_debug(const float * g, const float * beta, float * gate_values, float * beta_values, int num_heads)

gated_deltanet_pytorch_grouped_bf16_forward Forward
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)

Forward pass computation

gated_deltanet_pytorch_grouped_bf16_forward_debug 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)

Forward pass computation

gated_deltanet_pytorch_grouped_bf16_forward_impl Forward
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)

Forward pass computation

gated_deltanet_pytorch_grouped_bf16_prefill_forward Forward
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)

Forward pass computation

geglu_backward_bf16_mixed Backward
void geglu_backward_bf16_mixed(const uint16_t * x, const uint16_t * d_out, float * d_x, int tokens, int dim)

Backward pass / gradient computation

geglu_backward_fp32 Backward
void geglu_backward_fp32(const float * x, const float * d_out, float * d_x, int tokens, int dim)

Backward pass / gradient computation

geglu_forward_bf16 Forward
void geglu_forward_bf16(const uint16_t * x, uint16_t * out, int tokens, int dim, float * scratch)

Forward pass computation

geglu_forward_exact Forward
void geglu_forward_exact(const float * x, float * out, int tokens, int dim)

Forward pass computation

geglu_forward_fp32 Forward
void geglu_forward_fp32(const float * x, float * out, int tokens, int dim)

Forward pass computation

geglu_forward_ggml_native Forward
void geglu_forward_ggml_native(const float * x, float * out, int tokens, int dim)

Forward pass computation

gemv_bf16 Forward
void gemv_bf16(float * y, const void * W, const float * x, int M, int K)

gemv_bf16_bf16_storage Forward
void gemv_bf16_bf16_storage(float * y, const void * W, const float * x, int M, int K)

gemv_bf16_bf16_storage_parallel_dispatch Forward
void gemv_bf16_bf16_storage_parallel_dispatch(float * y, const void * W, const float * x, int M, int K)

gemv_bf16_bf16_storage_row_range Forward
void gemv_bf16_bf16_storage_row_range(float * y, const uint16_t * w, const float * x, int M, int K, int row_begin, int row_end)

gemv_bf16_parallel_dispatch Forward
void gemv_bf16_parallel_dispatch(float * y, const void * W, const float * x, int M, int K)

gemv_bf16_row_range Forward
void gemv_bf16_row_range(float * y, const uint16_t * w, const float * x, int M, int K, int row_begin, int row_end)

gemv_f16 Forward
void gemv_f16(float * y, const uint16_t * W, const float * x, int M, int K)

Auto-dispatch GEMV based on available SIMD.

gemv_f16_backward Backward
void gemv_f16_backward(float * dX, const uint16_t * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_f16_backward_ref Backward
void gemv_f16_backward_ref(float * dX, const uint16_t * W, const float * dY, int M, int K)

Backward pass: compute input gradient (scalar reference)

gemv_f16_ref Forward
void gemv_f16_ref(float * y, const uint16_t * W, const float * x, int M, int K)

Matrix-vector multiply with FP16 weights (scalar reference)

gemv_fused_q5_0_bias Forward
void gemv_fused_q5_0_bias(float * y, const void * W, const float * x, const float * bias, int M, int K)

gemv_fused_q5_0_bias_dispatch Forward
void gemv_fused_q5_0_bias_dispatch(float * y, const void * W, const float * x, const float * bias, int M, int K)

gemv_fused_q5_0_bias_parallel_omp Forward
void gemv_fused_q5_0_bias_parallel_omp(float * y, const void * W, const float * x, const float * bias, int M, int K)

gemv_fused_q8_0_bias Forward
void gemv_fused_q8_0_bias(float * y, const void * W, const float * x, const float * bias, int M, int K)

gemv_fused_q8_0_bias_dispatch Forward
void gemv_fused_q8_0_bias_dispatch(float * y, const void * W, const float * x, const float * bias, int M, int K)

gemv_nt_q5_0_head_major_output Forward
void gemv_nt_q5_0_head_major_output(float * output, const float * attn_out, const void * wo, const float * bias, int tokens, int embed_dim, int num_heads, int head_dim)

Output projection reading head-major attention output (Q5_0 weights)

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

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

gemv_q4_0 Forward
void gemv_q4_0(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV.

gemv_q4_0_backward Backward
void gemv_q4_0_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q4_0_backward_ref Backward
void gemv_q4_0_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient.

gemv_q4_0_ref Forward
void gemv_q4_0_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q4_0 weights (scalar reference)

gemv_q4_1 Forward
void gemv_q4_1(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV.

gemv_q4_1_backward Backward
void gemv_q4_1_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q4_1_backward_ref Backward
void gemv_q4_1_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient.

gemv_q4_1_ref Forward
void gemv_q4_1_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q4_1 weights (scalar reference)

gemv_q4_k Forward
void gemv_q4_k(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV based on available SIMD.

gemv_q4_k_backward Backward
void gemv_q4_k_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q4_k_backward_ref Backward
void gemv_q4_k_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient (scalar reference)

gemv_q4_k_q8_k Forward
void gemv_q4_k_q8_k(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_amx Forward
void gemv_q4_k_q8_k_amx(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_avx Forward
void gemv_q4_k_q8_k_avx(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_avx2 Forward
void gemv_q4_k_q8_k_avx2(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_parallel Forward
void gemv_q4_k_q8_k_parallel(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

gemv_q4_k_q8_k_parallel_simd Forward
void gemv_q4_k_q8_k_parallel_simd(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

gemv_q4_k_q8_k_parallel_vnni Forward
void gemv_q4_k_q8_k_parallel_vnni(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

gemv_q4_k_q8_k_ref Forward
void gemv_q4_k_q8_k_ref(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_sse Forward
void gemv_q4_k_q8_k_sse(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_q8_k_vnni Forward
void gemv_q4_k_q8_k_vnni(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q4_k_ref Forward
void gemv_q4_k_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q4_K weights (scalar reference)

gemv_q5_0 Forward
void gemv_q5_0(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV for Q5_0 weights based on CPU features.

gemv_q5_0_backward Backward
void gemv_q5_0_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q5_0_backward_ref Backward
void gemv_q5_0_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient.

gemv_q5_0_from_fp32 Forward
void gemv_q5_0_from_fp32(float * out, const void * W_q5_0, const float * x_fp32, const float * bias, int M, int K, block_q8_0 * x_q8_scratch)

gemv_q5_0_parallel Forward
void gemv_q5_0_parallel(float * y, const void * W, const float * x, int M, int K, int ith, int nth)

Parallel reference GEMV for Q5_0 × FP32.

gemv_q5_0_parallel_simd Forward
void gemv_q5_0_parallel_simd(float * y, const void * W, const float * x, int M, int K, int ith, int nth)

Parallel SIMD GEMV for Q5_0 × FP32 with prefetching.

gemv_q5_0_q8_0 Forward
void gemv_q5_0_q8_0(float * y, const void * W, const void * x_q8, int M, int K)

Matrix-vector multiply with Q5_0 weights and Q8_0 input.

gemv_q5_0_q8_0_parallel_omp Forward
void gemv_q5_0_q8_0_parallel_omp(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q5_0_q8_0_parallel_simd Forward
void gemv_q5_0_q8_0_parallel_simd(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

Parallel SIMD GEMV for Q5_0 x Q8_0 with prefetching.

gemv_q5_0_ref Forward
void gemv_q5_0_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q5_0 weights (scalar reference)

gemv_q5_1 Forward
void gemv_q5_1(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV.

gemv_q5_1_backward Backward
void gemv_q5_1_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q5_1_backward_ref Backward
void gemv_q5_1_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient.

gemv_q5_1_q8_1 Forward
void gemv_q5_1_q8_1(float * y, const void * W, const float * x, int M, int K)

gemv_q5_1_q8_1_ref Forward
void gemv_q5_1_q8_1_ref(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q5_1_ref Forward
void gemv_q5_1_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q5_1 weights (scalar reference)

gemv_q5_k Forward
void gemv_q5_k(float * y, const void * W, const float * x, int M, int K)

gemv_q5_k_q8_k Forward
void gemv_q5_k_q8_k(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q5_k_q8_k_ref Forward
void gemv_q5_k_q8_k_ref(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q5_k_ref Forward
void gemv_q5_k_ref(float * y, const void * W, const float * x, int M, int K)

gemv_q5_k_ref_fp32 Forward
void gemv_q5_k_ref_fp32(float * y, const void * W, const float * x, int M, int K)

gemv_q6_k Forward
void gemv_q6_k(float * y, const void * W, const float * x, int M, int K)

gemv_q6_k_q8_k Forward
void gemv_q6_k_q8_k(float * y, const void * W, const void * x_q8, int M, int K)

GEMV: y = W @ x where W is Q6_K and x is Q8_K.

gemv_q6_k_q8_k_avx Forward
void gemv_q6_k_q8_k_avx(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q6_k_q8_k_avx2 Forward
void gemv_q6_k_q8_k_avx2(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q6_k_q8_k_avx512 Forward
void gemv_q6_k_q8_k_avx512(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q6_k_q8_k_avx512_vbmi Forward
void gemv_q6_k_q8_k_avx512_vbmi(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q6_k_q8_k_parallel Forward
void gemv_q6_k_q8_k_parallel(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

Parallel reference GEMV for Q6_K × Q8_K.

gemv_q6_k_q8_k_parallel_simd Forward
void gemv_q6_k_q8_k_parallel_simd(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

Parallel SIMD GEMV for Q6_K × Q8_K.

gemv_q6_k_q8_k_ref Forward
void gemv_q6_k_q8_k_ref(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q6_k_q8_k_sse Forward
void gemv_q6_k_q8_k_sse(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q8_0 Forward
void gemv_q8_0(float * y, const void * W, const float * x, int M, int K)

Auto-dispatch GEMV for Q8_0 weights based on CPU features.

gemv_q8_0_backward Backward
void gemv_q8_0_backward(float * dX, const void * W, const float * dY, int M, int K)

Auto-dispatch backward.

gemv_q8_0_backward_ref Backward
void gemv_q8_0_backward_ref(float * dX, const void * W, const float * dY, int M, int K)

Backward pass: compute input gradient (scalar reference)

gemv_q8_0_from_fp32 Forward
void gemv_q8_0_from_fp32(float * out, const void * W_q8_0, const float * x_fp32, const float * bias, int M, int K, block_q8_0 * x_q8_scratch)

gemv_q8_0_parallel_simd Forward
void gemv_q8_0_parallel_simd(float * y, const void * W, const float * x, int M, int K, int ith, int nth)

Parallel SIMD GEMV for Q8_0 weights x FP32 input with prefetching.

gemv_q8_0_q8_0 Forward
void gemv_q8_0_q8_0(float * y, const void * W, const void * x_q8, int M, int K)

Matrix-vector multiply with Q8_0 weights and Q8_0 input.

gemv_q8_0_q8_0_contract Forward
void gemv_q8_0_q8_0_contract(float * y, const void * W, const float * x, int M, int K)

gemv_q8_0_q8_0_parallel Forward
void gemv_q8_0_q8_0_parallel(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

Parallel reference GEMV for Q8_0 x Q8_0.

gemv_q8_0_q8_0_parallel_omp Forward
void gemv_q8_0_q8_0_parallel_omp(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q8_0_q8_0_parallel_simd Forward
void gemv_q8_0_q8_0_parallel_simd(float * y, const void * W, const void * x_q8, int M, int K, int ith, int nth)

Parallel SIMD GEMV for Q8_0 x Q8_0 with prefetching.

gemv_q8_0_q8_0_ref_rows Forward
void gemv_q8_0_q8_0_ref_rows(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q8_0_q8_0_x4 Forward
void gemv_q8_0_q8_0_x4(float * y, const void * W, const void * x_q8, int M, int K)

gemv_q8_0_ref Forward
void gemv_q8_0_ref(float * y, const void * W, const float * x, int M, int K)

Matrix-vector multiply with Q8_0 weights (scalar reference)

get_cache_dir
const char * get_cache_dir(void)

get_cpu_info Forward
const CPUInfo * get_cpu_info(void)

get_current_cpu Forward
int get_current_cpu(void)

get_numa_node_for_cpu Forward
int get_numa_node_for_cpu(int cpu)

get_optimal_decode_threads Forward
int get_optimal_decode_threads(void)

get_time_ms Forward
double get_time_ms(void)

ggml_fp16_to_fp32 Forward
float ggml_fp16_to_fp32(ggml_fp16_t value)

ggml_fp32_to_fp16 Forward
ggml_fp16_t ggml_fp32_to_fp16(float value)

ggml_get_data Forward
void * ggml_get_data(const struct ggml_tensor * tensor)

ggml_get_data_f32 Forward
float * ggml_get_data_f32(const struct ggml_tensor * tensor)

ggml_nbytes Forward
size_t ggml_nbytes(const struct ggml_tensor * tensor)

global_pool_init Forward
void global_pool_init(void)

gpt2_byte_is_identity Forward
bool gpt2_byte_is_identity(unsigned int byte)

gpt2_byte_to_codepoint Forward
unsigned int gpt2_byte_to_codepoint(unsigned int byte)

gpt2_codepoint_to_byte Forward
int gpt2_codepoint_to_byte(int cp)

gpt2_pretokenize Forward
int gpt2_pretokenize(const char * text, int text_len, PretokChunk * chunks, int max_chunks, CKBPEPretokenizer pretokenizer)

gradient_accumulate_bf16 Forward
void gradient_accumulate_bf16(uint16_t * dst, const uint16_t * src, size_t numel)

Accumulate gradients: dst += src (bf16)

gradient_accumulate_f32 Forward
void gradient_accumulate_f32(float * dst, const float * src, size_t numel)

gradient_accumulate_f32_impl Forward
void gradient_accumulate_f32_impl(float * dst, const float * src, size_t numel)

Accumulate gradients: dst += src (fp32)

gradient_accumulate_multi_f32 Forward
void gradient_accumulate_multi_f32(float *const * dsts, const float *const * srcs, const size_t * numels, int tensor_count)

gradient_clip_norm_bf16 Forward
float gradient_clip_norm_bf16(uint16_t * grad, size_t numel, float max_norm)

Clip gradient norm (bf16)

gradient_clip_norm_f32 Forward
float gradient_clip_norm_f32(float * grad, size_t numel, float max_norm)

Clip gradient norm (fp32)

gradient_global_norm_multi_f32 Forward
float gradient_global_norm_multi_f32(const float *const * grads, const size_t * numels, int tensor_count)

gradient_scale_bf16 Forward
void gradient_scale_bf16(uint16_t * grad, size_t numel, float scale)

Scale gradients: grad *= scale (bf16)

gradient_scale_f32 Forward
void gradient_scale_f32(float * grad, size_t numel, float scale)

gradient_scale_f32_impl Forward
void gradient_scale_f32_impl(float * grad, size_t numel, float scale)

Scale gradients by a constant: grad *= scale (fp32)

gradient_sum_sq_f32_impl Forward
double gradient_sum_sq_f32_impl(const float * grad, size_t numel)

group_limited_topk_router_f32_impl Forward
void group_limited_topk_router_f32_impl(const float * scores, const float * correction_bias, int * indices, float * weights, int rows, int n_experts, int top_k, int n_group, int topk_group, int norm_topk_prob, float routed_scaling_factor, int apply_sigmoid)

handle_sigint Forward
void handle_sigint(int sig)

has_cpu_flag Forward
int has_cpu_flag(const char * flags, const char * flag)

hash_pair Forward
uint32_t hash_pair(int32_t left, int32_t right)

hash_string Forward
uint32_t hash_string(const char * s, int len)

hsum_epi32_sse Forward
int32_t hsum_epi32_sse(__m128i v)

hyper_connection_mix_bf16 Forward
void hyper_connection_mix_bf16(const float * hyper_input, const float * norm_weight, const uint16_t * mix_down_weight, const uint16_t * mix_up_weight, const uint16_t * inject_weight, float * mixed_output, float * injection_output, float * normalized_scratch, float * dynamic_scratch, float * mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)

hyper_connection_mix_q4k_q5_0_q4k Forward
void hyper_connection_mix_q4k_q5_0_q4k(const float * hyper_input, const float * norm_weight, const void * mix_down_weight, const void * mix_up_weight, const void * inject_weight, float * mixed_output, float * injection_output, float * normalized_scratch, float * dynamic_scratch, float * mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)

hyper_connection_mix_q6k_q5_0_q4k Forward
void hyper_connection_mix_q6k_q5_0_q4k(const float * hyper_input, const float * norm_weight, const void * mix_down_weight, const void * mix_up_weight, const void * inject_weight, float * mixed_output, float * injection_output, float * normalized_scratch, float * dynamic_scratch, float * mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection)

hyper_connection_mix_quantized Forward
void hyper_connection_mix_quantized(const float * hyper_input, const float * norm_weight, const void * mix_down_weight, const void * mix_up_weight, const void * inject_weight, float * mixed_output, float * injection_output, float * normalized_scratch, float * dynamic_scratch, float * mix_scratch, int rows, int streams, int hidden_dim, int dynamic_dim, float eps, int emit_injection, ck_hyper_q8k_gemm_fn injection_gemm, ck_hyper_q8k_gemm_fn down_gemm)

hyper_injection_q4k_q8k_llama_dispatch Forward
void hyper_injection_q4k_q8k_llama_dispatch(const void * input, const void * weight, const float * bias, float * output, int rows, int output_dim, int input_dim)

hyper_stream_expand_bf16 Forward
void hyper_stream_expand_bf16(const float * input, float * output, int rows, int streams, int hidden_dim)

hyper_stream_expand_f32 Forward
void hyper_stream_expand_f32(const float * input, float * output, int rows, int streams, int hidden_dim)

hyper_stream_inject_bf16 Forward
void hyper_stream_inject_bf16(const float * hyper_input, const float * block_output, const float * injection_weight, float * output, int rows, int streams, int hidden_dim)

hyper_stream_inject_f32 Forward
void hyper_stream_inject_f32(const float * hyper_input, const float * block_output, const float * injection_weight, float * output, int rows, int streams, int hidden_dim)

im2patch Forward
void im2patch(const float * image, float * patches, int C, int H, int W, int P)

im2patch: Transforms an image into a sequence of flattened patches.

im2patch_bf16 Forward
void im2patch_bf16(const uint16_t * image, uint16_t * patches, int C, int H, int W, int P)

init_tokens_from_text Forward
int init_tokens_from_text(CKTrueBPE * bpe, CKBPETokenList * list, const char * text, int text_len)

is_activation_role Forward
int is_activation_role(CKBufferRole role)

is_bpe_digit Forward
bool is_bpe_digit(const char * s, int len)

is_bpe_letter Forward
bool is_bpe_letter(const char * s, int len)

is_bpe_newline Forward
bool is_bpe_newline(const char * s, int len)

is_bpe_punct Forward
bool is_bpe_punct(const char * s, int len)

is_digit Forward
bool is_digit(unsigned char c)

is_eos_token Forward
bool is_eos_token(const CLIOptions * opt, int token)

is_footer_global Forward
int is_footer_global(const char * name)

is_gpt2_space Forward
bool is_gpt2_space(const char * s, int len)

is_letter Forward
bool is_letter(unsigned char c)

is_whitespace Forward
bool is_whitespace(unsigned char c)

is_word_prefix_char Forward
bool is_word_prefix_char(const char * s, int len)

json_match_char Forward
int json_match_char(JSONParser * p, char c)

json_parse_int Forward
int json_parse_int(JSONParser * p, int * out)

json_parse_string Forward
int json_parse_string(JSONParser * p, char * buf, int max_len)

json_skip_value Forward
void json_skip_value(JSONParser * p)

json_skip_whitespace Forward
void json_skip_whitespace(JSONParser * p)

kv_cache_repack_head_major_inplace
void kv_cache_repack_head_major_inplace(float * buf, int num_heads, int tokens, int cache_capacity, int aligned_head_dim)

kv_cache_store
void kv_cache_store(float *__restrict kv_cache_k, float *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int layer, int pos, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_batch_bf16
void kv_cache_store_batch_bf16(uint16_t *__restrict kv_cache_k, uint16_t *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int start_pos, int num_tokens, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_batch_f16
void kv_cache_store_batch_f16(uint16_t *__restrict kv_cache_k, uint16_t *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int start_pos, int num_tokens, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_batch_f32
void kv_cache_store_batch_f32(float *__restrict kv_cache_k, float *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int start_pos, int num_tokens, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_bf16
void kv_cache_store_bf16(uint16_t *__restrict kv_cache_k, uint16_t *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int layer, int pos, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_f16
void kv_cache_store_f16(uint16_t *__restrict kv_cache_k, uint16_t *__restrict kv_cache_v, const float *__restrict k, const float *__restrict v, int layer, int pos, int num_kv_heads, int head_dim, int max_seq_len)

kv_cache_store_shared_q
void kv_cache_store_shared_q(float *__restrict kv_cache_k, float *__restrict kv_cache_v, const float *__restrict q, int layer, int pos, int num_heads, int head_dim, int max_seq_len)

kv_cache_write_head_major
void kv_cache_write_head_major(const float *__restrict k_token, const float *__restrict v_token, float *__restrict k_cache, float *__restrict v_cache, int num_kv_heads, int token_index, int cache_capacity, int head_dim, int aligned_head_dim)

layout_transformer_from_ir Forward
void layout_transformer_from_ir(TransformerModel * m, const CKIRGraph * ir)

Compute a simple forward-only layout for TransformerModel based on: CKModelConfig (dims, heads, vocab, context) The IR graph structure (number of layers, op types)

list_available_models Forward
void list_available_models(void)

load_eos_from_vocab_json Forward
bool load_eos_from_vocab_json(const char * weights_path, CLIOptions * opt)

load_manifest Forward
int load_manifest(const char * path, ManifestEntry ** entries, int * num_entries)

load_model_api Forward
bool load_model_api(const char * lib_path, ModelAPI * api)

load_weights Forward
int load_weights(QWEN2_DECODEModel * model, const char * bump_path, const char * manifest_path)

load_weights_from_bump Forward
int load_weights_from_bump(void * model, const char * bump_path)

logits_copy_to_position Forward
void logits_copy_to_position(const float *__restrict src, float *__restrict dst, int position, int vocab_size)

Copy logits to position-indexed location in output buffer.

lookup_token_exact Forward
int32_t lookup_token_exact(const CKTrueBPE * bpe, const char * token)

main Forward
int main(int argc, char ** argv)

mamba2_conv1d_decode_f32 Forward
void mamba2_conv1d_decode_f32(const float * state_in, const float * x, const float * weight, const float * bias, float * conv_out, float * state_out, int rows, int conv_dim, int kernel_size)

mamba2_conv1d_f32_channel_range Forward
void mamba2_conv1d_f32_channel_range(const float * state_in, const float * x, const float * weight, const float * bias, float * conv_out, float * state_out, int rows, int conv_dim, int kernel_size, int channel_begin, int channel_end)

mamba2_conv1d_f32_parallel_dispatch Forward
void mamba2_conv1d_f32_parallel_dispatch(const float * state_in, const float * x, const float * weight, const float * bias, float * conv_out, float * state_out, int rows, int conv_dim, int kernel_size)

mamba2_dt_softplus_f32 Forward
void mamba2_dt_softplus_f32(const float * dt, const float * dt_bias, float * dt_out, int rows, int num_heads, float dt_min, float dt_max)

mamba2_in_proj_split_f32 Forward
void mamba2_in_proj_split_f32(const float * projected, float * gate, float * hidden_bc, float * dt, int rows, int d_mlp, int intermediate_dim, int conv_dim, int num_heads)

mamba2_selective_scan_f32 Forward
void mamba2_selective_scan_f32(const float * state_init, const float * x, const float * dt, const float * a, const float * b, const float * c, const float * d, float * state_out, float * y, int batch, int seq_len, int num_heads, int head_dim, int state_dim, int num_groups)

mamba2_selective_scan_f32_head_range Forward
void mamba2_selective_scan_f32_head_range(const float * state_init, const float * x, const float * dt, const float * a, const float * b, const float * c, const float * d, float * state_out, float * y, int batch, int seq_len, int num_heads, int head_dim, int state_dim, int num_groups, int head_begin, int head_end)

mamba2_selective_scan_f32_parallel_dispatch Forward
void mamba2_selective_scan_f32_parallel_dispatch(const float * state_init, const float * x, const float * dt, const float * a, const float * b, const float * c, const float * d, float * state_out, float * y, int batch, int seq_len, int num_heads, int head_dim, int state_dim, int num_groups)

mamba2_selective_state_update_decode_f32 Forward
void mamba2_selective_state_update_decode_f32(const float * state_in, const float * x, const float * dt, const float * a, const float * b, const float * c, const float * d, float * state_out, float * y, int rows, int num_heads, int head_dim, int state_dim, int num_groups)

mamba2_silu_f32 Forward
float mamba2_silu_f32(float x)

mamba2_softplus_f32 Forward
float mamba2_softplus_f32(float x)

map_forward_to_backward Backward
CKOpType map_forward_to_backward(CKOpType op)

Backward pass / gradient computation

match_special_token Forward
int match_special_token(const CKTrueBPE * bpe, const char * text, int text_len, int pos)

max_k_for_query Forward
int max_k_for_query(int t_q, int T_q, int T_k)

mega_fuse_get_optimal_tiles Forward
void mega_fuse_get_optimal_tiles(int * q_tile, int * kv_tile, int head_dim)

Get optimal tile sizes for current CPU.

mega_fuse_output_proj_residual Forward
void mega_fuse_output_proj_residual(const float * attn_token, const float * wo, const float * bo, const float * residual, float * output, int embed_dim, int aligned_embed_dim, int num_heads, int head_dim, int aligned_head_dim)

mega_fuse_report_stats Forward
void mega_fuse_report_stats(int hidden, int num_layers, int seq_len)

Report memory savings from mega-fusion.

merge_hash Forward
size_t merge_hash(uint64_t key, size_t num_buckets)

merge_key Forward
uint64_t merge_key(int32_t left_id, int32_t right_id)

merge_table_create Forward
CKMergeTable * merge_table_create(size_t num_buckets)

merge_table_free Forward
void merge_table_free(CKMergeTable * table)

merge_table_insert Forward
int merge_table_insert(CKMergeTable * table, const CKBPEMerge * merge)

merge_table_lookup Forward
const CKBPEMerge * merge_table_lookup(const CKMergeTable * table, int32_t left_id, int32_t right_id)

model_align_elems Forward
int model_align_elems(int elems, int elem_bytes, int align_bytes)

model_decode Forward
void model_decode(MODELModel * model, const int * token, int token_index)

model_decode_token Forward
void model_decode_token(MODELModel * model, const int * token, int token_index)

model_forward Forward
void model_forward(MODELModel * model, const int * tokens, int num_tokens)

Forward pass computation

model_forward_prefill_impl Forward
void model_forward_prefill_impl(MODELModel * model, const int * tokens, int num_tokens)

Forward pass computation

model_layer_0_decode Forward
void model_layer_0_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_0_prefill Forward
void model_layer_0_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_10_decode Forward
void model_layer_10_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_10_prefill Forward
void model_layer_10_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_11_decode Forward
void model_layer_11_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_11_prefill Forward
void model_layer_11_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_12_decode Forward
void model_layer_12_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_12_prefill Forward
void model_layer_12_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_13_decode Forward
void model_layer_13_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_13_prefill Forward
void model_layer_13_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_14_decode Forward
void model_layer_14_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_14_prefill Forward
void model_layer_14_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_15_decode Forward
void model_layer_15_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_15_prefill Forward
void model_layer_15_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_16_decode Forward
void model_layer_16_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_16_prefill Forward
void model_layer_16_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_17_decode Forward
void model_layer_17_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_17_prefill Forward
void model_layer_17_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_18_decode Forward
void model_layer_18_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_18_prefill Forward
void model_layer_18_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_19_decode Forward
void model_layer_19_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_19_prefill Forward
void model_layer_19_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_1_decode Forward
void model_layer_1_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_1_prefill Forward
void model_layer_1_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_20_decode Forward
void model_layer_20_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_20_prefill Forward
void model_layer_20_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_21_decode Forward
void model_layer_21_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_21_prefill Forward
void model_layer_21_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_22_decode Forward
void model_layer_22_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_22_prefill Forward
void model_layer_22_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_23_decode Forward
void model_layer_23_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_23_prefill Forward
void model_layer_23_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_2_decode Forward
void model_layer_2_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_2_prefill Forward
void model_layer_2_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_3_decode Forward
void model_layer_3_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_3_prefill Forward
void model_layer_3_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_4_decode Forward
void model_layer_4_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_4_prefill Forward
void model_layer_4_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_5_decode Forward
void model_layer_5_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_5_prefill Forward
void model_layer_5_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_6_decode Forward
void model_layer_6_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_6_prefill Forward
void model_layer_6_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_7_decode Forward
void model_layer_7_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_7_prefill Forward
void model_layer_7_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_8_decode Forward
void model_layer_8_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_8_prefill Forward
void model_layer_8_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_9_decode Forward
void model_layer_9_decode(MODELModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_layer_9_prefill Forward
void model_layer_9_prefill(MODELModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

model_model_allocate Forward
int model_model_allocate(MODELModel * model)

model_model_free Forward
void model_model_free(MODELModel * model)

model_residual_add_token_major Forward
void model_residual_add_token_major(const float * a, const float * b, float * out, int tokens, int aligned_embed_dim)

model_verify_canaries Forward
int model_verify_canaries(MODELModel * model)

moe_accumulate_expert_f32 Forward
void moe_accumulate_expert_f32(float * output, const float * expert_output, float routing_weight, int hidden_dim)

Accumulate expert output: output += routing_weight * expert_output.

moe_relu2_expert_backward_f32 Backward
void moe_relu2_expert_backward_f32(const float * d_output, const float * hidden, const int * indices, const float * routing_weights, const float * expert_up, const float * expert_down, float * d_hidden, float * d_routing_weights, float * d_expert_up, float * d_expert_down, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Backward pass / gradient computation

moe_relu2_expert_forward_f32 Forward
void moe_relu2_expert_forward_f32(const float * hidden, const int * indices, const float * routing_weights, const float * expert_up, const float * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_relu2_expert_forward_q5_0_q5_0 Forward
void moe_relu2_expert_forward_q5_0_q5_0(const float * hidden, const int * indices, const float * routing_weights, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_relu2_expert_forward_q5_0_q8_0 Forward
void moe_relu2_expert_forward_q5_0_q8_0(const float * hidden, const int * indices, const float * routing_weights, const void * expert_up, const void * expert_down, float * output, int rows, int hidden_dim, int intermediate_dim, int n_experts, int top_k)

Forward pass computation

moe_relu2_shared_forward_q5_1_q8_0 Forward
void moe_relu2_shared_forward_q5_1_q8_0(const float * hidden, const float * routed, const void * shared_up, const void * shared_down, float * output, int rows, int hidden_dim, int intermediate_dim)

Forward pass computation

monotonic_ns Forward
uint64_t monotonic_ns(void)

nemotron_group_limited_topk_router_f32 Forward
void nemotron_group_limited_topk_router_f32(const float * scores, const float * correction_bias, int * indices, float * weights, int rows, int n_experts, int top_k, int n_group, int topk_group, int norm_topk_prob, float routed_scaling_factor)

op_name Forward
const char * op_name(CKOpType op)

out_proj_head_major_q5_0_q8_0 Forward
void out_proj_head_major_q5_0_q8_0(const uint8_t * attn_q8, const void * wo, const float * bias, float * output, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

out_proj_head_major_q8_0_q8_0 Forward
void out_proj_head_major_q8_0_q8_0(const uint8_t * attn_q8, const void * wo, const float * bias, float * output, int tokens, int aligned_embed_dim, int num_heads, int aligned_head_dim)

output_append Forward
void output_append(char * buf, size_t * len, const char * text)

output_flush Forward
void output_flush(char * buf, size_t * len)

output_token Forward
void output_token(char * buf, size_t * len, const char * token)

pack_a_panel Forward
void pack_a_panel(const float * A, int lda, float * Ap, int mc, int kc, int mr)

pack_b_panel Forward
void pack_b_panel(const float * B, int ldb, float * Bp, int kc, int nc, int nr)

pack_q4_k_to_packed_meta Forward
void pack_q4_k_to_packed_meta(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_meta_x16 Forward
void pack_q4_k_to_packed_meta_x16(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_meta_x8 Forward
void pack_q4_k_to_packed_meta_x8(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_u8 Forward
void pack_q4_k_to_packed_u8(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_u8_x16 Forward
void pack_q4_k_to_packed_u8_x16(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_vnni_x16 Forward
void pack_q4_k_to_packed_vnni_x16(const void * src, void * dst, int N, int K)

pack_q4_k_to_packed_vnni_x8 Forward
void pack_q4_k_to_packed_vnni_x8(const void * src, void * dst, int N, int K)

parse_args Forward
bool parse_args(int argc, char ** argv, CLIOptions * opt)

parse_env_bool_on Forward
int parse_env_bool_on(const char * name)

parse_eos_ids Forward
bool parse_eos_ids(const char * arg, CLIOptions * opt)

parse_float_field_any Forward
int parse_float_field_any(const char * json, size_t len, const char *const * keys, float * out_value)

parse_float_field_in_range Forward
int parse_float_field_in_range(const char * json, size_t len, const char * key, float * out_value)

parse_int_field Forward
int parse_int_field(const char * json, const char * key, int * out_value)

parse_int_field_any Forward
int parse_int_field_any(const char * json, size_t len, const char *const * keys, int * out_value)

parse_int_field_in_range Forward
int parse_int_field_in_range(const char * json, size_t len, const char * key, int * out_value)

parse_manifest_entry Forward
bool parse_manifest_entry(const char * json, const char * name, size_t * offset, size_t * size)

parse_manifest_int Forward
int parse_manifest_int(const char * json, const char * key)

parse_op Forward
CKOpType parse_op(const char * s)

parse_u64 Forward
unsigned long long parse_u64(const char * s)

patch2im Forward
void patch2im(const float * d_patches, float * d_image, int C, int H, int W, int P)

patch2im: Accumulates gradients from patches back into the image. (Backward pass)

patch2im_bf16 Forward
void patch2im_bf16(const uint16_t * d_patches, uint16_t * d_image, int C, int H, int W, int P)

patch_projection_bf16_pytorch_onednn_conv3d_storage Forward
void patch_projection_bf16_pytorch_onednn_conv3d_storage(const float * input, const void * weights, const float * bias, float * output, int batch, int out_channels, int in_channels, int temporal, int patch_h, int patch_w)

patch_projection_image_bf16_native_storage Forward
void patch_projection_image_bf16_native_storage(const float * image, const void * weights_t0, const void * weights_t1, const float * bias, float * output, int channels, int image_h, int image_w, int patch_size, int out_channels, int merge_size)

patch_projection_image_bf16_pytorch_onednn_conv3d_storage Forward
void patch_projection_image_bf16_pytorch_onednn_conv3d_storage(const float * image, const void * weights_t0, const void * weights_t1, const float * bias, float * output, int channels, int image_h, int image_w, int patch_size, int out_channels, int merge_size)

pcie_bandwidth_gbs Forward
float pcie_bandwidth_gbs(int gen, int width)

plan_size Forward
size_t plan_size(const CKMemPlan * plan, int idx)

pool_new_block Forward
CKPoolBlock * pool_new_block(size_t capacity)

position_embeddings_add Forward
void position_embeddings_add(float * x, const float * position_embd, int num_tokens, int embed_dim, int num_positions)

Add learned absolute position embeddings in-place.

position_embeddings_add_at_offset Forward
void position_embeddings_add_at_offset(float * x, const float * position_embd, int num_tokens, int embed_dim, int num_positions, int start_position)

position_embeddings_add_tiled_2d Forward
void position_embeddings_add_tiled_2d(float * x, const float * position_embd, int grid_h, int grid_w, int embed_dim, int merge_size, int source_grid_size)

position_embeddings_add_tiled_2d_align_corners Forward
void position_embeddings_add_tiled_2d_align_corners(float * x, const float * position_embd, int grid_h, int grid_w, int embed_dim, int merge_size, int source_grid_size)

position_embeddings_add_tiled_2d_align_corners_bf16 Forward
void position_embeddings_add_tiled_2d_align_corners_bf16(float * x, const float * position_embd, int grid_h, int grid_w, int embed_dim, int merge_size, int source_grid_size)

position_embeddings_add_tiled_2d_align_corners_fp32_interp_bf16 Forward
void position_embeddings_add_tiled_2d_align_corners_fp32_interp_bf16(float * x, const float * position_embd, int grid_h, int grid_w, int embed_dim, int merge_size, int source_grid_size)

preprocess_bpe_spaces Forward
int preprocess_bpe_spaces(const char * text, int text_len, char * out, int out_max, CKSpacePrefixStyle style)

preprocess_spm_llama_text Forward
int preprocess_spm_llama_text(const char * text, int text_len, char * out, int out_max, bool add_space_prefix)

preprocess_spm_text Forward
int preprocess_spm_text(const char * text, int text_len, char * out, int out_max, bool add_space_prefix)

preprocess_text Forward
int preprocess_text(const CKTrueBPE * bpe, const char * text, int text_len, char * out, int out_max)

print_banner Forward
void print_banner(void)

print_cpu_info Forward
void print_cpu_info(void)

print_header Forward
void print_header(const char * title)

print_help Forward
void print_help(const char * prog)

print_ok Forward
void print_ok(const char * msg)

print_progress Forward
void print_progress(int token_id, float token_per_sec)

print_section Forward
void print_section(const char * title)

print_tree_item Forward
void print_tree_item(int level, int is_last, const char * fmt, ...)

print_usage Forward
void print_usage(const char * argv0)

print_version Forward
void print_version(void)

print_warning Forward
void print_warning(const char * msg)

process_repl_command Forward
bool process_repl_command(const char * line, CLIOptions * opt, ModelAPI * api)

q4_k_packed_meta_block_size Forward
size_t q4_k_packed_meta_block_size(void)

q4_k_packed_meta_x16_block_size Forward
size_t q4_k_packed_meta_x16_block_size(void)

q4_k_packed_meta_x8_block_size Forward
size_t q4_k_packed_meta_x8_block_size(void)

q4_k_packed_u8_block_size Forward
size_t q4_k_packed_u8_block_size(void)

q4_k_packed_u8_x16_block_size Forward
size_t q4_k_packed_u8_x16_block_size(void)

q4_k_packed_vnni_x16_block_size Forward
size_t q4_k_packed_vnni_x16_block_size(void)

q4_k_packed_vnni_x8_block_size Forward
size_t q4_k_packed_vnni_x8_block_size(void)

q5_k_quant_value Forward
uint8_t q5_k_quant_value(const block_q5_K * block, int subblock, int i)

q_norm_forward Forward
void q_norm_forward(float * q, const float * q_gamma, int num_heads, int num_tokens, int head_dim, float eps)

Forward pass for Gemma4-assistant q-only per-head RMSNorm.

qk_norm_backward Backward
void qk_norm_backward(const float * d_q_out, const float * d_k_out, const float * q_in, const float * k_in, const float * q_gamma, const float * k_gamma, float * d_q_in, float * d_k_in, float * d_q_gamma, float * d_k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Backward pass for per-head QK RMSNorm.

qk_norm_backward_last_isa Backward
int qk_norm_backward_last_isa(void)

Backward pass / gradient computation

qk_norm_compute_rstd Forward
void qk_norm_compute_rstd(const float * input, float * rstd_cache, int rows, int head_dim, float eps)

qk_norm_compute_rstd_scalar Forward
void qk_norm_compute_rstd_scalar(const float * input, float * rstd_cache, int rows, int head_dim, float eps)

qk_norm_forward Forward
void qk_norm_forward(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Per-head RMSNorm on Q and K.

qk_norm_forward_decode_exact Forward
void qk_norm_forward_decode_exact(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_fp64_sum Forward
void qk_norm_forward_fp64_sum(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_llama_production Forward
void qk_norm_forward_llama_production(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_parallel_dispatch Forward
void qk_norm_forward_parallel_dispatch(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_prefill_exact Forward
void qk_norm_forward_prefill_exact(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_pytorch_bf16_storage Forward
void qk_norm_forward_pytorch_bf16_storage(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_forward_qwen4_pytorch_bf16_storage Forward
void qk_norm_forward_qwen4_pytorch_bf16_storage(float * q, float * k, const float * q_gamma, const float * k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)

Forward pass computation

qk_norm_isa_compiled Forward
int qk_norm_isa_compiled(QKNormISA isa)

qk_norm_parse_forced_isa Forward
QKNormISA qk_norm_parse_forced_isa(void)

qk_norm_select_isa Forward
QKNormISA qk_norm_select_isa(void)

qkv_index Forward
size_t qkv_index(int h, int t, int d, int num_tokens, int aligned_head_dim)

qkv_projection_parallel Forward
void qkv_projection_parallel(const void * ln1_q8, const void * WQ, const void * WK, const void * WV, float * q_out, float * k_out, float * v_out, int H, int H_kv, int head_dim, int embed_dim, int num_threads)

Parallel Q/K/V projection for single token decode.

qkv_q8_0_dtype_supported Forward
int qkv_q8_0_dtype_supported(CKDataType dt)

qkv_q8_k_dtype_supported Forward
int qkv_q8_k_dtype_supported(CKDataType dt)

quantize_attn_out_head_major_q8_0 Forward
void quantize_attn_out_head_major_q8_0(const float * attn_out, uint8_t * dst, int tokens, int num_heads, int aligned_head_dim)

quantize_batch_q8_0 Forward
void quantize_batch_q8_0(const float * x, void * y, int num_rows, int k)

Batch quantize FP32 to Q8_0 format (row-major output)

quantize_batch_q8_k Forward
void quantize_batch_q8_k(const float * x, void * y, int num_rows, int k)

Batch quantize FP32 to Q8_K format (row-major output)

quantize_batch_q8_k_4row_nearest_even Forward
void quantize_batch_q8_k_4row_nearest_even(const float * x, void * vy, int num_rows, int k)

quantize_row_q8_0 Forward
void quantize_row_q8_0(const float * x, void * vy, int k)

Quantize FP32 to Q8_0 format (scalar reference)

quantize_row_q8_0_ref_local Forward
void quantize_row_q8_0_ref_local(const float * x, block_q8_0 * y, int k)

quantize_row_q8_1_scalar Forward
void quantize_row_q8_1_scalar(const float * x, block_q8_1 * y, int k)

quantize_row_q8_k Forward
void quantize_row_q8_k(const float * x, void * vy, int k)

quantize_row_q8_k_avx Forward
void quantize_row_q8_k_avx(const float * x, void * vy, int k)

quantize_row_q8_k_avx2 Forward
void quantize_row_q8_k_avx2(const float * x, void * vy, int k)

quantize_row_q8_k_avx512 Forward
void quantize_row_q8_k_avx512(const float * x, void * vy, int k)

quantize_row_q8_k_ref Forward
void quantize_row_q8_k_ref(const float * x, void * vy, int k)

quantize_row_q8_k_sse Forward
void quantize_row_q8_k_sse(const float * x, void * vy, int k)

qwen2_0_5b_decode_align_elems Forward
int qwen2_0_5b_decode_align_elems(int elems, int elem_bytes, int align_bytes)

qwen2_0_5b_decode_decode Forward
void qwen2_0_5b_decode_decode(QWEN2_0_5B_DECODEModel * model, const int * token, int token_index)

qwen2_0_5b_decode_decode_token Forward
void qwen2_0_5b_decode_decode_token(QWEN2_0_5B_DECODEModel * model, const int * token, int token_index)

qwen2_0_5b_decode_forward Forward
void qwen2_0_5b_decode_forward(QWEN2_0_5B_DECODEModel * model, const int * tokens, int num_tokens)

Forward pass computation

qwen2_0_5b_decode_forward_prefill_impl Forward
void qwen2_0_5b_decode_forward_prefill_impl(QWEN2_0_5B_DECODEModel * model, const int * tokens, int num_tokens)

Forward pass computation

qwen2_0_5b_decode_layer_0_decode Forward
void qwen2_0_5b_decode_layer_0_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_0_prefill Forward
void qwen2_0_5b_decode_layer_0_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_10_decode Forward
void qwen2_0_5b_decode_layer_10_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_10_prefill Forward
void qwen2_0_5b_decode_layer_10_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_11_decode Forward
void qwen2_0_5b_decode_layer_11_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_11_prefill Forward
void qwen2_0_5b_decode_layer_11_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_12_decode Forward
void qwen2_0_5b_decode_layer_12_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_12_prefill Forward
void qwen2_0_5b_decode_layer_12_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_13_decode Forward
void qwen2_0_5b_decode_layer_13_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_13_prefill Forward
void qwen2_0_5b_decode_layer_13_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_14_decode Forward
void qwen2_0_5b_decode_layer_14_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_14_prefill Forward
void qwen2_0_5b_decode_layer_14_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_15_decode Forward
void qwen2_0_5b_decode_layer_15_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_15_prefill Forward
void qwen2_0_5b_decode_layer_15_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_16_decode Forward
void qwen2_0_5b_decode_layer_16_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_16_prefill Forward
void qwen2_0_5b_decode_layer_16_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_17_decode Forward
void qwen2_0_5b_decode_layer_17_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_17_prefill Forward
void qwen2_0_5b_decode_layer_17_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_18_decode Forward
void qwen2_0_5b_decode_layer_18_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_18_prefill Forward
void qwen2_0_5b_decode_layer_18_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_19_decode Forward
void qwen2_0_5b_decode_layer_19_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_19_prefill Forward
void qwen2_0_5b_decode_layer_19_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_1_decode Forward
void qwen2_0_5b_decode_layer_1_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_1_prefill Forward
void qwen2_0_5b_decode_layer_1_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_20_decode Forward
void qwen2_0_5b_decode_layer_20_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_20_prefill Forward
void qwen2_0_5b_decode_layer_20_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_21_decode Forward
void qwen2_0_5b_decode_layer_21_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_21_prefill Forward
void qwen2_0_5b_decode_layer_21_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_22_decode Forward
void qwen2_0_5b_decode_layer_22_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_22_prefill Forward
void qwen2_0_5b_decode_layer_22_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_23_decode Forward
void qwen2_0_5b_decode_layer_23_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_23_prefill Forward
void qwen2_0_5b_decode_layer_23_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_2_decode Forward
void qwen2_0_5b_decode_layer_2_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_2_prefill Forward
void qwen2_0_5b_decode_layer_2_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_3_decode Forward
void qwen2_0_5b_decode_layer_3_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_3_prefill Forward
void qwen2_0_5b_decode_layer_3_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_4_decode Forward
void qwen2_0_5b_decode_layer_4_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_4_prefill Forward
void qwen2_0_5b_decode_layer_4_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_5_decode Forward
void qwen2_0_5b_decode_layer_5_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_5_prefill Forward
void qwen2_0_5b_decode_layer_5_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_6_decode Forward
void qwen2_0_5b_decode_layer_6_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_6_prefill Forward
void qwen2_0_5b_decode_layer_6_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_7_decode Forward
void qwen2_0_5b_decode_layer_7_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_7_prefill Forward
void qwen2_0_5b_decode_layer_7_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_8_decode Forward
void qwen2_0_5b_decode_layer_8_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_8_prefill Forward
void qwen2_0_5b_decode_layer_8_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_9_decode Forward
void qwen2_0_5b_decode_layer_9_decode(QWEN2_0_5B_DECODEModel * model, int token_index, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_layer_9_prefill Forward
void qwen2_0_5b_decode_layer_9_prefill(QWEN2_0_5B_DECODEModel * model, int num_tokens, int aligned_embed_dim, int aligned_head_dim, int aligned_intermediate_dim, int aligned_context_window)

qwen2_0_5b_decode_model_allocate Forward
int qwen2_0_5b_decode_model_allocate(QWEN2_0_5B_DECODEModel * model)

qwen2_0_5b_decode_model_free Forward
void qwen2_0_5b_decode_model_free(QWEN2_0_5B_DECODEModel * model)

qwen2_0_5b_decode_residual_add_token_major Forward
void qwen2_0_5b_decode_residual_add_token_major(const float * a, const float * b, float * out, int tokens, int aligned_embed_dim)

qwen2_0_5b_decode_verify_canaries Forward
int qwen2_0_5b_decode_verify_canaries(QWEN2_0_5B_DECODEModel * model)

qwen4_bf16_load Forward
float qwen4_bf16_load(const uint16_t * value)

qwen4_bf16_round Forward
float qwen4_bf16_round(float value)

qwen4_history_token Forward
int32_t qwen4_history_token(const int32_t * tokens, const float * state, int row, int shift, int state_len, int eos_token_id, int position)

qwen4_llama_mul_sum_rows Forward
float qwen4_llama_mul_sum_rows(const float * left, const float * right, int dim)

qwen4_ple_gate_conv_inject_bf16 Forward
void qwen4_ple_gate_conv_inject_bf16(const float * hyper_input, const float * key_projected, const float * value_projected, const float * norm_key_weight, const float * norm_query_weight, const float * norm_conv_weight, const uint16_t * conv_weight, float * hyper_output, float * key_norm_scratch, float * query_norm_scratch, float * gated_scratch, float * conv_norm_scratch, const float * conv_state_in, float * conv_state_out, int rows, int streams, int hidden_dim, int kernel_size, int dilation, float eps)

qwen4_ple_gate_conv_inject_fp16 Forward
void qwen4_ple_gate_conv_inject_fp16(const float * hyper_input, const float * key_projected, const float * value_projected, const float * norm_key_weight, const float * norm_query_weight, const float * norm_conv_weight, const uint16_t * conv_weight, float * hyper_output, float * key_norm_scratch, float * query_norm_scratch, float * gated_scratch, float * conv_norm_scratch, const float * conv_state_in, float * conv_state_out, int rows, int streams, int hidden_dim, int kernel_size, int dilation, float eps)

qwen4_ple_gate_conv_inject_impl Forward
void qwen4_ple_gate_conv_inject_impl(const float * hyper_input, const float * key_projected, const float * value_projected, const float * norm_key_weight, const float * norm_query_weight, const float * norm_conv_weight, const void * conv_weight, float * hyper_output, float * key_norm_scratch, float * query_norm_scratch, float * gated_scratch, float * conv_norm_scratch, const float * conv_state_in, float * conv_state_out, int rows, int streams, int hidden_dim, int kernel_size, int dilation, float eps, int conv_weight_is_fp16, int llama_fp32_arithmetic)

qwen4_ple_gate_conv_inject_llama_fp16 Forward
void qwen4_ple_gate_conv_inject_llama_fp16(const float * hyper_input, const float * key_projected, const float * value_projected, const float * norm_key_weight, const float * norm_query_weight, const float * norm_conv_weight, const uint16_t * conv_weight, float * hyper_output, float * key_norm_scratch, float * query_norm_scratch, float * gated_scratch, float * conv_norm_scratch, const float * conv_state_in, float * conv_state_out, int rows, int streams, int hidden_dim, int kernel_size, int dilation, float eps)

qwen4_ple_ngram_embed_bf16 Forward
void qwen4_ple_ngram_embed_bf16(const int32_t * token_ids, const uint16_t * embedding, const int64_t * layer_multipliers, const int64_t * head_offsets, const int64_t * head_vocab_sizes, float * output, const float * token_state_in, float * token_state_out, int rows, int ngram_size, int heads_per_ngram, int head_dim, int eos_token_id, int position_offset)

qwen4_ple_ngram_embed_impl Forward
void qwen4_ple_ngram_embed_impl(const int32_t * token_ids, const void * embedding, const int64_t * layer_multipliers, const int64_t * head_offsets, const int64_t * head_vocab_sizes, float * output, const float * token_state_in, float * token_state_out, int rows, int ngram_size, int heads_per_ngram, int head_dim, int eos_token_id, int position_offset, int embedding_is_q5_0)

qwen4_ple_ngram_embed_q5_0 Forward
void qwen4_ple_ngram_embed_q5_0(const int32_t * token_ids, const void * embedding, const int64_t * layer_multipliers, const int64_t * head_offsets, const int64_t * head_vocab_sizes, float * output, const float * token_state_in, float * token_state_out, int rows, int ngram_size, int heads_per_ngram, int head_dim, int eos_token_id, int position_offset)

qwen4_pytorch_bf16_dot Forward
float qwen4_pytorch_bf16_dot(const float * left, const float * right, int dim)

qwen4_qsa_index_select_bf16 Forward
void qwen4_qsa_index_select_bf16(const float * projected_qk, const float * index_key_cache_in, const float * q_norm_weight, const float * k_norm_weight, float * selected_indices, float * index_key_cache_out, float * q_norm_scratch, float * pooled_key_scratch, float * block_score_scratch, int32_t * block_index_scratch, int rows, int query_heads, int index_head_dim, int token_budget, int compress_ratio, int rotary_dim, int context_length, int position, float rope_theta, float eps)

read_file_int Forward
int read_file_int(const char * path)

read_file_string Forward
int read_file_string(const char * path, char * buf, size_t buf_size)

read_file_uint64 Forward
uint64_t read_file_uint64(const char * path)

read_floats Forward
int read_floats(FILE * f, float * dst, size_t count)

read_manifest_entry Forward
bool read_manifest_entry(const char * json_path, const char * entry_name, size_t * out_offset, size_t * out_size)

read_prompt_file Forward
char * read_prompt_file(const char * path)

read_u16_le Forward
uint16_t read_u16_le(const uint8_t * p)

read_u32_le Forward
uint32_t read_u32_le(const uint8_t * p)

record_allocation Forward
int record_allocation(void * ptr, size_t len, int was_mmap)

recurrent_ceil_log2 Forward
int recurrent_ceil_log2(int value)

recurrent_conv_backward_extents Backward
int recurrent_conv_backward_extents(int history_len, int num_seqs, int num_tokens, int q_dim, int k_dim, int v_dim, int * channels_out, int * total_len_out, size_t * elements_out)

Backward pass / gradient computation

recurrent_conv_state_update_backward Backward
void recurrent_conv_state_update_backward(const float * d_conv_x, const float * d_state_out, float * d_state_in, float * d_q, float * d_k, float * d_v, int history_len, int num_seqs, int num_tokens, int q_dim, int k_dim, int v_dim)

Backward pass / gradient computation

recurrent_conv_state_update_backward_workspace Backward
void recurrent_conv_state_update_backward_workspace(const float * d_conv_x, const float * d_state_out, float * d_state_in, float * d_q, float * d_k, float * d_v, float * d_conv_total, int history_len, int num_seqs, int num_tokens, int q_dim, int k_dim, int v_dim)

Backward pass / gradient computation

recurrent_conv_state_update_forward Forward
void recurrent_conv_state_update_forward(const float * state_in, const float * q, const float * k, const float * v, float * conv_x, float * state_out, int history_len, int num_seqs, int num_tokens, int q_dim, int k_dim, int v_dim)

Forward pass computation

recurrent_dt_gate_backward Backward
void recurrent_dt_gate_backward(const float * d_gate, const float * alpha, const float * dt_bias, const float * a, float * d_alpha, float * d_dt_bias, float * d_a, int rows, int dim)

Backward pass / gradient computation

recurrent_dt_gate_expanded_forward Forward
void recurrent_dt_gate_expanded_forward(const float * alpha, const float * dt_bias, const float * a, float * gate, int rows, int num_heads, int state_dim)

Forward pass computation

recurrent_dt_gate_forward Forward
void recurrent_dt_gate_forward(const float * alpha, const float * dt_bias, const float * a, float * gate, int rows, int num_heads, int state_dim)

Forward pass computation

recurrent_dt_gate_forward_pytorch_fp32 Forward
void recurrent_dt_gate_forward_pytorch_fp32(const float * alpha, const float * dt_bias, const float * a, float * gate, int rows, int num_heads, int state_dim)

Forward pass computation

recurrent_l2_norm_rows_backward_one Backward
void recurrent_l2_norm_rows_backward_one(const float * d_out, const float * x, float * d_x, int rows, int dim, int head_dim, float eps)

Backward pass / gradient computation

recurrent_l2_norm_rows_forward_one Forward
void recurrent_l2_norm_rows_forward_one(float * x, int rows, int dim, int head_dim, float eps)

Forward pass computation

recurrent_norm_gate_backward Backward
void recurrent_norm_gate_backward(const float * d_out, const float * x, const float * gate, const float * weight, float * d_x, float * d_gate, float * d_weight, int rows, int num_heads, int head_dim, float eps)

Backward pass / gradient computation

recurrent_norm_gate_forward Forward
void recurrent_norm_gate_forward(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

Forward pass computation

recurrent_norm_gate_llama_avx2_forward Forward
void recurrent_norm_gate_llama_avx2_forward(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

Forward pass computation

recurrent_norm_gate_pytorch_bf16_storage Forward
void recurrent_norm_gate_pytorch_bf16_storage(const float * x, const float * gate, const float * weight, float * out, int rows, int num_heads, int head_dim, float eps)

recurrent_pytorch_bf16_l2_rows Forward
void recurrent_pytorch_bf16_l2_rows(float * x, int rows, int dim, int expanded_heads, int head_dim, float eps)

recurrent_pytorch_bf16_square_sum Forward
float recurrent_pytorch_bf16_square_sum(const float * x, int dim)

recurrent_pytorch_fp32_l2_rows Forward
void recurrent_pytorch_fp32_l2_rows(float * x, int rows, int dim, int head_dim, float eps)

recurrent_pytorch_fp32_square_sum Forward
float recurrent_pytorch_fp32_square_sum(const float * x, int dim)

recurrent_qk_l2_norm_backward Backward
void recurrent_qk_l2_norm_backward(const float * d_q_out, const float * d_k_out, const float * q, const float * k, float * d_q, float * d_k, int rows, int q_dim, int k_dim, int head_dim, float eps)

Backward pass / gradient computation

recurrent_qk_l2_norm_forward Forward
void recurrent_qk_l2_norm_forward(float * q, float * k, int rows, int q_dim, int k_dim, int head_dim, float eps)

Forward pass computation

recurrent_qk_l2_norm_pytorch_bf16_storage Forward
void recurrent_qk_l2_norm_pytorch_bf16_storage(float * q, float * k, int rows, int q_dim, int k_dim, int expanded_heads, int head_dim, float eps)

recurrent_qk_l2_norm_pytorch_fp32_output Forward
void recurrent_qk_l2_norm_pytorch_fp32_output(float * q, float * k, int rows, int q_dim, int k_dim, int head_dim, float eps)

recurrent_silu_backward Backward
void recurrent_silu_backward(const float * d_out, const float * x, float * d_x, int rows, int dim)

Backward pass / gradient computation

recurrent_silu_forward Forward
void recurrent_silu_forward(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_silu_forward_ggml Forward
void recurrent_silu_forward_ggml(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_silu_forward_pytorch_bf16_input_fp32_output Forward
void recurrent_silu_forward_pytorch_bf16_input_fp32_output(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_silu_forward_pytorch_bf16_storage Forward
void recurrent_silu_forward_pytorch_bf16_storage(const float * x, float * out, int rows, int dim)

Forward pass computation

recurrent_softplus Forward
float recurrent_softplus(float x)

recurrent_split_conv_qkv_backward Backward
void recurrent_split_conv_qkv_backward(const float * d_q, const float * d_k, const float * d_v, float * d_packed_qkv, int rows, int q_dim, int k_dim, int v_dim)

Backward pass / gradient computation

recurrent_split_conv_qkv_forward Forward
void recurrent_split_conv_qkv_forward(const float * packed_qkv, float * q, float * k, float * v, int rows, int q_dim, int k_dim, int v_dim)

Forward pass computation

recurrent_split_qkv_backward Backward
void recurrent_split_qkv_backward(const float * d_q, const float * d_k, const float * d_v, float * d_packed_qkv, int rows, int q_dim, int k_dim, int v_dim)

Backward pass / gradient computation

recurrent_split_qkv_forward Forward
void recurrent_split_qkv_forward(const float * packed_qkv, float * q, float * k, float * v, int rows, int q_dim, int k_dim, int v_dim)

Forward pass computation

reflect_index Forward
int reflect_index(int index, int length)

relu2_backward Backward
void relu2_backward(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

relu2_forward Forward
void relu2_forward(const float * input, float * output, size_t n)

Forward pass computation

relu_backward Backward
void relu_backward(const float * input, const float * d_output, float * d_input, size_t n)

Backward pass / gradient computation

relu_backward_bf16 Backward
void relu_backward_bf16(const uint16_t * input, const uint16_t * d_output, uint16_t * d_input, size_t n)

Backward pass / gradient computation

relu_forward Forward
void relu_forward(const float * input, float * output, size_t n)

Forward pass computation

relu_forward_bf16 Forward
void relu_forward_bf16(const uint16_t * input, uint16_t * output, size_t n)

Forward pass computation

relu_forward_inplace Forward
void relu_forward_inplace(float * data, size_t n)

Forward pass computation

relu_forward_inplace_bf16 Forward
void relu_forward_inplace_bf16(uint16_t * data, size_t n)

Forward pass computation

residual_add Forward
void residual_add(float * residual, float * addend, int n)

residual_add_parallel Forward
void residual_add_parallel(const float * a, const float * b, float * out, int n, int ith, int nth)

Single-token decode with parallel SIMD kernels.

resolve_dim Forward
size_t resolve_dim(const CKModelConfig * cfg, const CKIRV2AlignInfo * align, CKDimKind kind, int tokens_override)

resolve_shape_elems Forward
size_t resolve_shape_elems(const CKModelConfig * cfg, const CKIRV2AlignInfo * align, const CKDimToken * shape, int tokens_override)

resolve_symbol Forward
bool resolve_symbol(void * handle, const char * name, void ** out_ptr, bool required)

rowwise_bias_add Forward
void rowwise_bias_add(float * x, const float * bias, int rows, int dim)

run_benchmark Forward
void run_benchmark(void * model, int num_tokens)

run_command Forward
int run_command(const char * cmd, char * output, size_t output_size)

run_generation_test Forward
void run_generation_test(void * model, int num_tokens)

run_inference Forward
int run_inference(const char * bump_path, const char * manifest_path, const char * tokenizer_path, const char * prompt, int max_tokens, float temperature, int topk)

run_prompt Forward
int run_prompt(ModelAPI * api, CKTrueBPE * tokenizer, CLIOptions * opt, const char * input)

sample_argmax Forward
int sample_argmax(const float * logits, int vocab_size)

sample_token Forward
int sample_token(float * logits, int vocab_size, float temp, int top_k)

sample_top_p Forward
int sample_top_p(float * logits, int vocab_size, float temperature, float top_p)

sample_topk Forward
int sample_topk(float * probs, int vocab_size, int topk)

scal_copy_f32 Forward
void scal_copy_f32(float * y, const float * x, float alpha, int n)

Scaled copy: y = alpha * x.

score_index Forward
size_t score_index(int h, int i, int j, int aligned_context_window)

sgd_momentum_update_bf16 Forward
void sgd_momentum_update_bf16(const uint16_t * grad, uint16_t * weight, float * velocity, size_t numel, float lr, float momentum, float weight_decay)

SGD with momentum (bf16 weights/gradients)

sgd_momentum_update_f32 Forward
void sgd_momentum_update_f32(const float * grad, float * weight, float * velocity, size_t numel, float lr, float momentum, float weight_decay)

SGD with momentum optimizer update (fp32 version)

silu Forward
void silu(float * x, int n)

silu_prefill Forward
float silu_prefill(float x)

silu_scalar Forward
float silu_scalar(float x)

simd_strcmp Forward
int simd_strcmp(const char * s1, const char * s2)

simple_embedding Forward
void simple_embedding(const int32_t * tokens, int num_tokens, const float * weight, float * output, int vocab_size, int embed_dim)

spatial_average_pool_contiguous Forward
void spatial_average_pool_contiguous(const float * input, float * output, int grid_h, int grid_w, int embed_dim, int merge_size)

spatial_merge_2x2 Forward
void spatial_merge_2x2(const float * input, float * output, int grid_h, int grid_w, int embed_dim)

Merge 2x2 neighboring tokens into a single wider token.

spatial_merge_contiguous_tiled Forward
void spatial_merge_contiguous_tiled(const float * input, float * output, int grid_h, int grid_w, int embed_dim, int merge_size)

speculative_commit_one_i32 Forward
void speculative_commit_one_i32(int accepted, int verified_token, int * token_buffer, int * token_count, int max_tokens, int * target_position, int * draft_position, int * accepted_count, int * rejected_count)

Commit one verified speculative token and update decode counters.

speculative_verify_greedy_f32 Forward
void speculative_verify_greedy_f32(const float * target_logits, int vocab_size, int draft_token, int * accepted, int * verified_token)

Greedy one-token speculative verification.

split_q_gate_backward Backward
void split_q_gate_backward(const float * d_q, const float * d_gate, float * d_packed_qg, int rows, int q_dim, int gate_dim, int group_dim)

Backward pass / gradient computation

split_q_gate_forward Forward
void split_q_gate_forward(const float * packed_qg, float * q, float * gate, int rows, int q_dim, int gate_dim, int group_dim)

Forward pass computation

split_qkv_packed_head_major_forward Forward
void split_qkv_packed_head_major_forward(const float * packed_qkv, float * q, float * k, float * v, int rows, int q_dim, int k_dim, int v_dim, int num_heads, int num_kv_heads)

Forward pass computation

spm_build_byte_lookup Forward
void spm_build_byte_lookup(CKTokenizer * tok, const char * strings, const int32_t * offsets, int vocab_size)

spm_count_unknown_run Forward
int spm_count_unknown_run(const CKTokenizer * tok, const char * text, int text_len, size_t pos)

spm_encode_byte_fallback Forward
int spm_encode_byte_fallback(const CKTokenizer * tok, const char * text, int text_len, int32_t * ids, int max_ids)

spm_find_candidates_at_pos Forward
int spm_find_candidates_at_pos(const CKTokenizer * tok, const char * text, int text_len, size_t pos, int32_t * candidates, int max_candidates)

spm_find_special_token_at_pos Forward
int32_t spm_find_special_token_at_pos(const CKTokenizer * tok, const char * text, int text_len, size_t pos, size_t * match_len)

spm_get_byte_token Forward
int32_t spm_get_byte_token(const CKTokenizer * tok, unsigned char byte_val)

spm_is_byte_token Forward
bool spm_is_byte_token(const CKTokenizer * tok, int32_t token_id)

spm_llama_resegment_node Forward
int spm_llama_resegment_node(const CKTokenizer * tok, const SpmLlamaNode * nodes, int node_id, int32_t * ids, int max_ids, int out_idx)

spm_token_allowed_in_dp Forward
bool spm_token_allowed_in_dp(const CKTokenizer * tok, int32_t token_id)

spm_token_is_byte_format Forward
bool spm_token_is_byte_format(const char * token)

ssm_conv1d_backward Backward
void ssm_conv1d_backward(const float * d_out, const float * conv_x, const float * kernel, float * d_conv_x, float * d_kernel, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Backward pass / gradient computation

ssm_conv1d_backward_ref Backward
void ssm_conv1d_backward_ref(const float * d_out, const float * conv_x, const float * kernel, float * d_conv_x, float * d_kernel, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Backward pass / gradient computation

ssm_conv1d_forward Forward
void ssm_conv1d_forward(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

ssm_conv1d_forward_llama_fma Forward
void ssm_conv1d_forward_llama_fma(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

ssm_conv1d_forward_llama_production Forward
void ssm_conv1d_forward_llama_production(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

ssm_conv1d_forward_llama_production_serial Forward
void ssm_conv1d_forward_llama_production_serial(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

ssm_conv1d_forward_pytorch_bf16_storage Forward
void ssm_conv1d_forward_pytorch_bf16_storage(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

ssm_conv1d_forward_ref Forward
void ssm_conv1d_forward_ref(const float * conv_x, const float * kernel, float * out, int kernel_size, int num_channels, int num_tokens, int num_seqs)

Forward pass computation

starts_with Forward
int starts_with(const char * s, const char * prefix)

tile_order_index_2d Forward
int tile_order_index_2d(int linear_idx, int grid_h, int grid_w, int merge_size)

tile_order_linear_index_2d Forward
int tile_order_linear_index_2d(int y, int x, int grid_h, int grid_w, int merge_size)

token_has_gpt2_bytes Forward
bool token_has_gpt2_bytes(const char * token)

token_list_append Forward
int token_list_append(CKBPETokenList * list, const char * str, size_t len, int32_t id)

token_list_clear Forward
void token_list_clear(CKBPETokenList * list)

token_list_create Forward
CKBPETokenList * token_list_create(size_t initial_capacity)

token_list_free Forward
void token_list_free(CKBPETokenList * list)

token_list_merge_at Forward
int token_list_merge_at(CKBPETokenList * list, size_t pos, const char * merged_str, size_t merged_len, int32_t merged_id)

tokenize Forward
int32_t * tokenize(const char * text, int * num_tokens)

topk_batched_f32 Forward
void topk_batched_f32(const float * scores, int num_tokens, int n_experts, int k, int * indices, float * weights)

Batched top-K selection for multiple tokens.

topk_f32 Forward
void topk_f32(const float * scores, int n, int k, int * indices, float * values)

Find top-K indices and values from a score vector.

topology_discover Forward
int topology_discover(SystemTopology * topo)

topology_discover_affinity Forward
int topology_discover_affinity(AffinityInfo * aff)

topology_discover_cache
int topology_discover_cache(CacheTopology * cache)

topology_discover_cpu Forward
int topology_discover_cpu(CPUInfo * cpu)

topology_discover_memory Forward
int topology_discover_memory(MemoryInfo * mem)

topology_discover_network Forward
int topology_discover_network(NetworkTopology * net)

topology_discover_numa Forward
int topology_discover_numa(NUMATopology * numa)

topology_discover_pcie Forward
int topology_discover_pcie(PCIeTopology * pcie)

topology_estimate_channels_from_bandwidth Forward
int topology_estimate_channels_from_bandwidth(float measured_bw_gbs, int memory_speed_mhz, const char * memory_type)

topology_estimate_memory_bandwidth Forward
float topology_estimate_memory_bandwidth(const MemoryInfo * mem)

topology_estimate_network_training_time Forward
float topology_estimate_network_training_time(const NetworkTopology * net, uint64_t model_size_mb)

topology_generate_recommendations Forward
int topology_generate_recommendations(const SystemTopology * topo, RecommendationList * recs)

topology_measure_memory_bandwidth Forward
float topology_measure_memory_bandwidth(void)

topology_measure_memory_bandwidth_ex Forward
float topology_measure_memory_bandwidth_ex(int * numa_node_out, int * num_threads_out)

topology_print_affinity Forward
void topology_print_affinity(const AffinityInfo * aff)

topology_print_cache
void topology_print_cache(const CacheTopology * cache, int logical_cores)

topology_print_cpu Forward
void topology_print_cpu(const CPUInfo * cpu)

topology_print_distributed_potential Forward
void topology_print_distributed_potential(const SystemTopology * topo)

topology_print_memory Forward
void topology_print_memory(const MemoryInfo * mem)

topology_print_network Forward
void topology_print_network(const NetworkTopology * net)

topology_print_numa Forward
void topology_print_numa(const NUMATopology * numa, int sockets)

topology_print_pcie Forward
void topology_print_pcie(const PCIeTopology * pcie)

topology_print_recommendations Forward
void topology_print_recommendations(const RecommendationList * recs)

topology_print_summary Forward
void topology_print_summary(const SystemTopology * topo)

trim_string Forward
void trim_string(char * str)

unpack_q4_k_scales Forward
void unpack_q4_k_scales(const uint8_t * scales, uint8_t * sc, uint8_t * m)

Unpack Q4_K sub-block scales and mins.

unpack_q5_k_scales Forward
void unpack_q5_k_scales(const uint8_t * scales, uint8_t * sc, uint8_t * m)

utf8_char_len Forward
int utf8_char_len(unsigned char c)

utf8_len Forward
int utf8_len(unsigned char c)

v6_prefill Forward
void v6_prefill(const float * embed_weight, const int32_t * tokens, int num_tokens, float * logits)

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

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

vec_dot_q5_0_q8_0 Forward
void vec_dot_q5_0_q8_0(int n, float * s, const void * vx, const void * vy)

Auto-dispatch quantized dot product Q5_0 x Q8_0.

vec_dot_q5_0_q8_0_ref Forward
void vec_dot_q5_0_q8_0_ref(int n, float * s, const void * vx, const void * vy)

Quantized dot product: Q5_0 weights x Q8_0 input (scalar reference)

vec_dot_q6_k_q8_k Forward
void vec_dot_q6_k_q8_k(int n, float * s, const void * vx, const void * vy)

Q6_K x Q8_K dot product (single row)

vec_dot_q8_0_q8_0 Forward
void vec_dot_q8_0_q8_0(int n, float * s, const void * vx, const void * vy)

Auto-dispatch quantized dot product Q8_0 x Q8_0.

vec_dot_q8_0_q8_0_ref Forward
void vec_dot_q8_0_q8_0_ref(int n, float * s, const void * vx, const void * vy)

Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)

vec_scale_parallel Forward
void vec_scale_parallel(float * y, float scale, int n, int ith, int nth)

vec_zero_parallel Forward
void vec_zero_parallel(float * y, int n, int ith, int nth)

vision_position_ids_2d_merge Forward
void vision_position_ids_2d_merge(int32_t * positions, int grid_h, int grid_w, int merge_size)

Build merged 2D vision position IDs in the layout expected by vision M-RoPE.

weighted_sum_f32 Forward
void weighted_sum_f32(float * y, const float ** vectors, const float * weights, int k, int n)

Weighted sum of k vectors: y = sum_i(weights[i] * vectors[i])

worker_main Forward
void * worker_main(void * arg)

yarn_correction_dim Forward
float yarn_correction_dim(float rotations, int rotary_dim, float freq_base, int original_context)

yarn_mscale Forward
float yarn_mscale(float factor, float scale)

zero_gradients_bf16 Forward
void zero_gradients_bf16(uint16_t * grad, size_t numel)

Zero out gradient buffer (bf16)

zero_gradients_f32 Forward
void zero_gradients_f32(float * grad, size_t numel)

Zero out gradient buffer (fp32)

zero_row_f32 Forward
void zero_row_f32(float * row, int cols)

Usage Example

Forward + Backward Pass

// Forward pass
rmsnorm_forward(input, gamma, norm_out, rstd_cache, tokens, d_model, d_model, eps);
attention_forward_causal_head_major_gqa(q, k, v, scores, attn_out, heads, kv_heads, tokens, head_dim, head_dim, ctx_len);
swiglu_forward(mlp_in, mlp_out, tokens, hidden_dim);

// Backward pass (reverse order)
swiglu_backward(mlp_in, d_mlp_out, d_mlp_in, tokens, hidden_dim);
attention_backward_causal_head_major_gqa(d_attn_out, q, k, v, scores, d_q, d_k, d_v, d_scores, heads, kv_heads, tokens, head_dim, head_dim, ctx_len);
rmsnorm_backward(d_norm_out, input, gamma, rstd_cache, d_input, d_gamma, tokens, d_model, d_model);
Image
100% | |
Scroll to zoom | Drag to pan | W/H to fit | 0 to reset | ESC to close