1use std::collections::HashMap;
23use std::hash::BuildHasher;
24
25use crate::cache_config::KvCacheType;
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub struct KvElemsPerToken {
35 pub k: u64,
36 pub v: u64,
37}
38
39fn lookup<S: BuildHasher>(
42 metadata: &HashMap<String, String, S>,
43 arch: &str,
44 suffix: &str,
45) -> Option<u64> {
46 metadata
47 .get(&format!("{arch}.{suffix}"))
48 .or_else(|| metadata.get(suffix))
49 .and_then(|v| v.trim().parse::<u64>().ok())
50}
51
52#[must_use]
79pub fn estimate_kv_elems_per_token<S: BuildHasher>(
80 metadata: &HashMap<String, String, S>,
81 architecture: Option<&str>,
82) -> Option<KvElemsPerToken> {
83 let arch = architecture
84 .map(str::to_owned)
85 .or_else(|| metadata.get("general.architecture").cloned())?;
86 let arch = arch.trim().to_ascii_lowercase();
87
88 let block_count = lookup(metadata, &arch, "block_count")?;
89 let head_count = lookup(metadata, &arch, "attention.head_count");
90 let head_count_kv = lookup(metadata, &arch, "attention.head_count_kv").or(head_count)?;
92
93 let derived_head_dim = || {
96 let embedding_length = lookup(metadata, &arch, "embedding_length")?;
97 let heads = head_count?;
98 (heads > 0).then(|| embedding_length / heads)
99 };
100 let key_length = lookup(metadata, &arch, "attention.key_length").or_else(derived_head_dim)?;
101 let value_length =
102 lookup(metadata, &arch, "attention.value_length").or_else(derived_head_dim)?;
103
104 if block_count == 0 || head_count_kv == 0 {
105 return None;
106 }
107
108 let per_head = block_count.saturating_mul(head_count_kv);
109 Some(KvElemsPerToken {
110 k: per_head.saturating_mul(key_length),
111 v: per_head.saturating_mul(value_length),
112 })
113}
114
115#[must_use]
117pub const fn kv_bytes_per_token(elems: KvElemsPerToken, k: KvCacheType, v: KvCacheType) -> u64 {
118 k.bytes_for_elems(elems.k) + v.bytes_for_elems(elems.v)
119}
120
121#[must_use]
127pub const fn estimate_kv_bytes_for_context(kv_bytes_per_token: u64, context_size: u64) -> u64 {
128 kv_bytes_per_token.saturating_mul(context_size)
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134
135 fn qwen_metadata() -> HashMap<String, String> {
137 HashMap::from([
138 ("general.architecture".to_string(), "qwen3".to_string()),
139 ("qwen3.block_count".to_string(), "64".to_string()),
140 ("qwen3.attention.head_count".to_string(), "40".to_string()),
141 ("qwen3.attention.head_count_kv".to_string(), "8".to_string()),
142 ("qwen3.embedding_length".to_string(), "5120".to_string()),
143 ("qwen3.attention.key_length".to_string(), "128".to_string()),
144 (
145 "qwen3.attention.value_length".to_string(),
146 "128".to_string(),
147 ),
148 ])
149 }
150
151 const QWEN_ELEMS: u64 = 64 * 8 * 128;
153
154 #[test]
155 fn computes_from_explicit_head_dims() {
156 let got = estimate_kv_elems_per_token(&qwen_metadata(), Some("qwen3"));
157 assert_eq!(
158 got,
159 Some(KvElemsPerToken {
160 k: QWEN_ELEMS,
161 v: QWEN_ELEMS
162 })
163 );
164 }
165
166 #[test]
167 fn architecture_falls_back_to_general_architecture_key() {
168 let got = estimate_kv_elems_per_token(&qwen_metadata(), None);
170 assert_eq!(
171 got,
172 Some(KvElemsPerToken {
173 k: QWEN_ELEMS,
174 v: QWEN_ELEMS
175 })
176 );
177 }
178
179 #[test]
180 fn architecture_lookup_is_case_insensitive() {
181 let got = estimate_kv_elems_per_token(&qwen_metadata(), Some("QWEN3"));
182 assert_eq!(
183 got,
184 Some(KvElemsPerToken {
185 k: QWEN_ELEMS,
186 v: QWEN_ELEMS
187 })
188 );
189 }
190
191 #[test]
192 fn derives_head_dim_from_embedding_length_when_absent() {
193 let mut md = qwen_metadata();
194 md.remove("qwen3.attention.key_length");
195 md.remove("qwen3.attention.value_length");
196 let got = estimate_kv_elems_per_token(&md, Some("qwen3"));
198 assert_eq!(
199 got,
200 Some(KvElemsPerToken {
201 k: QWEN_ELEMS,
202 v: QWEN_ELEMS
203 })
204 );
205 }
206
207 #[test]
210 fn falls_back_to_head_count_without_gqa() {
211 let mut md = qwen_metadata();
212 md.remove("qwen3.attention.head_count_kv");
213 let got = estimate_kv_elems_per_token(&md, Some("qwen3"));
214 let expected = 64 * 40 * 128;
215 assert_eq!(
216 got,
217 Some(KvElemsPerToken {
218 k: expected,
219 v: expected
220 })
221 );
222 }
223
224 #[test]
225 fn none_when_block_count_missing() {
226 let mut md = qwen_metadata();
227 md.remove("qwen3.block_count");
228 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
229 }
230
231 #[test]
232 fn none_when_head_counts_missing() {
233 let mut md = qwen_metadata();
234 md.remove("qwen3.attention.head_count");
235 md.remove("qwen3.attention.head_count_kv");
236 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
237 }
238
239 #[test]
241 fn none_when_head_dim_underivable() {
242 let mut md = qwen_metadata();
243 md.remove("qwen3.attention.key_length");
244 md.remove("qwen3.attention.value_length");
245 md.remove("qwen3.embedding_length");
246 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
247 }
248
249 #[test]
250 fn none_on_non_numeric_values() {
251 let mut md = qwen_metadata();
252 md.insert("qwen3.block_count".to_string(), "sixty-four".to_string());
253 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
254 }
255
256 #[test]
257 fn none_when_metadata_empty() {
258 assert_eq!(
259 estimate_kv_elems_per_token(&HashMap::new(), Some("llama")),
260 None
261 );
262 assert_eq!(estimate_kv_elems_per_token(&HashMap::new(), None), None);
263 }
264
265 #[test]
268 fn none_on_degenerate_zero_counts() {
269 let mut md = qwen_metadata();
270 md.insert("qwen3.block_count".to_string(), "0".to_string());
271 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
272
273 let mut md = qwen_metadata();
274 md.insert("qwen3.attention.head_count_kv".to_string(), "0".to_string());
275 assert_eq!(estimate_kv_elems_per_token(&md, Some("qwen3")), None);
276 }
277
278 #[test]
279 fn unprefixed_keys_are_accepted_as_a_fallback() {
280 let md = HashMap::from([
281 ("block_count".to_string(), "32".to_string()),
282 ("attention.head_count".to_string(), "32".to_string()),
283 ("attention.head_count_kv".to_string(), "8".to_string()),
284 ("embedding_length".to_string(), "4096".to_string()),
285 ]);
286 let expected = 32 * 8 * 128;
288 assert_eq!(
289 estimate_kv_elems_per_token(&md, Some("llama")),
290 Some(KvElemsPerToken {
291 k: expected,
292 v: expected
293 })
294 );
295 }
296
297 #[test]
300 fn kv_bytes_per_token_at_f16_matches_the_old_formula() {
301 let elems = KvElemsPerToken {
304 k: QWEN_ELEMS,
305 v: QWEN_ELEMS,
306 };
307 let got = kv_bytes_per_token(elems, KvCacheType::F16, KvCacheType::F16);
308 assert_eq!(got, (QWEN_ELEMS + QWEN_ELEMS) * 2);
309 }
310
311 #[test]
312 fn kv_bytes_per_token_at_q8_0_is_smaller_than_f16() {
313 let elems = KvElemsPerToken {
314 k: QWEN_ELEMS,
315 v: QWEN_ELEMS,
316 };
317 let f16 = kv_bytes_per_token(elems, KvCacheType::F16, KvCacheType::F16);
318 let q8_0 = kv_bytes_per_token(elems, KvCacheType::Q8_0, KvCacheType::Q8_0);
319 assert!(q8_0 < f16);
320 }
321
322 #[test]
323 fn kv_bytes_per_token_supports_asymmetric_k_v_types() {
324 let elems = KvElemsPerToken {
327 k: QWEN_ELEMS,
328 v: QWEN_ELEMS,
329 };
330 let mixed = kv_bytes_per_token(elems, KvCacheType::Q8_0, KvCacheType::F16);
331 let expected = KvCacheType::Q8_0.bytes_for_elems(QWEN_ELEMS)
332 + KvCacheType::F16.bytes_for_elems(QWEN_ELEMS);
333 assert_eq!(mixed, expected);
334 }
335
336 #[test]
337 fn context_multiplication_saturates() {
338 assert_eq!(estimate_kv_bytes_for_context(1024, 100), 102_400);
339 assert_eq!(estimate_kv_bytes_for_context(u64::MAX, u64::MAX), u64::MAX);
340 }
341}