1#ifndef CKERNEL_ENGINE_H
2#define CKERNEL_ENGINE_H
28 void (*sgemm)(
int M,
int N,
int K,
29 const float *A,
int lda,
30 const float *B,
int ldb,
69 const float *norm_weight,
70 const uint16_t *mix_down_weight,
71 const uint16_t *mix_up_weight,
72 const uint16_t *inject_weight,
74 float *injection_output,
75 float *normalized_scratch,
76 float *dynamic_scratch,
85 const float *norm_weight,
86 const void *mix_down_weight,
87 const void *mix_up_weight,
88 const void *inject_weight,
90 float *injection_output,
91 float *normalized_scratch,
92 float *dynamic_scratch,
101 const float *norm_weight,
102 const void *mix_down_weight,
103 const void *mix_up_weight,
104 const void *inject_weight,
106 float *injection_output,
107 float *normalized_scratch,
108 float *dynamic_scratch,
117 const float *block_output,
118 const float *injection_weight,
124 const float *block_output,
125 const float *injection_weight,
132 const uint16_t *embedding,
133 const int64_t *layer_multipliers,
134 const int64_t *head_offsets,
135 const int64_t *head_vocab_sizes,
137 const float *token_state_in,
138 float *token_state_out,
146 const void *embedding,
147 const int64_t *layer_multipliers,
148 const int64_t *head_offsets,
149 const int64_t *head_vocab_sizes,
151 const float *token_state_in,
152 float *token_state_out,
160 const float *hyper_input,
161 const float *key_projected,
162 const float *value_projected,
163 const float *norm_key_weight,
164 const float *norm_query_weight,
165 const float *norm_conv_weight,
166 const uint16_t *conv_weight,
168 float *key_norm_scratch,
169 float *query_norm_scratch,
170 float *gated_scratch,
171 float *conv_norm_scratch,
172 const float *conv_state_in,
173 float *conv_state_out,
181 const float *hyper_input,
182 const float *key_projected,
183 const float *value_projected,
184 const float *norm_key_weight,
185 const float *norm_query_weight,
186 const float *norm_conv_weight,
187 const uint16_t *conv_weight,
189 float *key_norm_scratch,
190 float *query_norm_scratch,
191 float *gated_scratch,
192 float *conv_norm_scratch,
193 const float *conv_state_in,
194 float *conv_state_out,
202 const float *hyper_input,
203 const float *key_projected,
204 const float *value_projected,
205 const float *norm_key_weight,
206 const float *norm_query_weight,
207 const float *norm_conv_weight,
208 const uint16_t *conv_weight,
210 float *key_norm_scratch,
211 float *query_norm_scratch,
212 float *gated_scratch,
213 float *conv_norm_scratch,
214 const float *conv_state_in,
215 float *conv_state_out,
223 const float *projected_qk,
const float *index_key_cache_in,
224 const float *q_norm_weight,
const float *k_norm_weight,
225 float *selected_indices,
float *index_key_cache_out,
226 float *q_norm_scratch,
float *pooled_key_scratch,
227 float *block_score_scratch, int32_t *block_index_scratch,
228 int rows,
int query_heads,
int index_head_dim,
int token_budget,
229 int compress_ratio,
int rotary_dim,
int context_length,
230 int position,
float rope_theta,
float eps);
236 int aligned_embed_dim);
242 int aligned_embed_dim);
248 int aligned_embed_dim);
253 int aligned_embed_dim);
260 int M,
int N,
int K);
265 int M,
int N,
int K);
271 int M,
int N,
int K);
277 int M,
int N,
int K);
283 int M,
int N,
int K);
289 int M,
int N,
int K);
305 int M,
int N,
int K);
348 float *
const *weights,
349 float *
const *m_states,
350 float *
const *v_states,
351 const size_t *numels,
363 const float *
const *srcs,
364 const size_t *numels,
369 const size_t *numels,
375 const uint16_t *bias,
377 int M,
int N,
int K);
381 const uint16_t *input,
382 const uint16_t *weight,
394 int M,
int N,
int K);
404 int M,
int N,
int K);
409 const float *input_min,
410 const float *input_max,
411 const float *output_min,
412 const float *output_max,
414 int M,
int N,
int K);
437 int M,
int N,
int K);
443 int row_begin,
int row_end);
448 int M,
int N,
int K);
453 int M,
int N,
int K);
459 int row_begin,
int row_end);
464 int M,
int N,
int K);
470 int M,
int N,
int K);
475 int M,
int N,
int K);
482 size_t a_bf16_bytes);
487 int M,
int N,
int K);
489 const float *A,
const void *B,
const float *bias,
float *
C,
490 int M,
int N,
int K, uint16_t *a_bf16,
size_t a_bf16_bytes);
495 int M,
int N,
int K);
497 const float *A,
const void *B,
const float *bias,
float *
C,
498 int M,
int N,
int K);
500 const float *input,
const void *weights,
const float *bias,
float *output,
501 int batch,
int out_channels,
int in_channels,
int temporal,
502 int patch_h,
int patch_w);
504 const float *image,
const void *weights_t0,
const void *weights_t1,
505 const float *bias,
float *output,
int channels,
int image_h,
int image_w,
506 int patch_size,
int out_channels,
int merge_size);
509 const float *image,
const void *weights_t0,
const void *weights_t1,
510 const float *bias,
float *output,
int channels,
int image_h,
int image_w,
511 int patch_size,
int out_channels,
int merge_size);
515 const float *q,
const float *k,
const float *v,
float *output,
516 int num_heads,
int num_kv_heads,
int num_tokens,
517 int head_dim,
int aligned_head_dim,
int kv_stride_tokens);
519 const float *q,
const float *k,
const float *v,
float *output,
520 int num_heads,
int num_kv_heads,
int num_tokens,
521 int head_dim,
int aligned_head_dim,
int kv_stride_tokens);
548 int M,
int N,
int K);
554 int M,
int N,
int K);
566 int M,
int N,
int K);
572 int M,
int N,
int K);
575void gemm_nt_q4_0(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
576void gemm_nt_q4_1(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
577void gemm_nt_q5_0(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
578void gemm_nt_q5_1(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
579void gemm_nt_q5_1_q8_1(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
580void 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);
581void gemm_nt_q5_k(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
582void gemm_nt_q5_k_q8_k(
const void *A_q8,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
583void gemm_nt_q8_0(
const float *A,
const void *B,
const float *bias,
float *
C,
int M,
int N,
int K);
587void gemv_q4_0(
float *y,
const void *W,
const float *x,
int M,
int K);
588void gemv_q5_0(
float *y,
const void *W,
const float *x,
int M,
int K);
589void gemv_q5_1(
float *y,
const void *W,
const float *x,
int M,
int K);
590void gemv_q5_1_q8_1(
float *y,
const void *W,
const float *x,
int M,
int K);
592void gemv_q5_k(
float *y,
const void *W,
const float *x,
int M,
int K);
593void gemv_q5_k_q8_k(
float *y,
const void *W,
const void *x_q8,
int M,
int K);
594void gemv_q8_0(
float *y,
const void *W,
const float *x,
int M,
int K);
600 int M,
int K,
int ith,
int nth);
602 int M,
int K,
int ith,
int nth);
625void gemv_q5_0_q8_0(
float *y,
const void *W,
const void *x_q8,
int M,
int K);
628void gemv_q8_0_q8_0(
float *y,
const void *W,
const void *x_q8,
int M,
int K);
632 const float *bias,
int M,
int K);
636 const float *bias,
int M,
int K);
644 int num_rows,
int k);
677 int M,
int N,
int K);
683 int M,
int N,
int K);
697 int M,
int K,
int ith,
int nth);
699 int M,
int K,
int ith,
int nth);
704 int M,
int N,
int K);
710 int M,
int N,
int K);
732 int M,
int N,
int K);
738 int M,
int N,
int K);
744 int M,
int N,
int K);
749 int M,
int N,
int K,
int ldc);
758 int M,
int N,
int K);
764 int M,
int N,
int K);
771 int M,
int N,
int K);
777 int M,
int N,
int K);
801 int tokens,
int heads,
int head_dim);
803 int heads,
int tokens,
int head_dim);
812 int M,
int N,
int K);
818 int M,
int N,
int K);
824 int M,
int N,
int K);
831 int M,
int N,
int K);
837 int M,
int N,
int K);
843 int M,
int N,
int K);
854 int M,
int N,
int K);
949 const float *Wq,
const float *Bq,
950 const float *Wk,
const float *Bk,
951 const float *Wv,
const float *Bv,
957 int aligned_embed_dim,
961 int aligned_head_dim,
962 int kv_stride_tokens,
974 const void *Wq,
const float *Bq,
CKDataType wq_dt,
975 const void *Wk,
const float *Bk,
CKDataType wk_dt,
976 const void *Wv,
const float *Bv,
CKDataType wv_dt,
982 int aligned_embed_dim,
986 int aligned_head_dim,
987 int kv_stride_tokens,
1023 const float *W_gate,
1025 const float *W_down,
1037 const float *W_gate,
1039 const float *W_down,
1040 const float *B_gate,
1042 const float *B_down,
1069 int aligned_embed_dim,
1070 int intermediate_dim,
1071 int aligned_intermediate_dim,
1076 int aligned_intermediate_dim);
1085 int M,
int N,
int K,
1092 int M,
int N,
int K);
1098 int M,
int N,
int K);
1105 int M,
int N,
int K);
1114 int tokens,
int d_model,
int aligned_embed_dim,
1118 const float *__restrict gamma,
1119 const float *__restrict beta,
1120 float *__restrict output_slice_base,
1121 float *__restrict mean_cache_slice,
1122 float *__restrict rstd_cache_slice,
1123 int num_tokens_in_slice,
1125 int aligned_embed_dim,
1130 const float *__restrict gamma,
1131 const float *__restrict beta,
1132 uint16_t *__restrict output_slice_base,
1133 float *__restrict mean_cache_slice,
1134 float *__restrict rstd_cache_slice,
1135 int num_tokens_in_slice,
1137 int aligned_embed_dim,
1139 float *scratch_input,
1140 float *scratch_output);
1143 const float *__restrict gamma,
1144 const float *__restrict beta,
1145 float *__restrict output_slice_base,
1146 float *__restrict mean_cache_slice,
1147 float *__restrict rstd_cache_slice,
1148 int num_tokens_in_slice,
1154 const float *__restrict gamma,
1155 const float *__restrict beta,
1156 uint16_t *__restrict output_slice_base,
1157 float *__restrict mean_cache_slice,
1158 float *__restrict rstd_cache_slice,
1159 int num_tokens_in_slice,
1162 float *scratch_input,
1163 float *scratch_output);
1171 int tokens,
int d_model,
float eps);
1178 int tokens,
int d_model,
float eps);
1185 int tokens,
int d_model,
float eps);
1195 int tokens,
int d_model,
int aligned_embed_dim);
1199 const uint16_t *input,
1206 int tokens,
int d_model,
int aligned_embed_dim,
1207 float *scratch_d_output,
1208 float *scratch_input,
1209 float *scratch_d_input);
1216 int dst_feature_offset);
1219 const float *branch_input,
1223 int branch_slice_dim,
1224 int num_branch_slices);
1233 int aligned_embed_dim,
1242 int aligned_embed_dim,
1250 int aligned_embed_dim,
1258 int aligned_embed_dim,
1266 int aligned_embed_dim,
1283 int aligned_embed_dim,
1291 int aligned_embed_dim,
1308 int aligned_embed_dim,
1315 int aligned_embed_dim,
1335 const float *rstd_cache,
1340 int aligned_embed_dim);
1345 const float *q_gamma,
1346 const float *k_gamma,
1354 const float *q_gamma,
1355 const float *k_gamma,
1363 const float *q_gamma,
1364 const float *k_gamma,
1372 const float *q_gamma,
1373 const float *k_gamma,
1381 const float *q_gamma,
1382 const float *k_gamma,
1390 const float *q_gamma,
1391 const float *k_gamma,
1398 const float *q_gamma,
const float *k_gamma,
1399 int num_heads,
int num_kv_heads,
1400 int num_tokens,
int head_dim,
float eps);
1402 const float *q_gamma,
const float *k_gamma,
1403 int num_heads,
int num_kv_heads,
1404 int num_tokens,
int head_dim,
float eps);
1406 const float *q_gamma,
1414 const float *hidden,
1415 const int32_t *token_ids,
1416 const void *per_layer_token_emb,
1417 const uint16_t *per_layer_model_proj,
1418 const float *per_layer_proj_norm,
1427 const float *hidden,
1428 const int32_t *token_ids,
1429 const uint16_t *per_layer_token_emb,
1430 const uint16_t *per_layer_model_proj,
1431 const float *per_layer_proj_norm,
1440 const float *per_layer_input,
1441 const float *inp_gate,
1443 const float *post_norm,
1444 const float *out_scale,
1466 const float *d_k_out,
1469 const float *q_gamma,
1470 const float *k_gamma,
1490 int aligned_embed_dim,
1494 const uint16_t *input,
1496 const float *rstd_cache,
1501 int aligned_embed_dim);
1510 int aligned_embed_dim,
1512 float *scratch_input,
1513 float *scratch_output);
1517 const int8_t *input,
1519 const float *rstd_cache,
1524 int aligned_embed_dim,
1525 float *scratch_d_output,
1526 float *scratch_input,
1527 float *scratch_d_input);
1536 int aligned_embed_dim,
1538 float *scratch_input,
1539 float *scratch_output);
1543 const uint8_t *input,
1545 const float *rstd_cache,
1550 int aligned_embed_dim,
1551 float *scratch_d_output,
1552 float *scratch_input,
1553 float *scratch_d_input);
1572 const float *d_output,
1579 const float *d_output,
1584 const float *d_output,
1592 const uint16_t *d_output,
1595 float *scratch_input,
1596 float *scratch_d_output,
1597 float *scratch_d_input);
1599 const uint16_t *d_output,
1602 float *scratch_input,
1603 float *scratch_d_output,
1604 float *scratch_d_input);
1612void geglu_forward_bf16(
const uint16_t *x, uint16_t *out,
int tokens,
int dim,
float *scratch);
1619 const uint16_t *d_out,
1625 void relu_forward(
const float *input,
float *output,
size_t n);
1628 const float *d_output,
1631 void relu2_forward(
const float *input,
float *output,
size_t n);
1633 const float *d_output,
1640 const uint16_t *d_output,
1648 int aligned_context_window);
1655 int aligned_context_window);
1658 const float *weights,
1661 int aligned_context_window);
1667 int aligned_context_window,
1672 const uint16_t *weights,
1675 int aligned_context_window,
1676 float *scratch_d_scores,
1677 float *scratch_weights);
1691 int aligned_head_dim,
1692 int aligned_context_window);
1703 int aligned_head_dim,
1704 int aligned_context_window);
1716 int aligned_head_dim,
1717 int aligned_context_window);
1729 int aligned_head_dim,
1730 int aligned_context_window);
1742 int aligned_head_dim,
1743 int aligned_context_window,
1761 int aligned_head_dim);
1772 int aligned_head_dim);
1795 size_t scores_bytes);
1807 size_t scores_bytes);
1832 const float *k_cache,
1833 const float *v_cache,
1843 const float *k_cache,
1844 const float *v_cache,
1855 size_t scores_bytes);
1865 int aligned_head_dim,
1866 int kv_stride_tokens);
1877 int aligned_head_dim,
1878 int kv_stride_tokens);
1888 int aligned_head_dim,
1889 int kv_stride_tokens);
1900 int aligned_head_dim,
1901 int kv_stride_tokens);
1912 int aligned_head_dim,
1913 int kv_stride_tokens);
1924 int aligned_head_dim,
1925 int kv_stride_tokens);
1928 const float *q,
const float *k,
const float *v,
float *output,
1929 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
1930 int aligned_head_dim,
int kv_stride_tokens);
1933 const float *q,
const float *k,
const float *v,
float *output,
1934 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
1935 int aligned_head_dim,
int kv_stride_tokens);
1945 int aligned_head_dim,
1946 int kv_stride_tokens);
1948 const float *q,
const float *k,
const float *v,
float *output,
1949 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
1950 int aligned_head_dim,
int kv_stride_tokens);
1956 int aligned_head_dim,
1957 int kv_stride_tokens);
1967 int aligned_head_dim,
1968 int kv_stride_tokens);
1978 int aligned_head_dim,
1979 int kv_stride_tokens,
1983 const float *q,
const float *k,
const float *v,
float *output,
1984 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
1985 int aligned_head_dim,
int kv_stride_tokens,
int visual_start,
1999 int aligned_head_dim,
2000 int kv_stride_tokens);
2013 int aligned_head_dim,
2014 int kv_stride_tokens);
2016 const float *q,
const float *k,
const float *v,
float *output,
2017 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
2018 int aligned_head_dim,
int kv_stride_tokens,
2019 float *score_rows,
size_t score_rows_bytes,
2020 float *v_columns,
size_t v_columns_bytes,
2021 float *probability_row,
size_t probability_row_bytes);
2032 int aligned_head_dim,
2033 int kv_stride_tokens);
2035 const float *q,
const float *k,
const float *v,
float *output,
2036 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
2037 int aligned_head_dim,
int kv_stride_tokens,
2038 float *rounded_kv,
size_t rounded_kv_bytes);
2045 const float *k_cache,
2046 const float *v_cache,
2053 int aligned_head_dim);
2055 const float *k_cache,
2056 const float *v_cache,
2063 int aligned_head_dim);
2065 const float *k_cache,
2066 const float *v_cache,
2072 int aligned_head_dim);
2078 const float *k_cache,
2079 const float *v_cache,
2087 int aligned_head_dim);
2091 const float *k_cache,
2092 const float *v_cache,
2099 int aligned_head_dim);
2102 const uint16_t *k_cache,
2103 const uint16_t *v_cache,
2110 int aligned_head_dim);
2142 const float *q_token,
2143 const uint16_t *k_cache,
2144 const uint16_t *v_cache,
2151 int aligned_head_dim,
2154 const float *q_token,
2155 const uint16_t *k_cache,
2156 const uint16_t *v_cache,
2163 int aligned_head_dim,
2167 const uint16_t *key_cache,
2168 const uint16_t *value_cache,
2169 const float *selected_indices,
2171 float *score_scratch,
2176 int selection_width,
2190 const uint16_t *k_cache,
2191 const uint16_t *v_cache,
2199 int aligned_head_dim,
2203 const uint16_t *k_cache,
2204 const uint16_t *v_cache,
2212 int aligned_head_dim,
2214 float *token_workspace,
2215 size_t token_workspace_bytes);
2218 const uint16_t *k_cache,
2219 const uint16_t *v_cache,
2227 int aligned_head_dim,
2229 float *token_workspace,
2230 size_t token_workspace_bytes,
2231 void *gqa_workspace,
2232 size_t gqa_workspace_bytes,
2233 int route_num_heads,
2234 int route_num_kv_heads,
2236 int route_query_tokens,
2237 int route_min_kv_tokens,
2239 int route_query_tile_size,
2240 int route_concurrent_query_tiles);
2246 const uint16_t *k_cache,
2247 const uint16_t *v_cache,
2255 int aligned_head_dim,
2266 int query_tile_size,
2267 int concurrent_query_tiles);
2270 const uint16_t *k_cache,
2271 const uint16_t *v_cache,
2279 int aligned_head_dim,
2280 int query_tile_size,
2281 int concurrent_query_tiles,
2283 size_t workspace_bytes);
2286 const float *q,
const uint16_t *k_cache,
const uint16_t *v_cache,
2287 float *output,
int num_heads,
int num_kv_heads,
int q_tokens,
2288 int past_tokens,
int cache_capacity,
int head_dim,
int aligned_head_dim,
2290 size_t token_workspace_bytes,
const int *segment_lengths,
2294 const uint16_t *k_cache,
2295 const uint16_t *v_cache,
2303 int aligned_head_dim,
2307 const uint16_t *k_cache,
2308 const uint16_t *v_cache,
2316 int aligned_head_dim,
2318 float *token_workspace,
2319 size_t token_workspace_bytes);
2326 const uint16_t *k_cache,
2327 const uint16_t *v_cache,
2335 int aligned_head_dim,
2344 const uint16_t *k_cache,
2345 const uint16_t *v_cache,
2352 int aligned_head_dim,
2361 const float *k_cache,
2362 const float *v_cache,
2369 int aligned_head_dim);
2383 int aligned_head_dim,
2384 int kv_stride_tokens,
2385 int sliding_window);
2387 const float *q,
const float *k,
const float *v,
float *output,
2388 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
2389 int aligned_head_dim,
int kv_stride_tokens,
int sliding_window,
2390 float *scores,
size_t scores_bytes,
2391 float *value_columns,
size_t value_columns_bytes,
2392 float *scaled_scores,
size_t scaled_scores_bytes);
2394 const float *q,
const float *k,
const float *v,
float *output,
2395 int num_heads,
int num_kv_heads,
int live_tokens,
int kv_stride_tokens,
2396 int head_dim,
int aligned_head_dim,
int sliding_window,
2397 float *scores,
size_t scores_bytes,
2398 float *value_columns,
size_t value_columns_bytes,
2399 float *scaled_scores,
size_t scaled_scores_bytes);
2410 int aligned_head_dim,
2411 int kv_stride_tokens,
2412 int sliding_window);
2422 int aligned_head_dim,
2423 int kv_stride_tokens,
2424 int sliding_window);
2431 int aligned_head_dim,
2432 int kv_stride_tokens,
2433 int sliding_window);
2449 const float *kernel,
2456 const float *kernel,
2463 const float *kernel,
2470 const float *kernel,
2477 const float *kernel,
2492 const float *conv_x,
2493 const float *kernel,
2528 const float *d_gate,
2571 float *d_packed_qkv,
2587 const float *dt_bias,
2597 const float *dt_bias,
2620 const float *dt_bias,
2629 const float *dt_bias,
2660 const float *d_state_out,
2672 const float *d_state_out,
2677 float *d_conv_total,
2736 float *d_packed_qkv,
2772 const float *d_k_out,
2791 const float *weight,
2799 const float *weight,
2810 const float *weight,
2828 int intermediate_dim,
2834 const float *weight,
2844 const float *weight,
2856 const float *weight,
2865 const float *dt_bias,
2939 const float *weight,
2986 const float *state_in,
2998 const float *state_in,
3011 const float *state_in,
3024 const float *state_in,
3038 const float *state_in,
3052 const float *state_in,
3066 const float *state_in,
3080 const float *state_in,
3091 const float *weight,
3099 const float *weight,
3107 const float *weight,
3124 const float *d_state_out,
3130 const float *state_in,
3131 const float *state_out,
3145 const float *q_token,
3146 const float *k_cache,
3147 const float *v_cache,
3154 int aligned_head_dim,
3155 int sliding_window);
3158 const float *q_token,
3159 const float *k_cache,
3160 const float *v_cache,
3167 int aligned_head_dim,
3168 int sliding_window);
3170 const float *q_token,
3171 const float *k_cache,
3172 const float *v_cache,
3178 int aligned_head_dim,
3179 int sliding_window);
3207 const float *k_cache,
3208 const float *v_cache,
3215 int aligned_head_dim);
3219 const float *__restrict v_token,
3220 float *__restrict k_cache,
3221 float *__restrict v_cache,
3226 int aligned_head_dim);
3229 float *__restrict kv_cache_v,
3230 const float *__restrict k,
3231 const float *__restrict v,
3238 float *__restrict kv_cache_v,
3239 const float *__restrict q,
3247 uint16_t *__restrict kv_cache_v,
3248 const float *__restrict k,
3249 const float *__restrict v,
3256 uint16_t *__restrict kv_cache_v,
3257 const float *__restrict k,
3258 const float *__restrict v,
3265 float *__restrict kv_cache_v,
3266 const float *__restrict k,
3267 const float *__restrict v,
3274 uint16_t *__restrict kv_cache_v,
3275 const float *__restrict k,
3276 const float *__restrict v,
3284 uint16_t *__restrict kv_cache_v,
3285 const float *__restrict k,
3286 const float *__restrict v,
3302 int aligned_head_dim);
3331 const uint16_t *W_fc1,
3332 const uint16_t *b_fc1,
3333 const uint16_t *W_fc2,
3334 const uint16_t *b_fc2,
3340 float *scratch_bias1_f,
3341 float *scratch_bias2_f,
3342 uint16_t *scratch_fc1_bf16);
3346 const uint16_t *W_fc1,
3347 const uint16_t *b_fc1,
3348 const uint16_t *W_fc2,
3349 const uint16_t *b_fc2,
3355 float *scratch_input_f,
3356 float *scratch_bias1_f,
3357 float *scratch_bias2_f,
3358 uint16_t *scratch_fc1_bf16);
3362 const uint16_t *W_fc1,
3363 const uint16_t *b_fc1,
3364 const uint16_t *W_fc2,
3365 const uint16_t *d_output,
3374 float *scratch_fc1_pre,
3375 uint16_t *scratch_fc1_act_bf16,
3376 float *scratch_d_fc1);
3380 const float *fc2_input,
3391 const float *fc1_input,
3409 const float *d_output,
3417 float *scratch_input,
3418 float *scratch_output);
3421 const uint16_t *d_output,
3424 float *scratch_input,
3425 float *scratch_d_output,
3426 float *scratch_d_input);
3442 const float *d_output,
3470 const float *d_output,
3481 const uint16_t *d_output,
3556 const float **vectors,
3557 const float *weights,
3578 const float *expert_output,
3579 float routing_weight,
3585 const float *routing_weights,
3586 const float *expert_up,
3587 const float *expert_down,
3591 int intermediate_dim,
3597 const float *routing_weights,
3598 const void *expert_up,
3599 const void *expert_down,
3603 int intermediate_dim,
3609 const float *routing_weights,
3610 const void *expert_up,
3611 const void *expert_down,
3615 int intermediate_dim,
3620 const float *routed,
3621 const void *shared_up,
3622 const void *shared_down,
3626 int intermediate_dim);
3631 const float *routing_weights,
3632 const float *expert_gate,
3633 const float *expert_up,
3634 const float *expert_down,
3638 int intermediate_dim,
3643 const float *routed,
3644 const float *shared_gate,
3645 const float *shared_up,
3646 const float *shared_down,
3650 int intermediate_dim);
3654 const float *routing_weights,
3655 const uint16_t *expert_gate,
3656 const uint16_t *expert_up,
3657 const uint16_t *expert_down,
3661 int intermediate_dim,
3666 const float *hidden,
const int *indices,
const float *routing_weights,
3667 const uint16_t *expert_gate,
const uint16_t *expert_up,
3668 const uint16_t *expert_down,
float *output,
int rows,
int hidden_dim,
3669 int intermediate_dim,
int n_experts,
int top_k,
3670 int row_begin,
int row_end);
3673 const float *hidden,
const int *indices,
const float *routing_weights,
3674 const uint16_t *expert_gate,
const uint16_t *expert_up,
3675 const uint16_t *expert_down,
float *output,
int rows,
int hidden_dim,
3676 int intermediate_dim,
int n_experts,
int top_k);
3679 const float *hidden,
const int *indices,
const float *routing_weights,
3680 const uint16_t *expert_gate_up,
const uint16_t *expert_down,
3681 float *output,
int rows,
int hidden_dim,
int intermediate_dim,
3682 int n_experts,
int top_k);
3685 int intermediate_dim);
3687 int intermediate_dim);
3689 int hidden_dim,
int intermediate_dim);
3692 const float *hidden,
3694 const float *routing_weights,
3695 const void *expert_gate,
3696 const void *expert_up,
3697 const void *expert_down,
3701 int intermediate_dim,
3705 size_t workspace_bytes);
3708 const float *hidden,
3710 const float *routing_weights,
3711 const void *expert_gate,
3712 const void *expert_up,
3713 const void *expert_down,
3717 int intermediate_dim,
3721 size_t workspace_bytes);
3724 const float *hidden,
3726 const float *routing_weights,
3727 const void *expert_gate,
3728 const void *expert_up,
3729 const void *expert_down,
3733 int intermediate_dim,
3737 size_t workspace_bytes);
3740 const float *hidden,
3742 const float *routing_weights,
3743 const void *expert_gate,
3744 const void *expert_up,
3745 const void *expert_down,
3749 int intermediate_dim,
3753 size_t workspace_bytes);
3756 const float *hidden,
3758 const float *routing_weights,
3759 const void *expert_gate,
3760 const void *expert_up,
3761 const void *expert_down,
3765 int intermediate_dim,
3769 size_t workspace_bytes);
3772 const float *hidden,
3774 const float *routing_weights,
3775 const void *expert_gate,
3776 const void *expert_up,
3777 const void *expert_down,
3781 int intermediate_dim,
3785 size_t workspace_bytes);
3788 const float *hidden,
3790 const float *routing_weights,
3791 const void *expert_gate,
3792 const void *expert_up,
3793 const void *expert_down,
3797 int intermediate_dim,
3801 size_t workspace_bytes);
3804 const float *hidden,
const int *indices,
const float *routing_weights,
3805 const void *expert_gate,
const void *expert_up,
const void *expert_down,
3806 float *output,
int rows,
int hidden_dim,
int intermediate_dim,
3807 int n_experts,
int top_k,
void *workspace,
size_t workspace_bytes);
3810 const float *hidden,
const int *indices,
const float *routing_weights,
3811 const void *expert_gate,
const void *expert_up,
const void *expert_down,
3812 float *output,
int rows,
int hidden_dim,
int intermediate_dim,
3813 int n_experts,
int top_k,
void *workspace,
size_t workspace_bytes);
3816 const float *hidden,
3817 const float *routed,
3818 const void *shared_gate,
3819 const void *shared_up,
3820 const void *shared_down,
3824 int intermediate_dim,
3826 size_t workspace_bytes);
3829 const float *hidden,
3830 const float *routed,
3831 const void *shared_gate,
3832 const void *shared_up,
3833 const void *shared_down,
3837 int intermediate_dim,
3839 size_t workspace_bytes);
3842 const float *hidden,
3843 const float *routed,
3844 const void *shared_gate,
3845 const void *shared_up,
3846 const void *shared_down,
3850 int intermediate_dim,
3852 size_t workspace_bytes);
3855 const float *hidden,
3856 const float *routed,
3857 const void *shared_gate,
3858 const void *shared_up,
3859 const void *shared_down,
3863 int intermediate_dim,
3865 size_t workspace_bytes);
3868 const float *hidden,
3870 const float *routing_weights,
3871 const void *expert_gate,
3872 const void *expert_up,
3873 const void *expert_down,
3877 int intermediate_dim,
3881 size_t workspace_bytes);
3884 const float *hidden,
3886 const float *routing_weights,
3887 const void *expert_gate,
3888 const void *expert_up,
3889 const void *expert_down,
3893 int intermediate_dim,
3897 size_t workspace_bytes);
3900 const float *hidden,
3902 const float *routing_weights,
3903 const void *expert_gate,
3904 const void *expert_up,
3905 const void *expert_down,
3906 const void *expert_gate_packed,
3907 const void *expert_up_packed,
3911 int intermediate_dim,
3915 size_t workspace_bytes);
3918 const float *hidden,
3920 const float *routing_weights,
3921 const void *expert_gate,
3922 const void *expert_up,
3923 const void *expert_down,
3927 int intermediate_dim,
3931 size_t workspace_bytes);
3934 const float *hidden,
3936 const float *routing_weights,
3937 const void *expert_gate,
3938 const void *expert_up,
3939 const void *expert_down,
3943 int intermediate_dim,
3947 size_t workspace_bytes);
3950 int intermediate_dim);
3953 const float *hidden,
3954 const float *routed,
3955 const void *shared_gate,
3956 const void *shared_up,
3957 const void *shared_down,
3958 const float *shared_gate_input,
3962 int intermediate_dim,
3964 size_t workspace_bytes);
3967 const float *hidden,
3968 const float *routed,
3969 const void *shared_gate,
3970 const void *shared_up,
3971 const void *shared_down,
3972 const float *shared_gate_input,
3976 int intermediate_dim,
3978 size_t workspace_bytes);
3981 const float *hidden,
const float *routed,
const void *shared_gate,
3982 const void *shared_up,
const void *shared_down,
3983 const float *shared_gate_input,
float *output,
int rows,
int hidden_dim,
3984 int intermediate_dim,
void *workspace,
size_t workspace_bytes);
3987 const float *hidden,
const float *routed,
const void *shared_gate,
3988 const void *shared_up,
const void *shared_down,
3989 const float *shared_gate_input,
float *output,
int rows,
int hidden_dim,
3990 int intermediate_dim,
void *workspace,
size_t workspace_bytes);
3993 const float *hidden,
const float *routed,
const void *shared_gate,
3994 const void *shared_up,
const void *shared_down,
3995 const float *shared_gate_input,
float *output,
int rows,
int hidden_dim,
3996 int intermediate_dim,
void *workspace,
size_t workspace_bytes);
3999 const float *hidden,
const float *routed,
const void *shared_gate,
4000 const void *shared_up,
const void *shared_down,
4001 const float *shared_gate_input,
float *output,
int rows,
int hidden_dim,
4002 int intermediate_dim,
void *workspace,
size_t workspace_bytes);
4005 const float *routed,
4006 const uint16_t *shared_gate,
4007 const uint16_t *shared_up,
4008 const uint16_t *shared_down,
4012 int intermediate_dim);
4014 const float *hidden,
const float *routed,
const uint16_t *shared_gate,
4015 const uint16_t *shared_up,
const uint16_t *shared_down,
float *output,
4016 int rows,
int hidden_dim,
int intermediate_dim,
4017 int row_begin,
int row_end);
4019 const float *hidden,
const float *routed,
const uint16_t *shared_gate,
4020 const uint16_t *shared_up,
const uint16_t *shared_down,
float *output,
4021 int rows,
int hidden_dim,
int intermediate_dim);
4023 const float *hidden,
const float *routed,
const uint16_t *shared_gate,
4024 const uint16_t *shared_up,
const uint16_t *shared_down,
4025 const uint16_t *shared_router,
float *output,
int rows,
4026 int hidden_dim,
int intermediate_dim);
4028 const float *hidden,
const float *routed,
const uint16_t *shared_gate,
4029 const uint16_t *shared_up,
const uint16_t *shared_down,
4030 const uint16_t *shared_router,
float *output,
int rows,
4031 int hidden_dim,
int intermediate_dim,
int row_begin,
int row_end);
4033 const float *hidden,
const float *routed,
const uint16_t *shared_gate,
4034 const uint16_t *shared_up,
const uint16_t *shared_down,
4035 const uint16_t *shared_router,
float *output,
int rows,
4036 int hidden_dim,
int intermediate_dim);
4038 const float *routed,
4039 const float *post_attn_residual,
4040 const uint16_t *shared_gate,
4041 const uint16_t *shared_up,
4042 const uint16_t *shared_down,
4044 float *routed_free_output,
4047 int intermediate_dim);
4049 const float *hidden,
const float *routed,
4050 const float *post_attn_residual,
const uint16_t *shared_gate,
4051 const uint16_t *shared_up,
const uint16_t *shared_down,
4052 float *main_output,
float *routed_free_output,
int rows,
4053 int hidden_dim,
int intermediate_dim,
int row_begin,
int row_end);
4055 const float *hidden,
const float *routed,
4056 const float *post_attn_residual,
const uint16_t *shared_gate,
4057 const uint16_t *shared_up,
const uint16_t *shared_down,
4058 float *main_output,
float *routed_free_output,
int rows,
4059 int hidden_dim,
int intermediate_dim);
4062 const float *correction_bias,
4071 float routed_scaling_factor);
4074 const float *hidden,
4076 const float *routing_weights,
4077 const float *expert_up,
4078 const float *expert_down,
4080 float *d_routing_weights,
4082 float *d_expert_down,
4085 int intermediate_dim,
4110 const float *logits,
4116 float routed_scaling_factor,
4118 size_t workspace_bytes);
4120 const float *logits,
4126 float routed_scaling_factor,
4128 size_t workspace_bytes);
4132 const float *weights,
4133 const float *d_weights,
4136 int n_experts_or_keys,
4149 const float *correction_bias,
4158 float routed_scaling_factor);
4169 int *verified_token);
4179 int *target_position,
4180 int *draft_position,
4181 int *accepted_count,
4182 int *rejected_count);
4196 const float *streams,
4213 const float *weights,
4214 const float *d_weights,
4239 float *score_scratch,
4251 float *score_scratch,
4252 float *key_transpose_scratch,
4301 const float *kv_b_proj,
4311 const uint16_t *kv_b_proj,
4320 const uint16_t *kv_b_proj,
4331 const float *compressed_kv,
4332 const uint16_t *kv_b_proj,
4343 const float *k_nope,
4355 const float *k_nope,
4356 const float *kv_a_packed,
4368 const float *q_packed,
4369 const float *k_nope,
4370 const float *kv_a_packed,
4383 const float *d_output,
4387 const float *attn_weights,
4396 int aligned_head_dim,
4397 int aligned_context_window);
4401 const float *d_output,
4405 const float *attn_weights,
4413 int aligned_head_dim,
4414 int aligned_context_window);
4418 const uint16_t *d_output,
4423 const float *attn_weights,
4432 int aligned_head_dim,
4433 int aligned_context_window,
4434 float *scratch_d_output,
4453 const char *scaling_type,
4454 float scaling_factor);
4460 const int32_t *positions,
4465 int original_context,
4469 float mscale_all_dim);
4476 int original_context,
4480 float mscale_all_dim);
4482 uint16_t *sin_cache,
4483 const int32_t *positions,
4488 int original_context,
4492 float mscale_all_dim);
4500 const char *scaling_type,
4501 float scaling_factor);
4505 const float *cos_cache,
4506 const float *sin_cache,
4510 int aligned_head_dim,
4514 const float *cos_cache,
4515 const float *sin_cache,
4519 int aligned_head_dim,
4526 const float *cos_cache,
4527 const float *sin_cache,
4531 int aligned_head_dim,
4536 const float *cos_cache,
4537 const float *sin_cache,
4541 int aligned_head_dim,
4548 const float *cos_cache,
4549 const float *sin_cache,
4553 int aligned_head_dim,
4555 float *scratch_d_out,
4556 float *scratch_d_x);
4560 const float *cos_cache,
4561 const float *sin_cache,
4565 int aligned_head_dim,
4569 const float *cos_cache,
4570 const float *sin_cache,
4574 int aligned_head_dim,
4576 int head_stride_tokens);
4579 const float *cos_cache,
4580 const float *sin_cache,
4584 int aligned_head_dim,
4586 int head_stride_tokens,
4592 const float *cos_cache,
4593 const float *sin_cache,
4598 int aligned_head_dim,
4603 const float *cos_cache,
4604 const float *sin_cache,
4609 int aligned_head_dim,
4614 float *q,
float *k,
const float *freq_factors,
int use_freq_factors,
4615 int num_heads,
int num_kv_heads,
int num_tokens,
int head_dim,
4616 int aligned_head_dim,
int pos_offset,
int rotary_dim,
float freq_base,
4617 int token_begin,
int token_end);
4621 const float *freq_factors,
4622 int use_freq_factors,
4627 int aligned_head_dim,
4634 const float *freq_factors,
4635 int use_freq_factors,
4640 int aligned_head_dim,
4647 const float *freq_factors,
4648 int use_freq_factors,
4652 int aligned_head_dim,
4659 const float *freq_factors,
4660 int use_freq_factors,
4665 int aligned_head_dim,
4671 const float *cos_cache,
4672 const float *sin_cache,
4677 int aligned_head_dim,
4680 int cache_rotary_dim);
4684 const float *cos_cache,
4685 const float *sin_cache,
4690 int aligned_head_dim,
4696 const float *cos_cache,
4697 const float *sin_cache,
4702 int aligned_head_dim,
4712 int aligned_head_dim,
4729 const int32_t *positions,
4734 int aligned_head_dim,
4750 int num_heads,
int num_kv_heads,
int num_tokens,
4751 int head_dim,
int aligned_head_dim,
int n_dims,
4752 int section_0,
int section_1,
int section_2,
int section_3,
4753 int n_ctx_orig,
float freq_base,
float freq_scale,
4754 float ext_factor,
float attn_factor,
4755 float beta_fast,
float beta_slow);
4757 int num_heads,
int num_kv_heads,
int num_tokens,
4758 int head_dim,
int aligned_head_dim,
int n_dims,
4759 int section_0,
int section_1,
int section_2,
int section_3,
4760 int n_ctx_orig,
float freq_base,
float freq_scale,
4761 float ext_factor,
float attn_factor,
4762 float beta_fast,
float beta_slow);
4765 float *q,
float *k,
const int32_t *positions,
4766 int num_heads,
int num_kv_heads,
int num_tokens,
4767 int head_dim,
int aligned_head_dim,
int n_dims,
4768 int section_0,
int section_1,
int section_2,
int section_3,
4769 int n_ctx_orig,
float freq_base,
float freq_scale,
4770 float ext_factor,
float attn_factor,
float beta_fast,
float beta_slow);
4773 int num_heads,
int num_kv_heads,
int num_tokens,
4774 int head_dim,
int aligned_head_dim,
int n_dims,
4775 int section_0,
int section_1,
int section_2,
int section_3,
4776 int n_ctx_orig,
float freq_base,
float freq_scale,
4777 float ext_factor,
float attn_factor,
4778 float beta_fast,
float beta_slow);
4786 int aligned_head_dim,
4793 const int32_t *positions,
4798 int aligned_head_dim,
4818 int aligned_head_dim,
4839 int aligned_head_dim,
4856 const float *cos_cache,
4857 const float *sin_cache,
4862 int aligned_head_dim,
4864 int q_stride_tokens,
4865 int k_stride_tokens);
4869 const float *cos_cache,
4870 const float *sin_cache,
4875 int aligned_head_dim,
4877 int q_stride_tokens,
4878 int k_stride_tokens,
4882 const float *d_k_out,
4885 const float *cos_cache,
4886 const float *sin_cache,
4891 int aligned_head_dim,
4895 const float *d_k_out,
4898 const float *cos_cache,
4899 const float *sin_cache,
4904 int aligned_head_dim,
4911 const float *cos_cache,
4912 const float *sin_cache,
4917 int aligned_head_dim,
4924 const float *cos_cache,
4925 const float *sin_cache,
4930 int aligned_head_dim,
4937 const float *cos_cache,
4938 const float *sin_cache,
4942 int aligned_head_dim,
4949 const uint16_t *d_k_out,
4952 const float *cos_cache,
4953 const float *sin_cache,
4958 int aligned_head_dim,
4960 float *scratch_dq_out,
4962 float *scratch_dk_out,
4972 const float *token_embeddings,
4973 const float *pos_embeddings,
4976 int aligned_embed_dim,
4983 const void *token_embeddings,
4984 const float *pos_embeddings,
4987 int aligned_embed_dim,
4994 const void *token_embeddings,
4995 const float *pos_embeddings,
4998 int aligned_embed_dim,
5005 const void *token_embeddings,
5006 const float *pos_embeddings,
5009 int aligned_embed_dim,
5016 const void *token_embeddings,
5017 const float *pos_embeddings,
5020 int aligned_embed_dim,
5027 const uint16_t *token_embeddings,
5028 const uint16_t *pos_embeddings,
5031 int aligned_embed_dim,
5038 const uint16_t *token_embeddings,
5039 const float *pos_embeddings,
5042 int aligned_embed_dim,
5052 const float *d_output,
5053 float *d_token_embeddings,
5054 float *d_pos_embeddings,
5057 int aligned_embed_dim,
5063 const uint16_t *d_output,
5064 uint16_t *d_token_embeddings,
5065 uint16_t *d_pos_embeddings,
5068 int aligned_embed_dim,
5074 const uint16_t *d_output,
5075 float *d_token_embeddings,
5076 float *d_pos_embeddings,
5079 int aligned_embed_dim,
5086 const int32_t *targets,
5092 const int32_t *targets,
5100 const int32_t *targets,
5105 float *scratch_logits,
5106 float *scratch_d_logits);
5111 int C,
int H,
int W,
int P);
5112 void patch2im(
const float *d_patches,
5114 int C,
int H,
int W,
int P);
5116 const float *position_embd,
5121 const float *position_embd,
5125 int start_position);
5127 const float *position_embd,
5132 int source_grid_size);
5134 const float *position_embd,
5139 int source_grid_size);
5141 const float *position_embd,
5146 int source_grid_size);
5149 const float *position_embd,
5154 int source_grid_size);
5156 const float *position_embd,
5160 int source_grid_size);
5196 const float *branch_input,
5200 int branch_slice_dim,
5201 int num_branch_slices);
5211 float *decoder_rows,
5213 int source_row_stride,
5214 int decoder_row_stride,
5217 int decoder_capacity);
5230 int C,
int H,
int W,
int P);
5233 int C,
int H,
int W,
int P);
CKDataType
Supported data types in C-Kernel-Engine.
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)
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.
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)
void add_stream_reorder_2d(float *main_inout, float *aux_scratch, int grid_h, int grid_w, int embed_dim, int merge_size)
ck_attention_prefill_schedule_t
@ CK_ATTN_PREFILL_SCHEDULE_QUERY_TILES
@ CK_ATTN_PREFILL_SCHEDULE_KV_HEADS
@ CK_ATTN_PREFILL_SCHEDULE_QUERY_HEADS
@ CK_ATTN_PREFILL_SCHEDULE_GQA_SHARED_KV_TILES
@ CK_ATTN_PREFILL_SCHEDULE_KV_GROUP_QUERY_TILES
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)
float ck_attention_pytorch_sdpa_scale_f32(int head_dim)
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)
void dequant_q4_0_row(const void *src, float *dst, size_t n_elements)
Dequantize Q4_0 row (multiple blocks)
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)
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.
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)
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)
void axpy_f32(float *y, const float *x, float alpha, int n)
In-place AXPY: y += alpha * x.
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)
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)
size_t moe_swiglu_expert_q4k_q5k_workspace_bytes(int hidden_dim, int intermediate_dim)
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)
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)
void softmax_cross_entropy_loss_ptref(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out)
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)
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)
void gemm_q6_k(float *Y, const void *W, const float *X, int M, int N, int K)
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)
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)
int ck_flash_attn_choose_tile_k(int D_h)
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)
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)
void * ck_memcpy_parallel_dispatch(void *dst, const void *src, size_t size)
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)
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)
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)
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.
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)
void swiglu_forward_exact(const float *input, float *output, int tokens, int dim)
int ck_gemm_nt_f16_ggml_oracle(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
void spatial_merge_2x2(const float *input, float *output, int grid_h, int grid_w, int embed_dim)
void gated_deltanet_autoregressive_backward(const float *d_out, const float *d_state_out, const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, const float *state_out, float *d_q, float *d_k, float *d_v, float *d_g, float *d_beta, float *d_state_in, int num_heads, int state_dim, float norm_eps)
void 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)
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)
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)
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.
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)
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)
void rmsnorm_forward_no_weight(const float *input, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void gemm_nt_fp32_exact_parallel_dispatch(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void swiglu_forward_pytorch_bf16_storage(const float *input, float *output, int tokens, int dim)
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)
void gemv_q5_1(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV.
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)
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)
void gemm_naive_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
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)
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)
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)
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)
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)
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)
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)
void gemv_fused_q8_0_bias_dispatch(float *y, const void *W, const float *x, const float *bias, int M, int K)
void gelu_erf_fp64_f32_inplace(float *data, size_t n)
void recurrent_qk_l2_norm_forward(float *q, float *k, int rows, int q_dim, int k_dim, int head_dim, float eps)
void ck_residual_add_token_major_bf16_storage(const float *a, const float *b, float *out, int tokens, int aligned_embed_dim)
void swiglu_forward(const float *input, float *output, int tokens, int dim)
void gemm_bias_silu_fused(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void geglu_backward_bf16_mixed(const uint16_t *x, const uint16_t *d_out, float *d_x, int tokens, int dim)
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)
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)
void add_scaled_forward_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, float alpha, size_t n)
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)
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)
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)
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)
void geglu_forward_exact(const float *x, float *out, int tokens, int dim)
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)
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)
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)
void ck_set_num_threads(int num_threads)
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)
void attn_gate_sigmoid_mul_pytorch_bf16_storage(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
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)
void swiglu_backward(const float *input, const float *d_output, float *d_input, int tokens, int dim)
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.
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)
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)
void hyper_stream_expand_f32(const float *input, float *output, int rows, int streams, int hidden_dim)
void gelu_ggml_inplace(float *data, size_t n)
void gelu_backward_exact(const float *input, const float *d_output, float *d_input, size_t n)
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)
void axpy_zero_f32(float *y, const float *x, float alpha, int n)
Zero output then accumulate: y = 0; y += alpha * x.
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)
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.
void gelu_exact_inplace(float *data, size_t n)
void gemm_nt_q4_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void recurrent_sigmoid_forward_ggml(const float *x, float *out, int rows, int dim)
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)
void add_inplace_bf16(uint16_t *a, const uint16_t *b, size_t n)
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)
void gemm_blocked_serial_train_parallel_dispatch(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void ck_residual_add_backward(const float *d_out, float *d_a, float *d_b, int tokens, int aligned_embed_dim)
void gemm_nn_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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.
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)
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)
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)
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)
void kv_cache_repack_head_major_inplace(float *buf, int num_heads, int tokens, int cache_capacity, int aligned_head_dim)
void gelu_pytorch_erf_f32_inplace(float *data, size_t n)
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)
void gemv_q5_1_q8_1_ref(float *y, const void *W, const void *x_q8, int M, int K)
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)
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)
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)
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)
void gemm_bias_relu_fused(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
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)
void gemv_q6_k(float *y, const void *W, const float *x, int M, int K)
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.
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)
void backward_causal_softmax_head_major(float *d_scores, const float *weights, int num_heads, int num_tokens, int aligned_context_window)
void gemv_bf16_bf16_storage(float *y, const void *W, const float *x, int M, int K)
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)
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.
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)
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)
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.
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)
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)
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)
void patch2im(const float *d_patches, float *d_image, int C, int H, int W, int P)
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)
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)
void swiglu_forward_bf16(const uint16_t *input, uint16_t *output, int tokens, int dim)
int ck_gemm_bf16_amx_available(void)
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)
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)
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)
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)
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)
void relu_backward(const float *input, const float *d_output, float *d_input, size_t n)
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)
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)
void geglu_backward_fp32(const float *x, const float *d_out, float *d_x, int tokens, int dim)
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)
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)
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)
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)
void gated_deltanet_llama_avx2_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
void 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)
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)
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.
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)
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)
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)
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.
const char * ck_q6_k_q8_k_provider_name(void)
void dequant_q5_0_row(const void *src, float *dst, size_t n_elements)
Dequantize Q5_0 row (multiple blocks)
void add_forward_f32(const float *a, const float *b, float *y, size_t n)
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)
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)
void quantize_batch_q8_k_4row_nearest_even(const float *x, void *y, int num_rows, int k)
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)
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)
float gradient_clip_norm_f32(float *grad, size_t numel, float max_norm)
Clip gradient norm (fp32)
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)
void gemv_q8_0_q8_0_contract(float *y, const void *W, const float *x, int M, int K)
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.
void recurrent_silu_forward_pytorch_bf16_input_fp32_output(const float *x, float *out, int rows, int dim)
void attn_gate_sigmoid_mul_forward(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
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)
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)
void add_inplace_f32(float *a, const float *b, size_t n)
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)
void gemv_q5_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
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)
void ck_strict_store_next_gemm_a(const float *data, size_t elems)
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)
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)
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)
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)
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)
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)
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)
void relu_forward_inplace_bf16(uint16_t *data, size_t n)
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)
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)
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)
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)
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.
void im2patch_bf16(const uint16_t *image, uint16_t *patches, int C, int H, int W, int P)
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)
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)
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)
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)
void attn_gate_softplus_mul_forward(const float *x, const float *gate, float *out, int rows, int num_heads, int state_dim)
void gemm_microkernel(const float *A, const float *B, float *C, int M, int N, int K, int B_transposed)
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)
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)
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)
void ck_layout_token_to_head_f32(const float *src, float *dst, int tokens, int heads, int head_dim)
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)
void deepseek_mhc_mix_f32(const float *streams, const float *mix, float *out, int tokens, int n_streams, int dim)
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)
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)
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)
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)
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)
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)
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)
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)
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)
void gemm_nt_q5_1_q8_1(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void feature_slice_copy(const float *src, float *dst, int rows, int src_dim, int dst_dim, int dst_feature_offset)
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)
void gemm_nt_bf16_native_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
void recurrent_sigmoid_forward_pytorch_bf16_input_fp32_output(const float *x, float *out, int rows, int dim)
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)
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)
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)
void ck_layout_head_to_token_f32(const float *src, float *dst, int heads, int tokens, int head_dim)
void recurrent_silu_forward_ggml(const float *x, float *out, int rows, int dim)
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])
void spatial_merge_contiguous_tiled(const float *input, float *output, int grid_h, int grid_w, int embed_dim, int merge_size)
void gradient_accumulate_multi_f32(float *const *dsts, const float *const *srcs, const size_t *numels, int tensor_count)
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)
int ck_gemm_nt_f16_simd_lanes(void)
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)
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)
CKMathBackend ckernel_backend_native(void)
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)
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)
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)
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)
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)
void causal_softmax_head_major_exact(float *scores, int num_heads, int num_tokens, int aligned_context_window)
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)
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.
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)
void gemv_q4_k_q8_k_parallel(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
void gemm_nt_q5_0(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
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)
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)
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)
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)
void gemm_nt_q6_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
const float * ck_strict_consume_next_gemm_a(size_t elems)
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)
void swiglu_backward_exact(const float *input, const float *d_output, float *d_input, int tokens, int dim)
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)
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)
void swiglu_backward_bf16(const uint16_t *input, const uint16_t *d_output, uint16_t *d_input, int tokens, int dim)
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)
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)
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)
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)
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)
void relu2_backward(const float *input, const float *d_output, float *d_input, size_t n)
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)
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)
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)
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.
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)
void gemv_q5_k(float *y, const void *W, const float *x, int M, int K)
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)
void gemv_q8_0_q8_0_x4(float *y, const void *W, const void *x_q8, int M, int K)
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)
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)
int ck_get_physical_cores(void)
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)
void gated_deltanet_llama_avx2_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps)
void geglu_forward_fp32(const float *x, float *out, int tokens, int dim)
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)
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)
void relu_forward_bf16(const uint16_t *input, uint16_t *output, size_t n)
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)
void gemm_nt_bf16_amx_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
void gemm_tn_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
void gemm_nt_q8_0(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void relu_forward(const float *input, float *output, size_t n)
void im2patch(const float *image, float *patches, int C, int H, int W, int P)
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.
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)
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)
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)
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)
void quantize_row_q8_k(const float *x, void *y, int k)
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)
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,:].
void ck_set_strict_parity(int enabled)
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)
void recurrent_silu_forward(const float *x, float *out, int rows, int dim)
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)
void gemv_bf16_bf16_storage_parallel_dispatch(float *y, const void *W, const float *x, int M, int K)
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)
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)
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)
size_t moe_swiglu_shared_q8_0_gated_workspace_bytes(int hidden_dim, int intermediate_dim)
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)
void sigmoid_forward_bf16(const uint16_t *input, uint16_t *output, size_t n, float *scratch_input, float *scratch_output)
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)
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)
void geglu_forward_ggml_native(const float *x, float *out, int tokens, int dim)
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)
void gemm_nn_avx512(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void causal_softmax_head_major(float *scores, int num_heads, int num_tokens, int aligned_context_window)
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)
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)
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.
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!
int argmax_f32(const float *scores, int n)
Find index of maximum value.
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)
void assistant_layer_scale_forward(float *hidden, const float *scale, int tokens, int embed_dim)
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)
void gemv_q5_1_q8_1(float *y, const void *W, const float *x, int M, int K)
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)
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)
void gemma4_vision_projector_prep_forward(const float *input, float *output, int tokens, int dim, float scale, float eps)
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.
void add_forward_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, size_t n)
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)
void add_forward_2d_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, int tokens, int dim, int aligned_dim)
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)
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)
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)
void gemma4_final_logit_softcap_forward(float *logits, int tokens, int vocab_size, float cap)
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)
void gemm_avx512_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemv_q4_k(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV based on available SIMD.
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)
void gemm_nt_bf16(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gated_deltanet_pytorch_grouped_bf16_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int group_count, int state_dim, float norm_eps)
void 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)
void gelu_fast_inplace_bf16(uint16_t *data, size_t n, float *scratch)
void swiglu_forward_q8_k(const float *input, void *output_q8, int tokens, int dim)
void dequant_q8_0_row(const void *src, float *dst, size_t n_elements)
Dequantize Q8_0 row (multiple blocks)
float gradient_global_norm_multi_f32(const float *const *grads, const size_t *numels, int tensor_count)
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)
void gemm_nt_bf16_parallel_dispatch(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void gated_deltanet_pytorch_grouped_bf16_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
void q_norm_forward(float *q, const float *q_gamma, int num_heads, int num_tokens, int head_dim, float eps)
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)
void gemm_nn_simd(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
size_t moe_softmax_topk_router_workspace_bytes(int n_experts)
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)
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)
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)
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)
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)
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)
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.
int ck_strict_mtmd_clip_encode_planar_f32(const float *planar, int channels, int height, int width, float *out, size_t out_elems)
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)
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)
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)
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.
void recurrent_silu_forward_pytorch_bf16_storage(const float *x, float *out, int rows, int dim)
float sigmoid_scalar(float x)
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)
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)
void ck_residual_add_token_major_parallel_dispatch(const float *a, const float *b, float *out, int tokens, int aligned_embed_dim)
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)
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.
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)
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)
void sigmoid_backward(const float *input, const float *d_output, float *d_input, size_t n)
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)
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)
void quantize_row_q8_0(const float *x, void *y, int k)
Quantize FP32 to Q8_0 format (scalar reference)
void gemv_bf16(float *y, const void *W, const float *x, int M, int K)
void gelu_pytorch_erf_sleef_bf16_storage(float *data, size_t n)
void gated_deltanet_autoregressive_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int num_heads, int state_dim, float norm_eps)
void ck_residual_add_token_major(const float *a, const float *b, float *out, int tokens, int aligned_embed_dim)
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)
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)
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)
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)
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)
void gated_deltanet_llama_chunk64_head_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int head, int state_dim)
void 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)
size_t fused_rmsnorm_qkv_scratch_size(int hidden)
Get scratch buffer size for fused_rmsnorm_qkv_prefill.
void gelu_ggml_native_inplace(float *data, size_t n)
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.
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)
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)
void gradient_scale_f32(float *grad, size_t numel, float scale)
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)
void gemm_nt_q8_0_q8_0(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
gemm_nt_q8_0_q8_0 with optional bias (matches header signature)
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)
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)
void relu2_forward(const float *input, float *output, size_t n)
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)
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)
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)
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)
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)
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)
void gemm_fine_grained_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
size_t moe_swiglu_expert_q4k_q8_0_workspace_bytes(int hidden_dim, int intermediate_dim)
size_t moe_swiglu_shared_q4k_q8_0_gated_workspace_bytes(int hidden_dim, int intermediate_dim)
void spatial_average_pool_contiguous(const float *input, float *output, int grid_h, int grid_w, int embed_dim, int merge_size)
const char * ck_q6_k_prepared_provider_name(void)
void relu_forward_inplace(float *data, size_t n)
void gelu_fast_inplace(float *data, size_t n)
void gemv_q4_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
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)
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.
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)
void deepseek_dsa_topk_softmax_f32(const float *scores, int *indices, float *weights, int tokens, int heads, int key_count, int top_k)
void gemma4_v_norm_forward(const float *input, float *output, float *rstd_cache, int tokens, int num_kv_heads, int head_dim, float eps)
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)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)
int qk_norm_backward_last_isa(void)
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)
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)
void geglu_forward_bf16(const uint16_t *x, uint16_t *out, int tokens, int dim, float *scratch)
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)
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.
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)
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)
void gradient_accumulate_f32(float *dst, const float *src, size_t numel)
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)
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)
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)
void gelu_backward_scalar(const float *input, const float *d_output, float *d_input, size_t n)
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)
@ CK_ATTN_REDUCTION_BF16_PYTORCH_SDPA
@ CK_ATTN_REDUCTION_F16_ONLINE_FP32_MERGE
@ CK_ATTN_REDUCTION_F16_FLASH_AUTO_QTILE64
@ CK_ATTN_REDUCTION_F16_ONLINE_SINGLE_RANGE
@ CK_ATTN_REDUCTION_FP32_ONLINE
void gemm_q4_k_q8_k(float *Y, const void *W, const void *X_q8, int M, int N, int K)
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)
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)
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)
void gemm_microkernel_blocked(const float *A, const float *B, float *C, int M, int N, int K)
void dequant_q5_1_row(const void *src, float *dst, size_t n_elements)
Dequantize Q5_1 row (multiple blocks)
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)
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)
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.
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)
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)
void dequant_q4_1_row(const void *src, float *dst, size_t n_elements)
Dequantize Q4_1 row (multiple blocks)
void swiglu_forward_ggml(const float *input, float *output, int tokens, int dim)
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)
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)
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)
void gemv_bf16_parallel_dispatch(float *y, const void *W, const float *x, int M, int K)
void gemv_fused_q5_0_bias_dispatch(float *y, const void *W, const float *x, const float *bias, int M, int K)
void gated_deltanet_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int state_dim, float norm_eps)
void 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.
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)
void position_embeddings_add_at_offset(float *x, const float *position_embd, int num_tokens, int embed_dim, int num_positions, int start_position)
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)
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)
void swiglu_forward_ggml_split(const float *gate, const float *up, float *output, int tokens, int dim)
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)
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)
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)
int ck_flash_attn_fast_exp_kind(void)
void gated_deltanet_llama_chunk64_prefill_forward(const float *q, const float *k, const float *v, const float *g, const float *beta, const float *state_in, float *state_out, float *out, int rows, int num_heads, int group_count, int state_dim, float norm_eps)
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)
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)
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)
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)
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)
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)
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)
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)
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)
void gelu_erf_bf16_storage(float *data, size_t n)
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)
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)
void gemm_bias_gelu_fused(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
int ck_attention_sparse_bf16_pytorch_gqa_available(void)
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)
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)
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)
@ CK_ATTENTION_STATUS_UNSUPPORTED_CONTRACT
@ CK_ATTENTION_STATUS_INVALID_ARGUMENT
@ CK_ATTENTION_STATUS_INSUFFICIENT_WORKSPACE
void topk_f32(const float *scores, int n, int k, int *indices, float *values)
Find top-K indices and values from a score vector.
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.
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)
void gemm_microkernel_blocked_bt(const float *A, const float *B, float *C, int M, int N, int K)
void dequant_q6_k_row(const void *src, float *dst, size_t n_elements)
Dequantize Q6_K row (multiple blocks)
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)
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)
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)
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)
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)
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)
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)
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.
void gemm_nt_bf16_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void add_stream_inplace(float *a, const float *b, size_t n)
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)
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)
void split_q_gate_forward(const float *packed_qg, float *q, float *gate, int rows, int q_dim, int gate_dim, int group_dim)
void gemm_tn_blocked(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
void recurrent_silu_backward(const float *d_out, const float *x, float *d_x, int rows, int dim)
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)
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)
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)
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)
void gemm_nt_f32_llama_production(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void topk_softmax_f32(const float *scores, int n, int k, int *indices, float *weights)
Find top-K indices with softmax-normalized weights.
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)
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)
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.
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.
void rowwise_bias_add(float *x, const float *bias, int rows, int dim)
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)
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)
void gelu_backward_fast(const float *input, const float *d_output, float *d_input, size_t n)
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)
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)
void softmax_cross_entropy_loss(const float *logits, const int32_t *targets, int tokens, int vocab_size, float *d_logits, float *loss_out)
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)
void causal_softmax_head_major_bf16(uint16_t *scores, int num_heads, int num_tokens, int aligned_context_window, float *scratch)
void position_embeddings_add(float *x, const float *position_embd, int num_tokens, int embed_dim, int num_positions)
int ck_get_num_threads(void)
void patch2im_bf16(const uint16_t *d_patches, uint16_t *d_image, int C, int H, int W, int P)
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)
int ck_strict_parity_enabled(void)
void vision_position_ids_2d_merge(int32_t *positions, int grid_h, int grid_w, int merge_size)
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)
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)
void hyper_stream_expand_bf16(const float *input, float *output, int rows, int streams, int hidden_dim)
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)
void rope_precompute_cache_split(float *cos_cache, float *sin_cache, int max_seq_len, int head_dim, float base)
void add_backward_bf16(const uint16_t *d_y, uint16_t *d_a, uint16_t *d_b, size_t n)
void gemm_nt_q5_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
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)
void dequant_q4_k_row(const void *src, float *dst, size_t n_elements)
Dequantize Q4_K row (multiple blocks)
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)
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)
void scal_copy_f32(float *y, const float *x, float alpha, int n)
Scaled copy: y = alpha * x.
void gemm_microkernel_packed(const float *A, const float *B, float *C, int M, int N, int K)
void gemm_blocked_serial(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
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)
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)
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)
size_t fused_mlp_swiglu_scratch_size(int intermediate)
Get scratch buffer size for fused_mlp_swiglu_prefill.
void sigmoid_forward(const float *input, float *output, size_t n)
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)
void relu_backward_bf16(const uint16_t *input, const uint16_t *d_output, uint16_t *d_input, size_t n)
void final_logit_scale_f32(float *logits, int tokens, int vocab_size, float scale)
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)
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)
void gelu_pytorch_tanh_bf16_storage(float *data, size_t n)
int ck_attention_bf16_pytorch_gqa_available(void)
void gemv_q4_0(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV.
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)
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)
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)
void gemm_nn_blocked(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
void add_scaled_inplace_bf16(uint16_t *a, const uint16_t *b, float alpha, size_t n)
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)
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)
void gemm_tn_avx512(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
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)
Quantization block structures for weight-only quantization.
Mega-Fused Attention Kernel.