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 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 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 pub const F16: Self = Self {
88 k: GgufKvCacheType::F16,
89 v: GgufKvCacheType::F16,
90 };
91
92 pub const Q8_0: Self = Self {
94 k: GgufKvCacheType::Q8_0,
95 v: GgufKvCacheType::Q8_0,
96 };
97
98 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 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}