← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
kv_cache_kernels.c File Reference

KV-cache helper kernels (head-major layout) More...

#include "ckernel_engine.h"
#include "bf16_utils.h"
#include <stddef.h>
#include <string.h>

Go to the source code of this file.

Functions

static void ck_local_fp32_to_bf16_row (const float *src, uint16_t *dst, int n)
 
static void ck_local_fp32_to_fp16_row (const float *src, uint16_t *dst, int n)
 
void kv_cache_repack_head_major_inplace (float *buf, int num_heads, int tokens, int cache_capacity, int aligned_head_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 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 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 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)
 
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 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 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)
 
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 logits_copy_to_position (const float *__restrict src, float *__restrict dst, int position, int vocab_size)
 Copy logits to position-indexed location in output buffer.
 

Detailed Description

KV-cache helper kernels (head-major layout)

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

After changes: make test && make llamacpp-parity-full

Small, explicit helpers used by the runtime/orchestrator to maintain per-layer KV caches during autoregressive decoding.

Layout: k_cache[kv_head, token, aligned_head_dim] v_cache[kv_head, token, aligned_head_dim] with contiguous row-major storage and stride aligned_head_dim.

Definition in file kv_cache_kernels.c.

Function Documentation

◆ ck_local_fp32_to_bf16_row()

static void ck_local_fp32_to_bf16_row ( const float *  src,
uint16_t *  dst,
int  n 
)
inlinestatic

Definition at line 39 of file kv_cache_kernels.c.

40{
41 if (!src || !dst || n <= 0) {
42 return;
43 }
44 for (int i = 0; i < n; ++i) {
45 dst[i] = float_to_bf16(src[i]);
46 }
47}
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90

References float_to_bf16().

Referenced by kv_cache_store_batch_bf16(), and kv_cache_store_bf16().

◆ ck_local_fp32_to_fp16_row()

static void ck_local_fp32_to_fp16_row ( const float *  src,
uint16_t *  dst,
int  n 
)
inlinestatic

Definition at line 29 of file kv_cache_kernels.c.

30{
31 if (!src || !dst || n <= 0) {
32 return;
33 }
34 for (int i = 0; i < n; ++i) {
35 dst[i] = CK_FP32_TO_FP16(src[i]);
36 }
37}
#define CK_FP32_TO_FP16(x)

References CK_FP32_TO_FP16.

Referenced by kv_cache_store_batch_f16(), and kv_cache_store_f16().

◆ kv_cache_repack_head_major_inplace()

void kv_cache_repack_head_major_inplace ( float *  buf,
int  num_heads,
int  tokens,
int  cache_capacity,
int  aligned_head_dim 
)

Definition at line 49 of file kv_cache_kernels.c.

54{
55 if (!buf) {
56 return;
57 }
58 if (num_heads <= 0 || tokens <= 0 || cache_capacity <= 0 || aligned_head_dim <= 0) {
59 return;
60 }
61 if (tokens > cache_capacity) {
62 tokens = cache_capacity;
63 }
64 if (tokens == cache_capacity) {
65 return;
66 }
67
68 const size_t old_head_stride = (size_t)tokens * (size_t)aligned_head_dim;
69 const size_t new_head_stride = (size_t)cache_capacity * (size_t)aligned_head_dim;
70 const size_t bytes = (size_t)tokens * (size_t)aligned_head_dim * sizeof(float);
71
72 // Move head blocks from high to low to avoid overwriting source data
73 // for heads that have not yet been moved.
74 for (int h = num_heads - 1; h >= 0; --h) {
75 float *src = buf + (size_t)h * old_head_stride;
76 float *dst = buf + (size_t)h * new_head_stride;
77 memmove(dst, src, bytes);
78 }
79}

Referenced by qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), and qwen2_0_5b_decode_forward_prefill_impl().

◆ kv_cache_store()

