Skip to main content

ferrum_kernels/backend/
kv_layer.rs

1//! `KvLayer<B>` — per-K-dtype trait that picks the cache layout type and
2//! the K-specific paged write / read launchers (Dim 5 PR C trait-based
3//! dispatch).
4//!
5//! ## Why a trait, not an enum
6//!
7//! - `K` carries an associated `Layer` type — FP16 → `KvCache<B, KvFp16>`,
8//!   INT8 → `KvCacheQuant<B, KvInt8>`.
9//! - K-specific launchers (paged write + paged decode attention; contig
10//!   write + contig decode for FP16) are trait methods. The model bound
11//!   `where K: KvLayer<B>` lets `K::method(layer, ...)` dispatch directly
12//!   to the right backend launcher per (B, K) at monomorphization time —
13//!   no runtime tag, no enum match, no panicking accessors.
14//! - `LlamaFamilyModel<CpuBackend, KvInt8>` is a compile error because
15//!   `KvInt8: KvLayer<CpuBackend>` doesn't hold (CPU backend has no
16//!   `BackendInt8KvOps` impl).
17
18use ferrum_types::{FerrumError, Result};
19
20use crate::backend::{Backend, BackendInt8KvOps, KvCache, KvCacheQuant};
21use ferrum_interfaces::kv_dtype::{KvDtypeKind, KvFp16, KvInt8};
22
23/// Per-K-dtype dispatch trait.
24#[allow(clippy::too_many_arguments)]
25pub trait KvLayer<B: Backend>: KvDtypeKind {
26    /// Per-layer cache type (FP16 → `KvCache`, INT8 → `KvCacheQuant`).
27    type Layer: Send + Sync;
28
29    /// Allocate a paged cache layer for one sequence.
30    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    /// Allocate a contiguous cache layer (FP16 only; INT8 panics).
38    fn alloc_contig(capacity: usize, num_kv_heads: usize, head_dim: usize) -> Self::Layer;
39
40    // Metadata accessors (variant-agnostic).
41    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    /// Ensure a contiguous cache has physical room for `required_capacity`
55    /// KV positions. Paged caches and fixed-size cache dtypes may no-op.
56    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    /// Paged write: split QKV → norm → RoPE → write K/V into the paged
71    /// pool. FP16 uses `B::split_qkv_norm_rope_into_paged_cache`. INT8
72    /// uses `B::split_qkv_norm_rope` + `B::int8_kv_append_paged`.
73    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    /// Paged decode attention. Reads from the per-layer cache, writes the
96    /// attended output to `output`. FP16 reads from `pool_k`/`pool_v`;
97    /// INT8 reads from layer-internal INT8 buffers (pool args ignored).
98    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    /// Contig write: FP16 only. INT8 inherits the panic default —
113    /// `KvInt8::alloc_contig` panics in `ensure_kv`, so this branch is
114    /// dead code on the INT8 path.
115    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    /// Contig decode attention: FP16 only.
141    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
154// ─────────────────────────────────────────────────────────────────────
155// FP16 impl
156// ─────────────────────────────────────────────────────────────────────
157
158impl<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        // SAFETY: block_table outlives the call.
350        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
516// ─────────────────────────────────────────────────────────────────────
517// INT8 impl
518// ─────────────────────────────────────────────────────────────────────
519
520impl<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        // 1. split + norm + RoPE → FP16 head-major scratch (k/v_scratch).
595        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        // 2. quantize FP16 → INT8 + per-token scales, paged append.
614        // `paged_block_indices` is the host-side mirror populated at
615        // `ensure_kv` time — passing it directly avoids the D2H + sync
616        // barrier that would otherwise dominate per-token overhead.
617        let cache_len_before = layer.len;
618        let block_size = layer.block_size;
619        // Clone the host indices (small Vec<u32>) so we don't hold a
620        // borrow on `layer` while passing &mut layer.k/v/scales below.
621        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}