1use crate::{
15 codec::{FibCodeV1, FibQuantizer},
16 scoring::FibScorer,
17 FibQuantError, Result,
18};
19
20#[derive(Debug, Clone)]
22pub struct CompressedAttentionOutput {
23 pub logits: Vec<f32>,
25 pub probabilities: Vec<f32>,
27 pub output: Vec<f32>,
29 pub top_k_indices: Vec<usize>,
31 pub decompression_count: usize,
33}
34
35pub fn compressed_attention_logits(
55 query: &[f32],
56 compressed_keys: &[FibCodeV1],
57 scorer: &FibScorer,
58) -> Result<Vec<f32>> {
59 if query.is_empty() {
60 return Err(FibQuantError::ZeroDimension);
61 }
62 if compressed_keys.is_empty() {
63 return Ok(Vec::new());
64 }
65 check_finite(query)?;
66
67 let head_dim = query.len();
68 let scale = (head_dim as f64).sqrt() as f32;
69
70 let prepared = scorer.prepare_query(query)?;
72
73 let block_count = scorer.quantizer().profile().block_count() as usize;
75 let gram = scorer.gram_table();
76 let gram_size = gram.n();
77 let gram_values = gram.values();
78
79 let n_keys = compressed_keys.len();
80 let mut all_key_indices = Vec::with_capacity(n_keys * block_count);
81 let mut key_norms = Vec::with_capacity(n_keys);
82 for code in compressed_keys {
83 let stored_indices = crate::bitpack::unpack_indices(
84 &code.indices,
85 block_count,
86 scorer.quantizer().profile().wire_index_bits,
87 )?;
88 let stored_norm =
89 crate::scoring::decode_stored_norm(code, scorer.quantizer().profile())? as f32;
90 for &idx in &stored_indices {
91 all_key_indices.push(idx as u16);
92 }
93 key_norms.push(stored_norm);
94 }
95
96 let query_indices_u16: Vec<u16> = prepared.query_indices.iter().map(|&i| i as u16).collect();
98
99 let logits = crate::ffi::c_compressed_attention_logits(
100 &all_key_indices,
101 n_keys,
102 &key_norms,
103 gram_values,
104 gram_size,
105 &query_indices_u16,
106 block_count,
107 prepared.query_norm as f32,
108 scale,
109 );
110
111 check_finite(&logits)?;
113 Ok(logits)
114}
115
116pub fn compressed_attention_topk(
143 query: &[f32],
144 compressed_keys: &[FibCodeV1],
145 compressed_values: &[FibCodeV1],
146 scorer: &FibScorer,
147 quantizer: &FibQuantizer,
148 top_k: usize,
149) -> Result<CompressedAttentionOutput> {
150 if query.is_empty() {
151 return Err(FibQuantError::ZeroDimension);
152 }
153 if compressed_keys.is_empty() {
154 return Err(FibQuantError::CorruptPayload(
155 "compressed_attention_topk: empty keys".into(),
156 ));
157 }
158 if compressed_keys.len() != compressed_values.len() {
159 return Err(FibQuantError::CorruptPayload(format!(
160 "compressed_attention_topk: {} keys but {} values",
161 compressed_keys.len(),
162 compressed_values.len()
163 )));
164 }
165 check_finite(query)?;
166
167 let logits = compressed_attention_logits(query, compressed_keys, scorer)?;
169
170 let probabilities = softmax(&logits)?;
172
173 let n = compressed_keys.len();
175 let k = top_k.min(n).max(1);
176 let top_k_indices = topk_indices_by_probability(&probabilities, k);
177
178 let head_dim = query.len();
180 let mut output = vec![0.0f64; head_dim];
181 let mut decompression_count = 0usize;
182
183 for &idx in &top_k_indices {
184 let decoded = quantizer.decode(&compressed_values[idx])?;
185 decompression_count += 1;
186 let prob = f64::from(probabilities[idx]);
187 for (channel, acc) in decoded.iter().zip(output.iter_mut()) {
188 *acc += prob * f64::from(*channel);
189 }
190 }
191
192 let output: Vec<f32> = output.into_iter().map(|v| v as f32).collect();
193
194 Ok(CompressedAttentionOutput {
195 logits,
196 probabilities,
197 output,
198 top_k_indices,
199 decompression_count,
200 })
201}
202
203fn softmax(logits: &[f32]) -> Result<Vec<f32>> {
210 if logits.is_empty() {
211 return Err(FibQuantError::ZeroDimension);
212 }
213 check_finite(logits)?;
214 let mut logits_mut = logits.to_vec();
215 crate::ffi::c_softmax(&mut logits_mut).map_err(|_| {
216 FibQuantError::NumericalFailure("compressed attention softmax underflow".into())
217 })?;
218 Ok(logits_mut)
219}
220
221fn topk_indices_by_probability(probabilities: &[f32], k: usize) -> Vec<usize> {
224 let mut indexed: Vec<(usize, f32)> = probabilities.iter().copied().enumerate().collect();
225 indexed.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
227 indexed.into_iter().take(k).map(|(idx, _)| idx).collect()
228}
229
230fn check_finite(values: &[f32]) -> Result<()> {
232 if values.iter().any(|v| !v.is_finite()) {
233 return Err(FibQuantError::CorruptPayload(
234 "compressed attention input contains non-finite value".into(),
235 ));
236 }
237 Ok(())
238}
239
240#[cfg(test)]
241mod tests {
242 use super::*;
243 use crate::profile::FibQuantProfileV1;
244
245 fn build_test_quantizer() -> Result<FibQuantizer> {
247 let profile = FibQuantProfileV1::paper_default(8, 2, 32, 7)?;
248 FibQuantizer::new(profile)
249 }
250
251 fn mse(a: &[f32], b: &[f32]) -> f64 {
253 assert_eq!(a.len(), b.len(), "mse length mismatch");
254 if a.is_empty() {
255 return 0.0;
256 }
257 let sum: f64 = a
258 .iter()
259 .zip(b)
260 .map(|(x, y)| {
261 let d = f64::from(*x) - f64::from(*y);
262 d * d
263 })
264 .sum();
265 sum / a.len() as f64
266 }
267
268 #[test]
269 fn test_compressed_attention_vs_reference() -> Result<()> {
270 let quantizer = build_test_quantizer()?;
271 let scorer = FibScorer::new(quantizer.clone())?;
272 let head_dim = 8usize;
273
274 let query: Vec<f32> = vec![0.1, -0.2, 0.3, 0.4, -0.5, 0.6, -0.7, 0.8];
276
277 let raw_keys: Vec<Vec<f32>> = vec![
279 vec![0.8, -0.1, 0.2, 0.3, -0.4, 0.5, -0.6, 0.7],
280 vec![-0.3, 0.4, -0.5, 0.6, 0.7, -0.8, 0.1, -0.2],
281 vec![0.5, 0.5, -0.5, 0.1, 0.2, -0.3, 0.4, 0.5],
282 vec![-0.2, 0.3, 0.4, -0.6, 0.5, -0.1, 0.2, -0.7],
283 vec![0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],
284 vec![0.6, -0.5, 0.4, -0.3, 0.2, -0.1, 0.8, -0.6],
285 ];
286 let raw_values: Vec<Vec<f32>> = vec![
287 vec![0.2, 0.3, -0.1, 0.5, 0.4, -0.2, 0.6, 0.1],
288 vec![-0.5, 0.4, 0.3, -0.2, 0.6, 0.1, -0.3, 0.5],
289 vec![0.7, -0.3, 0.2, 0.4, -0.1, 0.5, 0.3, -0.4],
290 vec![0.1, -0.6, 0.3, 0.2, -0.4, 0.7, -0.1, 0.3],
291 vec![0.3, 0.3, 0.3, 0.3, 0.3, 0.3, 0.3, 0.3],
292 vec![-0.2, 0.5, -0.4, 0.6, -0.3, 0.2, 0.7, -0.5],
293 ];
294
295 let compressed_keys: Vec<FibCodeV1> = raw_keys
297 .iter()
298 .map(|k| quantizer.encode(k))
299 .collect::<Result<Vec<_>>>()?;
300 let compressed_values: Vec<FibCodeV1> = raw_values
301 .iter()
302 .map(|v| quantizer.encode(v))
303 .collect::<Result<Vec<_>>>()?;
304
305 let top_k = 4usize;
307 let compressed_out = compressed_attention_topk(
308 &query,
309 &compressed_keys,
310 &compressed_values,
311 &scorer,
312 &quantizer,
313 top_k,
314 )?;
315
316 assert_eq!(
318 compressed_out.decompression_count, top_k,
319 "decompression_count should be {}, got {}",
320 top_k, compressed_out.decompression_count
321 );
322 assert_eq!(compressed_out.top_k_indices.len(), top_k);
323 assert_eq!(compressed_out.output.len(), head_dim);
324 assert_eq!(compressed_out.logits.len(), raw_keys.len());
325 assert_eq!(compressed_out.probabilities.len(), raw_keys.len());
326
327 let prob_sum: f64 = compressed_out
329 .probabilities
330 .iter()
331 .map(|p| f64::from(*p))
332 .sum();
333 assert!(
334 (prob_sum - 1.0).abs() < 1e-5,
335 "probabilities should sum to 1.0, got {}",
336 prob_sum
337 );
338
339 let decoded_keys: Vec<Vec<f32>> = compressed_keys
341 .iter()
342 .map(|c| quantizer.decode(c))
343 .collect::<Result<Vec<_>>>()?;
344 let decoded_values: Vec<Vec<f32>> = compressed_values
345 .iter()
346 .map(|c| quantizer.decode(c))
347 .collect::<Result<Vec<_>>>()?;
348
349 let flat_decoded_keys: Vec<f32> = decoded_keys.iter().flatten().copied().collect();
350 let flat_decoded_values: Vec<f32> = decoded_values.iter().flatten().copied().collect();
351
352 let ref_logits = super::super::attention_ref::reference_attention_logits(
353 &query,
354 &flat_decoded_keys,
355 head_dim,
356 )?;
357 let ref_probs = softmax_local(&ref_logits)?;
358 let ref_output = super::super::attention_ref::reference_value_aggregation(
359 &ref_probs,
360 &flat_decoded_values,
361 head_dim,
362 )?;
363
364 let logit_mse = mse(&compressed_out.logits, &ref_logits);
366 assert!(logit_mse < 2.0, "logit MSE too large: {}", logit_mse);
367
368 let output_mse = mse(&compressed_out.output, &ref_output);
372 assert!(output_mse < 0.5, "output MSE too large: {}", output_mse);
373
374 let ref_topk = topk_indices_by_probability(&ref_probs, top_k);
376 let overlap = compressed_out
377 .top_k_indices
378 .iter()
379 .filter(|idx| ref_topk.contains(idx))
380 .count();
381 let agreement = overlap as f64 / top_k as f64;
382 assert!(
383 agreement >= 0.5,
384 "top-K agreement too low: {}/{} (compressed={:?}, ref={:?})",
385 overlap,
386 top_k,
387 compressed_out.top_k_indices,
388 ref_topk
389 );
390
391 Ok(())
392 }
393
394 #[test]
395 fn test_empty_keys_returns_empty_logits() -> Result<()> {
396 let quantizer = build_test_quantizer()?;
397 let scorer = FibScorer::new(quantizer.clone())?;
398 let query: Vec<f32> = vec![0.1, -0.2, 0.3, 0.4, -0.5, 0.6, -0.7, 0.8];
399
400 let logits = compressed_attention_logits(&query, &[], &scorer)?;
401 assert!(logits.is_empty());
402 Ok(())
403 }
404
405 #[test]
406 fn test_single_key_logit_finite() -> Result<()> {
407 let quantizer = build_test_quantizer()?;
408 let scorer = FibScorer::new(quantizer.clone())?;
409
410 let query: Vec<f32> = vec![0.1, -0.2, 0.3, 0.4, -0.5, 0.6, -0.7, 0.8];
411 let key: Vec<f32> = vec![0.5, 0.5, -0.5, 0.1, 0.2, -0.3, 0.4, 0.5];
412 let compressed_key = quantizer.encode(&key)?;
413
414 let logits = compressed_attention_logits(&query, &[compressed_key], &scorer)?;
415 assert_eq!(logits.len(), 1);
416 assert!(logits[0].is_finite());
417 Ok(())
418 }
419
420 #[test]
421 fn test_topk_exceeds_n_clamps() -> Result<()> {
422 let quantizer = build_test_quantizer()?;
423 let scorer = FibScorer::new(quantizer.clone())?;
424 let head_dim = 8usize;
425
426 let query: Vec<f32> = vec![0.1, -0.2, 0.3, 0.4, -0.5, 0.6, -0.7, 0.8];
427 let keys: Vec<Vec<f32>> = vec![
428 vec![0.8, -0.1, 0.2, 0.3, -0.4, 0.5, -0.6, 0.7],
429 vec![-0.3, 0.4, -0.5, 0.6, 0.7, -0.8, 0.1, -0.2],
430 vec![0.5, 0.5, -0.5, 0.1, 0.2, -0.3, 0.4, 0.5],
431 ];
432 let compressed_keys: Vec<FibCodeV1> = keys
433 .iter()
434 .map(|k| quantizer.encode(k))
435 .collect::<Result<Vec<_>>>()?;
436 let compressed_values: Vec<FibCodeV1> = compressed_keys.clone();
437
438 let out = compressed_attention_topk(
440 &query,
441 &compressed_keys,
442 &compressed_values,
443 &scorer,
444 &quantizer,
445 10,
446 )?;
447 assert_eq!(out.decompression_count, 3);
448 assert_eq!(out.top_k_indices.len(), 3);
449 assert_eq!(out.output.len(), head_dim);
450 Ok(())
451 }
452
453 fn softmax_local(logits: &[f32]) -> Result<Vec<f32>> {
455 use crate::FibQuantError;
456 if logits.is_empty() {
457 return Err(FibQuantError::ZeroDimension);
458 }
459 let max = logits
460 .iter()
461 .copied()
462 .fold(f32::NEG_INFINITY, |acc, v| acc.max(v));
463 let mut sum = 0.0f64;
464 let mut exps = Vec::with_capacity(logits.len());
465 for &v in logits {
466 let exp = f64::from(v - max).exp();
467 sum += exp;
468 exps.push(exp);
469 }
470 Ok(exps.into_iter().map(|e| (e / sum) as f32).collect())
471 }
472}