Skip to main content

model_artifact/gguf/
kv_cache.rs

1use super::GgufCompactMeta;
2
3impl GgufCompactMeta {
4    pub fn k_cache_bytes_per_token_f16(&self) -> Option<u64> {
5        GgufKvCacheQuant::f16().k_cache_bytes_per_token(self)
6    }
7
8    pub fn v_cache_bytes_per_token_f16(&self) -> Option<u64> {
9        GgufKvCacheQuant::f16().v_cache_bytes_per_token(self)
10    }
11
12    pub fn kv_cache_bytes_per_token_f16(&self) -> Option<u64> {
13        GgufKvCacheQuant::f16().kv_cache_bytes_per_token(self)
14    }
15
16    fn kv_cache_head_count(&self) -> Option<u32> {
17        // GLM-DSA uses absorbed MLA: cache one compressed KV group rather
18        // than one expanded vector for every attention head.
19        if self.architecture == "glm-dsa" {
20            Some(1)
21        } else {
22            self.effective_kv_head_count()
23        }
24    }
25
26    fn kv_cache_value_length(&self) -> u32 {
27        // The cached V row is the compressed KV latent. The regular
28        // attention value length describes the expanded per-head value.
29        if self.architecture == "glm-dsa" && self.kv_lora_rank > 0 {
30            self.kv_lora_rank
31        } else {
32            self.value_length
33        }
34    }
35}
36
37#[derive(Clone, Copy, Debug, Eq, PartialEq)]
38pub enum GgufKvCacheType {
39    F16,
40    Q8_0,
41    Q4_0,
42}
43
44impl GgufKvCacheType {
45    pub fn from_llama_arg(value: &str) -> Option<Self> {
46        match value.to_ascii_lowercase().as_str() {
47            "f16" => Some(Self::F16),
48            "q8_0" => Some(Self::Q8_0),
49            "q4_0" => Some(Self::Q4_0),
50            _ => None,
51        }
52    }
53
54    pub const fn as_llama_arg(self) -> &'static str {
55        match self {
56            Self::F16 => "f16",
57            Self::Q8_0 => "q8_0",
58            Self::Q4_0 => "q4_0",
59        }
60    }
61
62    fn block_shape(self) -> (u64, u64) {
63        match self {
64            Self::F16 => (1, 2),
65            Self::Q8_0 => (32, 34),
66            Self::Q4_0 => (32, 18),
67        }
68    }
69
70    fn bytes_for_elements(self, elements: u64) -> Option<u64> {
71        let (block_elements, block_bytes) = self.block_shape();
72        let blocks = elements
73            .checked_add(block_elements.checked_sub(1)?)?
74            .checked_div(block_elements)?;
75        blocks.checked_mul(block_bytes)
76    }
77}
78
79#[derive(Clone, Copy, Debug, Eq, PartialEq)]
80pub struct GgufKvCacheQuant {
81    pub k: GgufKvCacheType,
82    pub v: GgufKvCacheType,
83}
84
85impl GgufKvCacheQuant {
86    /// f16 K + f16 V — highest quality, largest KV cache.
87    pub const F16: Self = Self {
88        k: GgufKvCacheType::F16,
89        v: GgufKvCacheType::F16,
90    };
91
92    /// q8_0 K + q8_0 V — moderate compression.
93    pub const Q8_0: Self = Self {
94        k: GgufKvCacheType::Q8_0,
95        v: GgufKvCacheType::Q8_0,
96    };
97
98    /// q4_0 K + q4_0 V — most aggressive compression, smallest KV cache.
99    pub const Q4_0: Self = Self {
100        k: GgufKvCacheType::Q4_0,
101        v: GgufKvCacheType::Q4_0,
102    };
103
104    pub const fn new(k: GgufKvCacheType, v: GgufKvCacheType) -> Self {
105        Self { k, v }
106    }
107
108    pub const fn f16() -> Self {
109        Self::F16
110    }
111
112    /// Returns `true` if `self` uses more aggressive (smaller) quantisation
113    /// than `other`.
114    pub const fn is_more_aggressive_than(self, other: Self) -> bool {
115        Self::aggressiveness(self) > Self::aggressiveness(other)
116    }
117
118    const fn aggressiveness(q: Self) -> u8 {
119        Self::type_aggressiveness(q.k) + Self::type_aggressiveness(q.v)
120    }
121
122    const fn type_aggressiveness(t: GgufKvCacheType) -> u8 {
123        match t {
124            GgufKvCacheType::F16 => 0,
125            GgufKvCacheType::Q8_0 => 1,
126            GgufKvCacheType::Q4_0 => 2,
127        }
128    }
129
130    pub fn from_llama_args(cache_type_k: &str, cache_type_v: &str) -> Option<Self> {
131        Some(Self {
132            k: GgufKvCacheType::from_llama_arg(cache_type_k)?,
133            v: GgufKvCacheType::from_llama_arg(cache_type_v)?,
134        })
135    }
136
137    pub fn k_cache_bytes_per_token(self, meta: &GgufCompactMeta) -> Option<u64> {
138        cache_bytes_per_token(meta, meta.key_length, self.k)
139    }
140
141    pub fn v_cache_bytes_per_token(self, meta: &GgufCompactMeta) -> Option<u64> {
142        cache_bytes_per_token(meta, meta.kv_cache_value_length(), self.v)
143    }
144
145    pub fn kv_cache_bytes_per_token(self, meta: &GgufCompactMeta) -> Option<u64> {
146        self.k_cache_bytes_per_token(meta)?
147            .checked_add(self.v_cache_bytes_per_token(meta)?)
148    }
149}
150
151fn cache_bytes_per_token(
152    meta: &GgufCompactMeta,
153    vector_length: u32,
154    cache_type: GgufKvCacheType,
155) -> Option<u64> {
156    let kv_heads = u64::from(meta.kv_cache_head_count()?);
157    let vector_length = u64::from((vector_length > 0).then_some(vector_length)?);
158    let layers = u64::from((meta.layer_count > 0).then_some(meta.layer_count)?);
159    let elements_per_layer = kv_heads.checked_mul(vector_length)?;
160    cache_type
161        .bytes_for_elements(elements_per_layer)?
162        .checked_mul(layers)
163}
164
165#[cfg(test)]
166mod tests {
167    use super::*;
168
169    #[test]
170    fn prices_key_and_value_types_independently() {
171        let meta = GgufCompactMeta {
172            head_count: 32,
173            kv_head_count: 8,
174            layer_count: 24,
175            key_length: 128,
176            value_length: 128,
177            ..Default::default()
178        };
179        let quant = GgufKvCacheQuant::new(GgufKvCacheType::Q8_0, GgufKvCacheType::Q4_0);
180
181        assert_eq!(quant.k_cache_bytes_per_token(&meta), Some(26_112));
182        assert_eq!(quant.v_cache_bytes_per_token(&meta), Some(13_824));
183        assert_eq!(quant.kv_cache_bytes_per_token(&meta), Some(39_936));
184    }
185
186    #[test]
187    fn prices_key_and_value_widths_independently() {
188        let meta = GgufCompactMeta {
189            head_count: 32,
190            kv_head_count: 8,
191            layer_count: 24,
192            key_length: 64,
193            value_length: 256,
194            ..Default::default()
195        };
196        let quant = GgufKvCacheQuant::new(GgufKvCacheType::Q8_0, GgufKvCacheType::Q4_0);
197
198        assert_eq!(quant.k_cache_bytes_per_token(&meta), Some(13_056));
199        assert_eq!(quant.v_cache_bytes_per_token(&meta), Some(27_648));
200        assert_eq!(quant.kv_cache_bytes_per_token(&meta), Some(40_704));
201    }
202
203    #[test]
204    fn prices_glm_dsa_absorbed_mla_shape() {
205        let meta = GgufCompactMeta {
206            architecture: "glm-dsa".to_string(),
207            head_count: 64,
208            kv_head_count: 64,
209            layer_count: 79,
210            key_length: 576,
211            value_length: 256,
212            kv_lora_rank: 512,
213            ..Default::default()
214        };
215
216        assert_eq!(
217            GgufKvCacheQuant::Q4_0.k_cache_bytes_per_token(&meta),
218            Some(25_596)
219        );
220        assert_eq!(
221            GgufKvCacheQuant::Q4_0.v_cache_bytes_per_token(&meta),
222            Some(22_752)
223        );
224        assert_eq!(
225            GgufKvCacheQuant::Q4_0.kv_cache_bytes_per_token(&meta),
226            Some(48_348)
227        );
228    }
229
230    #[test]
231    fn returns_none_when_required_fields_are_missing() {
232        let meta = GgufCompactMeta {
233            head_count: 32,
234            layer_count: 24,
235            key_length: 128,
236            ..Default::default()
237        };
238
239        assert_eq!(meta.k_cache_bytes_per_token_f16(), Some(196_608));
240        assert_eq!(meta.v_cache_bytes_per_token_f16(), None);
241        assert_eq!(
242            GgufKvCacheQuant::f16().kv_cache_bytes_per_token(&meta),
243            None
244        );
245    }
246}