← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
kv_cache_kernels.c
Go to the documentation of this file.
1/**
2 * @file kv_cache_kernels.c
3 * @brief KV-cache helper kernels (head-major layout)
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * Small, explicit helpers used by the runtime/orchestrator to maintain
15 * per-layer KV caches during autoregressive decoding.
16 *
17 * Layout:
18 * k_cache[kv_head, token, aligned_head_dim]
19 * v_cache[kv_head, token, aligned_head_dim]
20 * with contiguous row-major storage and stride aligned_head_dim.
21 */
22
23#include "ckernel_engine.h"
24#include "bf16_utils.h"
25
26#include <stddef.h>
27#include <string.h>
28
29static inline void ck_local_fp32_to_fp16_row(const float *src, uint16_t *dst, int n)
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}
38
39static inline void ck_local_fp32_to_bf16_row(const float *src, uint16_t *dst, int n)
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}
48
50 int num_heads,
51 int tokens,
52 int cache_capacity,
53 int aligned_head_dim)
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}
80
81void kv_cache_write_head_major(const float *__restrict k_token,
82 const float *__restrict v_token,
83 float *__restrict k_cache,
84 float *__restrict v_cache,
85 int num_kv_heads,
86 int token_index,
87 int cache_capacity,
88 int head_dim,
89 int aligned_head_dim)
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}
121
122void kv_cache_store(float *__restrict kv_cache_k,
123 float *__restrict kv_cache_v,
124 const float *__restrict k,
125 const float *__restrict v,
126 int layer,
127 int pos,
128 int num_kv_heads,
129 int head_dim,
130 int max_seq_len)
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}
141
142void kv_cache_store_shared_q(float *__restrict kv_cache_k,
143 float *__restrict kv_cache_v,
144 const float *__restrict q,
145 int layer,
146 int pos,
147 int num_heads,
148 int head_dim,
149 int max_seq_len)
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}
160
161void kv_cache_store_f16(uint16_t *__restrict kv_cache_k,
162 uint16_t *__restrict kv_cache_v,
163 const float *__restrict k,
164 const float *__restrict v,
165 int layer,
166 int pos,
167 int num_kv_heads,
168 int head_dim,
169 int max_seq_len)
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}
194
195void kv_cache_store_bf16(uint16_t *__restrict kv_cache_k,
196 uint16_t *__restrict kv_cache_v,
197 const float *__restrict k,
198 const float *__restrict v,
199 int layer,
200 int pos,
201 int num_kv_heads,
202 int head_dim,
203 int max_seq_len)
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}
222
223void kv_cache_store_batch_f32(float *__restrict kv_cache_k,
224 float *__restrict kv_cache_v,
225 const float *__restrict k,
226 const float *__restrict v,
227 int start_pos,
228 int num_tokens,
229 int num_kv_heads,
230 int head_dim,
231 int max_seq_len)
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}
254
255void kv_cache_store_batch_f16(uint16_t *__restrict kv_cache_k,
256 uint16_t *__restrict kv_cache_v,
257 const float *__restrict k,
258 const float *__restrict v,
259 int start_pos,
260 int num_tokens,
261 int num_kv_heads,
262 int head_dim,
263 int max_seq_len)
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}
290
291void kv_cache_store_batch_bf16(uint16_t *__restrict kv_cache_k,
292 uint16_t *__restrict kv_cache_v,
293 const float *__restrict k,
294 const float *__restrict v,
295 int start_pos,
296 int num_tokens,
297 int num_kv_heads,
298 int head_dim,
299 int max_seq_len)
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}
323
324/**
325 * @brief Copy logits to position-indexed location in output buffer.
326 *
327 * Used in decode mode to copy single-token logits from position 0 to
328 * the correct sequence position. This moves buffer management logic
329 * from codegen to the IR layer, making codegen "dumb" - just emit
330 * kernel calls, no runtime if-statements.
331 *
332 * @param src Source logits buffer (single token) [vocab_size]
333 * @param dst Destination logits buffer [max_seq_len, vocab_size]
334 * @param position Token position index (0-based)
335 * @param vocab_size Number of logits per token
336 */
337void logits_copy_to_position(const float *__restrict src,
338 float *__restrict dst,
339 int position,
340 int vocab_size)
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}
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90
#define CK_FP32_TO_FP16(x)
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_repack_head_major_inplace(float *buf, int num_heads, int tokens, int cache_capacity, int aligned_head_dim)
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_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)
static void ck_local_fp32_to_bf16_row(const float *src, uint16_t *dst, int n)
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_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 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(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 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.
static void ck_local_fp32_to_fp16_row(const float *src, uint16_t *dst, int n)
int vocab_size
Definition true_bpe.h:193