void kv_cache_store ( float *__restrict  kv_cache_k,
float *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  layer,
int  pos,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 122 of file kv_cache_kernels.c.

131{
132 (void)layer;
134 kv_cache_k, kv_cache_v,
135 num_kv_heads,
136 pos,
137 max_seq_len,
138 head_dim,
139 head_dim);
140}
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)

References kv_cache_write_head_major().

◆ kv_cache_store_batch_bf16()

void kv_cache_store_batch_bf16 ( uint16_t *__restrict  kv_cache_k,
uint16_t *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  start_pos,
int  num_tokens,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 291 of file kv_cache_kernels.c.

300{
301 if (!kv_cache_k || !kv_cache_v || !k || !v ||
302 start_pos < 0 || num_tokens <= 0 || num_kv_heads <= 0 ||
303 head_dim <= 0 || max_seq_len <= 0 ||
304 start_pos > max_seq_len - num_tokens) {
305 return;
306 }
307
308 const size_t compact_head_stride = (size_t)num_tokens * (size_t)head_dim;
309 const size_t cache_head_stride = (size_t)max_seq_len * (size_t)head_dim;
310 for (int h = 0; h < num_kv_heads; ++h) {
311 const float *k_head = k + (size_t)h * compact_head_stride;
312 const float *v_head = v + (size_t)h * compact_head_stride;
313 uint16_t *k_head_cache = kv_cache_k + (size_t)h * cache_head_stride;
314 uint16_t *v_head_cache = kv_cache_v + (size_t)h * cache_head_stride;
315 for (int t = 0; t < num_tokens; ++t) {
316 const size_t src_offset = (size_t)t * (size_t)head_dim;
317 const size_t dst_offset = (size_t)(start_pos + t) * (size_t)head_dim;
318 ck_local_fp32_to_bf16_row(k_head + src_offset, k_head_cache + dst_offset, head_dim);
319 ck_local_fp32_to_bf16_row(v_head + src_offset, v_head_cache + dst_offset, head_dim);
320 }
321 }
322}
static void ck_local_fp32_to_bf16_row(const float *src, uint16_t *dst, int n)

References ck_local_fp32_to_bf16_row().

◆ kv_cache_store_batch_f16()

void kv_cache_store_batch_f16 ( uint16_t *__restrict  kv_cache_k,
uint16_t *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  start_pos,
int  num_tokens,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 255 of file kv_cache_kernels.c.

264{
265 if (!kv_cache_k || !kv_cache_v || !k || !v) {
266 return;
267 }
268 if (start_pos < 0 || num_tokens <= 0 || num_kv_heads <= 0 ||
269 head_dim <= 0 || max_seq_len <= 0 ||
270 start_pos > max_seq_len - num_tokens) {
271 return;
272 }
273
274 const size_t compact_head_stride = (size_t)num_tokens * (size_t)head_dim;
275 const size_t cache_head_stride = (size_t)max_seq_len * (size_t)head_dim;
276
277 for (int h = 0; h < num_kv_heads; ++h) {
278 const float *k_head = k + (size_t)h * compact_head_stride;
279 const float *v_head = v + (size_t)h * compact_head_stride;
280 uint16_t *k_head_cache = kv_cache_k + (size_t)h * cache_head_stride;
281 uint16_t *v_head_cache = kv_cache_v + (size_t)h * cache_head_stride;
282 for (int t = 0; t < num_tokens; ++t) {
283 const size_t src_offset = (size_t)t * (size_t)head_dim;
284 const size_t dst_offset = (size_t)(start_pos + t) * (size_t)head_dim;
285 ck_local_fp32_to_fp16_row(k_head + src_offset, k_head_cache + dst_offset, head_dim);
286 ck_local_fp32_to_fp16_row(v_head + src_offset, v_head_cache + dst_offset, head_dim);
287 }
288 }
289}
static void ck_local_fp32_to_fp16_row(const float *src, uint16_t *dst, int n)

References ck_local_fp32_to_fp16_row().

◆ kv_cache_store_batch_f32()

void kv_cache_store_batch_f32 ( float *__restrict  kv_cache_k,
float *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  start_pos,
int  num_tokens,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 223 of file kv_cache_kernels.c.

232{
233 if (!kv_cache_k || !kv_cache_v || !k || !v ||
234 start_pos < 0 || num_tokens <= 0 || num_kv_heads <= 0 ||
235 head_dim <= 0 || max_seq_len <= 0 ||
236 start_pos > max_seq_len - num_tokens) {
237 return;
238 }
239
240 const size_t compact_head_stride = (size_t)num_tokens * (size_t)head_dim;
241 const size_t cache_head_stride = (size_t)max_seq_len * (size_t)head_dim;
242 const size_t token_bytes = (size_t)num_tokens * (size_t)head_dim * sizeof(float);
243 for (int h = 0; h < num_kv_heads; ++h) {
244 const float *k_head = k + (size_t)h * compact_head_stride;
245 const float *v_head = v + (size_t)h * compact_head_stride;
246 float *k_head_cache = kv_cache_k + (size_t)h * cache_head_stride
247 + (size_t)start_pos * (size_t)head_dim;
248 float *v_head_cache = kv_cache_v + (size_t)h * cache_head_stride
249 + (size_t)start_pos * (size_t)head_dim;
250 memcpy(k_head_cache, k_head, token_bytes);
251 memcpy(v_head_cache, v_head, token_bytes);
252 }
253}

◆ kv_cache_store_bf16()

void kv_cache_store_bf16 ( uint16_t *__restrict  kv_cache_k,
uint16_t *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  layer,
int  pos,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 195 of file kv_cache_kernels.c.

204{
205 (void)layer;
206 if (!kv_cache_k || !kv_cache_v || !k || !v ||
207 num_kv_heads <= 0 || pos < 0 || pos >= max_seq_len ||
208 head_dim <= 0 || max_seq_len <= 0) {
209 return;
210 }
211
212 const size_t head_stride = (size_t)max_seq_len * (size_t)head_dim;
213 for (int h = 0; h < num_kv_heads; ++h) {
214 const float *k_src = k + (size_t)h * (size_t)head_dim;
215 const float *v_src = v + (size_t)h * (size_t)head_dim;
216 uint16_t *k_dst = kv_cache_k + (size_t)h * head_stride + (size_t)pos * (size_t)head_dim;
217 uint16_t *v_dst = kv_cache_v + (size_t)h * head_stride + (size_t)pos * (size_t)head_dim;
218 ck_local_fp32_to_bf16_row(k_src, k_dst, head_dim);
219 ck_local_fp32_to_bf16_row(v_src, v_dst, head_dim);
220 }
221}

References ck_local_fp32_to_bf16_row().

◆ kv_cache_store_f16()

void kv_cache_store_f16 ( uint16_t *__restrict  kv_cache_k,
uint16_t *__restrict  kv_cache_v,
const float *__restrict  k,
const float *__restrict  v,
int  layer,
int  pos,
int  num_kv_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 161 of file kv_cache_kernels.c.

170{
171 (void)layer;
172 if (!kv_cache_k || !kv_cache_v || !k || !v) {
173 return;
174 }
175 if (num_kv_heads <= 0 || pos < 0 || head_dim <= 0 || max_seq_len <= 0) {
176 return;
177 }
178 if (pos >= max_seq_len) {
179 return;
180 }
181
182 const size_t head_stride = (size_t)max_seq_len * (size_t)head_dim;
183 const size_t token_stride = (size_t)head_dim;
184
185 for (int h = 0; h < num_kv_heads; ++h) {
186 const float *k_src = k + (size_t)h * token_stride;
187 const float *v_src = v + (size_t)h * token_stride;
188 uint16_t *k_dst = kv_cache_k + (size_t)h * head_stride + (size_t)pos * token_stride;
189 uint16_t *v_dst = kv_cache_v + (size_t)h * head_stride + (size_t)pos * token_stride;
190 ck_local_fp32_to_fp16_row(k_src, k_dst, head_dim);
191 ck_local_fp32_to_fp16_row(v_src, v_dst, head_dim);
192 }
193}

References ck_local_fp32_to_fp16_row().

◆ kv_cache_store_shared_q()

void kv_cache_store_shared_q ( float *__restrict  kv_cache_k,
float *__restrict  kv_cache_v,
const float *__restrict  q,
int  layer,
int  pos,
int  num_heads,
int  head_dim,
int  max_seq_len 
)

Definition at line 142 of file kv_cache_kernels.c.

150{
151 (void)layer;
153 kv_cache_k, kv_cache_v,
154 num_heads,
155 pos,
156 max_seq_len,
157 head_dim,
158 head_dim);
159}

References kv_cache_write_head_major().

◆ kv_cache_write_head_major()

void kv_cache_write_head_major ( const float *__restrict  k_token,
const float *__restrict  v_token,
float *__restrict  k_cache,
float *__restrict  v_cache,
int  num_kv_heads,
int  token_index,
int  cache_capacity,
int  head_dim,
int  aligned_head_dim 
)

Definition at line 81 of file kv_cache_kernels.c.

90{
91 if (!k_token || !v_token || !k_cache || !v_cache) {
92 return;
93 }
94 if (num_kv_heads <= 0 || token_index < 0 || cache_capacity <= 0) {
95 return;
96 }
97 if (token_index >= cache_capacity || head_dim <= 0 || aligned_head_dim <= 0) {
98 return;
99 }
100
101 const size_t head_stride = (size_t)cache_capacity * (size_t)aligned_head_dim;
102 const size_t token_stride = (size_t)aligned_head_dim;
103
104 for (int h = 0; h < num_kv_heads; ++h) {
105 const float *k_src = k_token + (size_t)h * token_stride;
106 const float *v_src = v_token + (size_t)h * token_stride;
107
108 float *k_dst = k_cache + (size_t)h * head_stride + (size_t)token_index * token_stride;
109 float *v_dst = v_cache + (size_t)h * head_stride + (size_t)token_index * token_stride;
110
111 for (int d = 0; d < head_dim; ++d) {
112 k_dst[d] = k_src[d];
113 v_dst[d] = v_src[d];
114 }
115 for (int d = head_dim; d < aligned_head_dim; ++d) {
116 k_dst[d] = 0.0f;
117 v_dst[d] = 0.0f;
118 }
119 }
120}

Referenced by ck_layer_forward_rmsnorm_swiglu_decode(), ck_layer_forward_rmsnorm_swiglu_decode_fused(), ck_layer_forward_rmsnorm_swiglu_decode_fused_attn_impl(), ck_layer_forward_rmsnorm_swiglu_decode_q4_k(), ck_layer_forward_rmsnorm_swiglu_decode_quant(), kv_cache_store(), kv_cache_store_shared_q(), mega_fused_attention_decode_workspace(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), and qwen2_0_5b_decode_layer_9_decode().

◆ logits_copy_to_position()

void logits_copy_to_position ( const float *__restrict  src,
float *__restrict  dst,
int  position,
int  vocab_size 
)

Copy logits to position-indexed location in output buffer.

Used in decode mode to copy single-token logits from position 0 to the correct sequence position. This moves buffer management logic from codegen to the IR layer, making codegen "dumb" - just emit kernel calls, no runtime if-statements.

Parameters
srcSource logits buffer (single token) [vocab_size]
dstDestination logits buffer [max_seq_len, vocab_size]
positionToken position index (0-based)
vocab_sizeNumber of logits per token

Definition at line 337 of file kv_cache_kernels.c.

341{
342 if (!src || !dst || position < 0 || vocab_size <= 0) {
343 return;
344 }
345
346 // Copy logits to dst[position * vocab_size : (position+1) * vocab_size]
347 // Use memmove for safety in case src and dst overlap (e.g., src == dst)
348 float *dst_pos = dst + (size_t)position * (size_t)vocab_size;
349 memmove(dst_pos, src, (size_t)vocab_size * sizeof(float));
350}
int vocab_size
Definition true_bpe.h:193

References vocab_size.