31 if (!src || !dst || n <= 0) {
34 for (
int i = 0; i < n; ++i) {
41 if (!src || !dst || n <= 0) {
44 for (
int i = 0; i < n; ++i) {
58 if (num_heads <= 0 || tokens <= 0 || cache_capacity <= 0 || aligned_head_dim <= 0) {
61 if (tokens > cache_capacity) {
62 tokens = cache_capacity;
64 if (tokens == cache_capacity) {
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);
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);
82 const float *__restrict v_token,
83 float *__restrict k_cache,
84 float *__restrict v_cache,
91 if (!k_token || !v_token || !k_cache || !v_cache) {
94 if (num_kv_heads <= 0 || token_index < 0 || cache_capacity <= 0) {
97 if (token_index >= cache_capacity || head_dim <= 0 || aligned_head_dim <= 0) {
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;
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;
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;
111 for (
int d = 0; d < head_dim; ++d) {
115 for (
int d = head_dim; d < aligned_head_dim; ++d) {
123 float *__restrict kv_cache_v,
124 const float *__restrict k,
125 const float *__restrict v,
134 kv_cache_k, kv_cache_v,
143 float *__restrict kv_cache_v,
144 const float *__restrict q,
153 kv_cache_k, kv_cache_v,
162 uint16_t *__restrict kv_cache_v,
163 const float *__restrict k,
164 const float *__restrict v,
172 if (!kv_cache_k || !kv_cache_v || !k || !v) {
175 if (num_kv_heads <= 0 || pos < 0 || head_dim <= 0 || max_seq_len <= 0) {
178 if (pos >= max_seq_len) {
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;
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;
196 uint16_t *__restrict kv_cache_v,
197 const float *__restrict k,
198 const float *__restrict v,
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) {
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;
224 float *__restrict kv_cache_v,
225 const float *__restrict k,
226 const float *__restrict v,
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) {
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);
256 uint16_t *__restrict kv_cache_v,
257 const float *__restrict k,
258 const float *__restrict v,
265 if (!kv_cache_k || !kv_cache_v || !k || !v) {
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) {
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;
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;
292 uint16_t *__restrict kv_cache_v,
293 const float *__restrict k,
294 const float *__restrict v,
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) {
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;
338 float *__restrict dst,
342 if (!src || !dst || position < 0 ||
vocab_size <= 0) {
348 float *dst_pos = dst + (size_t)position * (
size_t)
vocab_size;
349 memmove(dst_pos, src, (
size_t)
vocab_size *
sizeof(
float));
static uint16_t float_to_bf16(float f)
#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)