1use ferrum_types::{FerrumError, Result};
19
20use crate::backend::{Backend, BackendInt8KvOps, KvCache, KvCacheQuant};
21use ferrum_interfaces::kv_dtype::{KvDtypeKind, KvFp16, KvInt8};
22
23#[allow(clippy::too_many_arguments)]
25pub trait KvLayer<B: Backend>: KvDtypeKind {
26 type Layer: Send + Sync;
28
29 fn alloc_paged(
31 max_blocks_per_seq: usize,
32 block_size: usize,
33 num_kv_heads: usize,
34 head_dim: usize,
35 ) -> Self::Layer;
36
37 fn alloc_contig(capacity: usize, num_kv_heads: usize, head_dim: usize) -> Self::Layer;
39
40 fn len(layer: &Self::Layer) -> usize;
42 fn set_len(layer: &mut Self::Layer, new_len: usize);
43 fn capacity(layer: &Self::Layer) -> usize;
44 fn block_size(layer: &Self::Layer) -> usize;
45 fn num_kv_heads(layer: &Self::Layer) -> usize;
46 fn head_dim(layer: &Self::Layer) -> usize;
47 fn block_table(layer: &Self::Layer) -> Option<&B::Buffer>;
48 fn block_table_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer>;
49 fn context_lens(layer: &Self::Layer) -> Option<&B::Buffer>;
50 fn context_lens_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer>;
51 fn paged_block_indices(layer: &Self::Layer) -> &[u32];
52 fn paged_block_indices_mut(layer: &mut Self::Layer) -> &mut Vec<u32>;
53
54 fn ensure_contig_capacity(
57 _ctx: &mut B::Context,
58 _layer: &mut Self::Layer,
59 _required_capacity: usize,
60 _num_kv_heads: usize,
61 _head_dim: usize,
62 ) -> Result<()> {
63 Ok(())
64 }
65
66 fn is_paged(layer: &Self::Layer) -> bool {
67 Self::block_size(layer) > 0
68 }
69
70 fn paged_write(
74 ctx: &mut B::Context,
75 layer: &mut Self::Layer,
76 qkv: &B::Buffer,
77 q_norm_w: &B::Buffer,
78 k_norm_w: &B::Buffer,
79 cos: &B::Buffer,
80 sin: &B::Buffer,
81 q_out: &mut B::Buffer,
82 k_scratch: &mut B::Buffer,
83 v_scratch: &mut B::Buffer,
84 pool_k: &mut B::Buffer,
85 pool_v: &mut B::Buffer,
86 tokens: usize,
87 num_q_heads: usize,
88 num_kv_heads: usize,
89 head_dim: usize,
90 pos_offset: usize,
91 eps: f32,
92 qk_mode: i32,
93 ) -> Result<()>;
94
95 fn paged_decode_attention(
99 ctx: &mut B::Context,
100 layer: &mut Self::Layer,
101 q: &B::Buffer,
102 pool_k: &B::Buffer,
103 pool_v: &B::Buffer,
104 output: &mut B::Buffer,
105 num_q_heads: usize,
106 num_kv_heads: usize,
107 head_dim: usize,
108 final_kv_len: usize,
109 tokens: usize,
110 ) -> Result<()>;
111
112 fn contig_write(
116 _ctx: &mut B::Context,
117 _layer: &mut Self::Layer,
118 _qkv: &B::Buffer,
119 _q_norm_w: &B::Buffer,
120 _k_norm_w: &B::Buffer,
121 _cos: &B::Buffer,
122 _sin: &B::Buffer,
123 _q_out: &mut B::Buffer,
124 _k_scratch: &mut B::Buffer,
125 _v_scratch: &mut B::Buffer,
126 _q_buf: &mut B::Buffer,
127 _k_buf: &mut B::Buffer,
128 _v_buf: &mut B::Buffer,
129 _tokens: usize,
130 _num_q_heads: usize,
131 _num_kv_heads: usize,
132 _head_dim: usize,
133 _pos_offset: usize,
134 _eps: f32,
135 _qk_mode: i32,
136 ) -> Result<()> {
137 unimplemented!("contig_write: not supported for this K dtype")
138 }
139
140 fn contig_decode_attention(
142 _ctx: &mut B::Context,
143 _layer: &Self::Layer,
144 _q: &B::Buffer,
145 _output: &mut B::Buffer,
146 _attn_cfg: crate::backend::AttnConfig,
147 _tokens: usize,
148 _pos_offset: usize,
149 ) -> Result<()> {
150 unimplemented!("contig_decode_attention: not supported for this K dtype")
151 }
152}
153
154impl<B: Backend + crate::backend::BackendPagedKv> KvLayer<B> for KvFp16 {
159 type Layer = KvCache<B, KvFp16>;
160
161 fn alloc_paged(
162 max_blocks_per_seq: usize,
163 block_size: usize,
164 num_kv_heads: usize,
165 head_dim: usize,
166 ) -> Self::Layer {
167 let block_table = B::alloc_typed(crate::backend::Dtype::U32, max_blocks_per_seq);
168 let mut context_lens = B::alloc_typed(crate::backend::Dtype::U32, 1);
169 let mut bt_ctx = B::new_context();
170 B::write_typed::<u32>(&mut bt_ctx, &mut context_lens, &[0u32]);
171 B::sync(&mut bt_ctx);
172 KvCache {
173 k: B::alloc(1),
174 v: B::alloc(1),
175 len: 0,
176 capacity: max_blocks_per_seq * block_size,
177 num_kv_heads,
178 head_dim,
179 block_size,
180 block_table: Some(block_table),
181 context_lens: Some(context_lens),
182 paged_block_indices: Vec::new(),
183 _kv_dtype: std::marker::PhantomData,
184 }
185 }
186
187 fn alloc_contig(capacity: usize, num_kv_heads: usize, head_dim: usize) -> Self::Layer {
188 KvCache {
189 k: B::alloc(num_kv_heads * capacity * head_dim),
190 v: B::alloc(num_kv_heads * capacity * head_dim),
191 len: 0,
192 capacity,
193 num_kv_heads,
194 head_dim,
195 block_size: 0,
196 block_table: None,
197 context_lens: None,
198 paged_block_indices: Vec::new(),
199 _kv_dtype: std::marker::PhantomData,
200 }
201 }
202
203 fn len(layer: &Self::Layer) -> usize {
204 layer.len
205 }
206 fn set_len(layer: &mut Self::Layer, new_len: usize) {
207 layer.len = new_len;
208 }
209 fn capacity(layer: &Self::Layer) -> usize {
210 layer.capacity
211 }
212 fn block_size(layer: &Self::Layer) -> usize {
213 layer.block_size
214 }
215 fn num_kv_heads(layer: &Self::Layer) -> usize {
216 layer.num_kv_heads
217 }
218 fn head_dim(layer: &Self::Layer) -> usize {
219 layer.head_dim
220 }
221 fn block_table(layer: &Self::Layer) -> Option<&B::Buffer> {
222 layer.block_table.as_ref()
223 }
224 fn block_table_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer> {
225 layer.block_table.as_mut()
226 }
227 fn context_lens(layer: &Self::Layer) -> Option<&B::Buffer> {
228 layer.context_lens.as_ref()
229 }
230 fn context_lens_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer> {
231 layer.context_lens.as_mut()
232 }
233 fn paged_block_indices(layer: &Self::Layer) -> &[u32] {
234 &layer.paged_block_indices
235 }
236 fn paged_block_indices_mut(layer: &mut Self::Layer) -> &mut Vec<u32> {
237 &mut layer.paged_block_indices
238 }
239
240 fn ensure_contig_capacity(
241 ctx: &mut B::Context,
242 layer: &mut Self::Layer,
243 required_capacity: usize,
244 num_kv_heads: usize,
245 head_dim: usize,
246 ) -> Result<()> {
247 if layer.block_size > 0 || required_capacity <= layer.capacity {
248 return Ok(());
249 }
250 let old_capacity = layer.capacity;
251 let old_len = layer.len;
252 let mut new_k = B::alloc(num_kv_heads * required_capacity * head_dim);
253 let mut new_v = B::alloc(num_kv_heads * required_capacity * head_dim);
254 if old_len > 0 {
255 let copy_len = old_len * head_dim;
256 for head in 0..num_kv_heads {
257 let old_offset = head * old_capacity * head_dim;
258 let new_offset = head * required_capacity * head_dim;
259 B::copy_slice(ctx, &layer.k, old_offset, &mut new_k, new_offset, copy_len);
260 B::copy_slice(ctx, &layer.v, old_offset, &mut new_v, new_offset, copy_len);
261 }
262 B::sync(ctx);
263 }
264 layer.k = new_k;
265 layer.v = new_v;
266 layer.capacity = required_capacity;
267 Ok(())
268 }
269
270 fn paged_write(
271 ctx: &mut B::Context,
272 layer: &mut Self::Layer,
273 qkv: &B::Buffer,
274 q_norm_w: &B::Buffer,
275 k_norm_w: &B::Buffer,
276 cos: &B::Buffer,
277 sin: &B::Buffer,
278 q_out: &mut B::Buffer,
279 _k_scratch: &mut B::Buffer,
280 _v_scratch: &mut B::Buffer,
281 pool_k: &mut B::Buffer,
282 pool_v: &mut B::Buffer,
283 tokens: usize,
284 num_q_heads: usize,
285 num_kv_heads: usize,
286 head_dim: usize,
287 pos_offset: usize,
288 eps: f32,
289 qk_mode: i32,
290 ) -> Result<()> {
291 let block_size = layer.block_size;
292 let cache_len_before = layer.len;
293 let num_blocks_per_seq = layer.capacity / block_size;
294 let bt = layer
295 .block_table
296 .as_ref()
297 .ok_or_else(|| FerrumError::model("FP16 paged_write: missing block_table"))?;
298 B::split_qkv_norm_rope_into_paged_cache(
299 ctx,
300 qkv,
301 0,
302 q_norm_w,
303 k_norm_w,
304 cos,
305 sin,
306 q_out,
307 0,
308 pool_k,
309 pool_v,
310 bt,
311 tokens,
312 num_q_heads,
313 num_kv_heads,
314 head_dim,
315 pos_offset,
316 eps,
317 qk_mode,
318 cache_len_before,
319 block_size,
320 num_blocks_per_seq,
321 )
322 }
323
324 fn paged_decode_attention(
325 ctx: &mut B::Context,
326 layer: &mut Self::Layer,
327 q: &B::Buffer,
328 pool_k: &B::Buffer,
329 pool_v: &B::Buffer,
330 output: &mut B::Buffer,
331 num_q_heads: usize,
332 num_kv_heads: usize,
333 head_dim: usize,
334 final_kv_len: usize,
335 tokens: usize,
336 ) -> Result<()> {
337 let block_size = layer.block_size;
338 let num_blocks_per_seq = layer.capacity / block_size;
339 let bt_ptr = layer
340 .block_table
341 .as_ref()
342 .ok_or_else(|| FerrumError::model("FP16 paged_decode: missing block_table"))?
343 as *const B::Buffer;
344 let cl_buf = layer
345 .context_lens
346 .as_mut()
347 .ok_or_else(|| FerrumError::model("FP16 paged_decode: missing context_lens"))?;
348 B::write_typed::<u32>(ctx, cl_buf, &[final_kv_len as u32]);
349 let bt = unsafe { &*bt_ptr };
351 let cl = layer.context_lens.as_ref().unwrap();
352 B::paged_decode_attention(
353 ctx,
354 q,
355 pool_k,
356 pool_v,
357 output,
358 bt,
359 cl,
360 1,
361 num_q_heads,
362 num_kv_heads,
363 head_dim,
364 block_size,
365 num_blocks_per_seq,
366 tokens,
367 )
368 }
369
370 fn contig_write(
371 ctx: &mut B::Context,
372 layer: &mut Self::Layer,
373 qkv: &B::Buffer,
374 q_norm_w: &B::Buffer,
375 k_norm_w: &B::Buffer,
376 cos: &B::Buffer,
377 sin: &B::Buffer,
378 q_out: &mut B::Buffer,
379 k_scratch: &mut B::Buffer,
380 v_scratch: &mut B::Buffer,
381 q_buf: &mut B::Buffer,
382 k_buf: &mut B::Buffer,
383 v_buf: &mut B::Buffer,
384 tokens: usize,
385 num_q_heads: usize,
386 num_kv_heads: usize,
387 head_dim: usize,
388 pos_offset: usize,
389 eps: f32,
390 qk_mode: i32,
391 ) -> Result<()> {
392 let cache_len_before = layer.len;
393 let cache_capacity = layer.capacity;
394 let used_into_cache = B::split_qkv_norm_rope_into_cache(
395 ctx,
396 qkv,
397 q_norm_w,
398 k_norm_w,
399 cos,
400 sin,
401 q_out,
402 &mut layer.k,
403 &mut layer.v,
404 tokens,
405 num_q_heads,
406 num_kv_heads,
407 head_dim,
408 pos_offset,
409 eps,
410 qk_mode,
411 cache_len_before,
412 cache_capacity,
413 )
414 .is_ok();
415 if used_into_cache {
416 return Ok(());
417 }
418 let used_fused_qkv = B::split_qkv_norm_rope(
419 ctx,
420 qkv,
421 q_norm_w,
422 k_norm_w,
423 cos,
424 sin,
425 q_out,
426 k_scratch,
427 v_scratch,
428 tokens,
429 num_q_heads,
430 num_kv_heads,
431 head_dim,
432 pos_offset,
433 eps,
434 qk_mode,
435 )
436 .is_ok();
437 if !used_fused_qkv {
438 let q_dim = num_q_heads * head_dim;
439 let kv_dim = num_kv_heads * head_dim;
440 B::split_qkv(ctx, qkv, q_buf, k_buf, v_buf, tokens, q_dim, kv_dim);
441 B::qk_norm_rope(
442 ctx,
443 q_buf,
444 q_norm_w,
445 cos,
446 sin,
447 q_out,
448 tokens,
449 num_q_heads,
450 head_dim,
451 pos_offset,
452 eps,
453 qk_mode,
454 );
455 B::qk_norm_rope(
456 ctx,
457 k_buf,
458 k_norm_w,
459 cos,
460 sin,
461 k_scratch,
462 tokens,
463 num_kv_heads,
464 head_dim,
465 pos_offset,
466 eps,
467 qk_mode,
468 );
469 B::qk_norm_rope(
470 ctx,
471 v_buf,
472 q_norm_w,
473 cos,
474 sin,
475 v_scratch,
476 tokens,
477 num_kv_heads,
478 head_dim,
479 pos_offset,
480 eps,
481 0,
482 );
483 }
484 B::kv_cache_append_head_major(
485 ctx,
486 &mut layer.k,
487 &mut layer.v,
488 cache_len_before,
489 cache_capacity,
490 k_scratch,
491 v_scratch,
492 tokens,
493 num_kv_heads,
494 head_dim,
495 );
496 Ok(())
497 }
498
499 fn contig_decode_attention(
500 ctx: &mut B::Context,
501 layer: &Self::Layer,
502 q: &B::Buffer,
503 output: &mut B::Buffer,
504 attn_cfg: crate::backend::AttnConfig,
505 tokens: usize,
506 pos_offset: usize,
507 ) -> Result<()> {
508 let kv_len = layer.len;
509 B::flash_attention(
510 ctx, q, &layer.k, &layer.v, output, 1, tokens, kv_len, pos_offset, &attn_cfg,
511 );
512 Ok(())
513 }
514}
515
516impl<B: Backend + BackendInt8KvOps> KvLayer<B> for KvInt8 {
521 type Layer = KvCacheQuant<B, KvInt8>;
522
523 fn alloc_paged(
524 max_blocks_per_seq: usize,
525 block_size: usize,
526 num_kv_heads: usize,
527 head_dim: usize,
528 ) -> Self::Layer {
529 B::alloc_paged_int8_layer(max_blocks_per_seq, block_size, num_kv_heads, head_dim)
530 }
531
532 fn alloc_contig(_capacity: usize, _num_kv_heads: usize, _head_dim: usize) -> Self::Layer {
533 panic!("KvInt8::alloc_contig: INT8 KV is paged-only")
534 }
535
536 fn len(layer: &Self::Layer) -> usize {
537 layer.len
538 }
539 fn set_len(layer: &mut Self::Layer, new_len: usize) {
540 layer.len = new_len;
541 }
542 fn capacity(layer: &Self::Layer) -> usize {
543 layer.capacity
544 }
545 fn block_size(layer: &Self::Layer) -> usize {
546 layer.block_size
547 }
548 fn num_kv_heads(layer: &Self::Layer) -> usize {
549 layer.num_kv_heads
550 }
551 fn head_dim(layer: &Self::Layer) -> usize {
552 layer.head_dim
553 }
554 fn block_table(layer: &Self::Layer) -> Option<&B::Buffer> {
555 layer.block_table.as_ref()
556 }
557 fn block_table_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer> {
558 layer.block_table.as_mut()
559 }
560 fn context_lens(layer: &Self::Layer) -> Option<&B::Buffer> {
561 layer.context_lens.as_ref()
562 }
563 fn context_lens_mut(layer: &mut Self::Layer) -> Option<&mut B::Buffer> {
564 layer.context_lens.as_mut()
565 }
566 fn paged_block_indices(layer: &Self::Layer) -> &[u32] {
567 &layer.paged_block_indices
568 }
569 fn paged_block_indices_mut(layer: &mut Self::Layer) -> &mut Vec<u32> {
570 &mut layer.paged_block_indices
571 }
572
573 fn paged_write(
574 ctx: &mut B::Context,
575 layer: &mut Self::Layer,
576 qkv: &B::Buffer,
577 q_norm_w: &B::Buffer,
578 k_norm_w: &B::Buffer,
579 cos: &B::Buffer,
580 sin: &B::Buffer,
581 q_out: &mut B::Buffer,
582 k_scratch: &mut B::Buffer,
583 v_scratch: &mut B::Buffer,
584 _pool_k: &mut B::Buffer,
585 _pool_v: &mut B::Buffer,
586 tokens: usize,
587 num_q_heads: usize,
588 num_kv_heads: usize,
589 head_dim: usize,
590 pos_offset: usize,
591 eps: f32,
592 qk_mode: i32,
593 ) -> Result<()> {
594 B::split_qkv_norm_rope(
596 ctx,
597 qkv,
598 q_norm_w,
599 k_norm_w,
600 cos,
601 sin,
602 q_out,
603 k_scratch,
604 v_scratch,
605 tokens,
606 num_q_heads,
607 num_kv_heads,
608 head_dim,
609 pos_offset,
610 eps,
611 qk_mode,
612 )?;
613 let cache_len_before = layer.len;
618 let block_size = layer.block_size;
619 let paged_indices: Vec<u32> = layer.paged_block_indices.clone();
622 B::int8_kv_append_paged(
623 ctx,
624 k_scratch,
625 v_scratch,
626 &mut layer.k,
627 &mut layer.v,
628 &mut layer.k_scales,
629 &mut layer.v_scales,
630 &paged_indices,
631 cache_len_before,
632 tokens,
633 block_size,
634 num_kv_heads,
635 head_dim,
636 )
637 }
638
639 fn paged_decode_attention(
640 ctx: &mut B::Context,
641 layer: &mut Self::Layer,
642 q: &B::Buffer,
643 _pool_k: &B::Buffer,
644 _pool_v: &B::Buffer,
645 output: &mut B::Buffer,
646 num_q_heads: usize,
647 num_kv_heads: usize,
648 head_dim: usize,
649 final_kv_len: usize,
650 _tokens: usize,
651 ) -> Result<()> {
652 let block_size = layer.block_size;
653 let cl_buf = layer
654 .context_lens
655 .as_mut()
656 .ok_or_else(|| FerrumError::model("INT8 paged_decode: missing context_lens"))?;
657 B::write_typed::<u32>(ctx, cl_buf, &[final_kv_len as u32]);
658 let bt = layer
659 .block_table
660 .as_ref()
661 .ok_or_else(|| FerrumError::model("INT8 paged_decode: missing block_table"))?;
662 let scale = (head_dim as f32).sqrt().recip();
663 B::int8_paged_decode_attention(
664 ctx,
665 q,
666 &layer.k,
667 &layer.v,
668 &layer.k_scales,
669 &layer.v_scales,
670 bt,
671 output,
672 num_q_heads,
673 num_kv_heads,
674 head_dim,
675 final_kv_len,
676 block_size,
677 scale,
678 )
679 }
680}