Skip to main content

fib_quant/kv/
compressed_attention.rs

1//! Compressed-domain attention: approximate logits on compressed keys,
2//! top-K value decode only.
3//!
4//! This module implements the core insight of compressed attention: you
5//! don't need to decompress every key vector to compute attention logits.
6//! The [`FibScorer`] can estimate `<query, key>` directly from the packed
7//! codeword indices via the Gram table, avoiding full decompression of the
8//! rotation-inverse + norm-scaling pipeline. Only the top-K value vectors
9//! (selected by approximate probability) need to be decompressed.
10//!
11//! This trades a small amount of logit accuracy for a large reduction in
12//! decode work: instead of `N` decompressions, only `top_k` are needed.
13
14use crate::{
15    codec::{FibCodeV1, FibQuantizer},
16    scoring::FibScorer,
17    FibQuantError, Result,
18};
19
20/// Output of compressed-domain attention with top-K value decode.
21#[derive(Debug, Clone)]
22pub struct CompressedAttentionOutput {
23    /// Approximate attention logits (pre-softmax), one per key.
24    pub logits: Vec<f32>,
25    /// Softmax probabilities derived from the approximate logits.
26    pub probabilities: Vec<f32>,
27    /// Weighted-aggregated output vector (length = head_dim).
28    pub output: Vec<f32>,
29    /// Indices of the top-K positions selected by probability (descending).
30    pub top_k_indices: Vec<usize>,
31    /// Number of value vectors actually decompressed (should be ≤ top_k).
32    pub decompression_count: usize,
33}
34
35/// Compute approximate attention logits on compressed keys WITHOUT full
36/// decompression.
37///
38/// Uses [`FibScorer::prepare_query`] + [`FibScorer::score_prepared`] for
39/// efficient batch scoring: the query is rotated and quantized once, then
40/// each compressed key is scored via Gram-table lookup only — no
41/// rotation-inverse or codeword reconstruction is needed.
42///
43/// The logits are scaled by `1/sqrt(head_dim)` as in standard scaled
44/// dot-product attention, where `head_dim = query.len()`.
45///
46/// # Arguments
47/// * `query` — Query vector, length = ambient_dim (= head_dim).
48/// * `compressed_keys` — Compressed key codes (`FibCodeV1`).
49/// * `scorer` — [`FibScorer`] wrapping the quantizer and Gram table.
50///
51/// # Errors
52/// Returns [`FibQuantError::ZeroDimension`] if the query is empty,
53/// [`FibQuantError::CorruptPayload`] if any input is non-finite.
54pub 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    // Prepare the query once for batch scoring (rotation + argmin).
71    let prepared = scorer.prepare_query(query)?;
72
73    // Gather all key indices and norms for the C kernel.
74    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    // Convert query indices to u16 for the C kernel.
97    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    // Validate logits are finite.
112    check_finite(&logits)?;
113    Ok(logits)
114}
115
116/// Compute compressed-domain attention with top-K value decode.
117///
118/// Pipeline:
119/// 1. Compute approximate logits on compressed keys (no decompression).
120/// 2. Softmax the logits to get attention probabilities.
121/// 3. Select top-K positions by probability (descending).
122/// 4. Decode ONLY the top-K value vectors via [`FibQuantizer::decode`].
123/// 5. Weighted-aggregate the top-K decoded values by their probabilities.
124///
125/// The `decompression_count` in the output will be `min(top_k, len)` — NOT
126/// the total number of values. This is the key efficiency win: with `N`
127/// stored positions and `top_k << N`, only `top_k` decode operations are
128/// performed instead of `N`.
129///
130/// # Arguments
131/// * `query` — Query vector, length = ambient_dim (= head_dim).
132/// * `compressed_keys` — Compressed key codes.
133/// * `compressed_values` — Compressed value codes (same length as keys).
134/// * `scorer` — [`FibScorer`] for approximate inner product scoring.
135/// * `quantizer` — [`FibQuantizer`] for decoding value vectors.
136/// * `top_k` — Number of top-probability positions to decompress and aggregate.
137///
138/// # Errors
139/// Returns [`FibQuantError::ZeroDimension`] if the query is empty,
140/// [`FibQuantError::CorruptPayload`] if keys/values length mismatch or
141/// any input is non-finite.
142pub 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    // 1. Compute approximate logits on compressed keys (no decompression).
168    let logits = compressed_attention_logits(query, compressed_keys, scorer)?;
169
170    // 2. Softmax → attention probabilities.
171    let probabilities = softmax(&logits)?;
172
173    // 3. Select top-K positions by descending probability.
174    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    // 4. Decode ONLY the top-K value vectors and weighted-aggregate.
179    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
203// ──────────────────────────────────────────────────────────────────────
204//  Internal helpers
205// ──────────────────────────────────────────────────────────────────────
206
207/// Numerically stable softmax with max-subtraction (f64 accumulator).
208/// Delegates the inner loop to the C kernel (`fq_softmax`).
209fn 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
221/// Select top-K indices by descending probability.
222/// Ties are broken by ascending index for determinism.
223fn topk_indices_by_probability(probabilities: &[f32], k: usize) -> Vec<usize> {
224    let mut indexed: Vec<(usize, f32)> = probabilities.iter().copied().enumerate().collect();
225    // Sort by descending probability, ties broken by ascending index.
226    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
230/// Check that all values are finite.
231fn 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    /// Build a test quantizer: ambient_dim=8, block_dim=2, codebook_size=32.
246    fn build_test_quantizer() -> Result<FibQuantizer> {
247        let profile = FibQuantProfileV1::paper_default(8, 2, 32, 7)?;
248        FibQuantizer::new(profile)
249    }
250
251    /// Simple MSE between two slices (for test assertions only).
252    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        // Synthetic query
275        let query: Vec<f32> = vec![0.1, -0.2, 0.3, 0.4, -0.5, 0.6, -0.7, 0.8];
276
277        // 6 synthetic key/value positions
278        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        // Encode keys and values
296        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        // --- Compressed attention (top-K = 4) ---
306        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        // Verify structural properties
317        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        // Probabilities should sum to ~1.0
328        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        // --- Reference: decode ALL keys and values, compute exact attention ---
340        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        // Logits should be in the same ballpark (approximate scoring).
365        let logit_mse = mse(&compressed_out.logits, &ref_logits);
366        assert!(logit_mse < 2.0, "logit MSE too large: {}", logit_mse);
367
368        // Output should be in the same ballpark.
369        // The compressed path uses top-K (4 of 6) with approximate probabilities,
370        // while the reference uses all 6 with exact probabilities on decoded keys.
371        let output_mse = mse(&compressed_out.output, &ref_output);
372        assert!(output_mse < 0.5, "output MSE too large: {}", output_mse);
373
374        // Top-K indices should have meaningful overlap with reference top-K.
375        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        // top_k=10 but only 3 keys — should clamp to 3
439        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    /// Local softmax for test comparisons (avoids importing private fns).
454    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}