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.
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);