Skip to main content

xz_embed/quantize/
scalar.rs

1use crate::quantize::VectorQuantizer;
2
3/// 标量量化(Scalar Quantization)
4///
5/// 将 float32 压缩到 u8(每维度 1 byte,压缩比 4:1)
6#[derive(Debug)]
7pub struct ScalarQuantizer {
8    /// 每个维度的 min/max
9    ranges: Vec<(f32, f32)>,
10    /// 量化位数 (8 = u8)
11    bits: usize,
12}
13
14impl ScalarQuantizer {
15    pub fn new(ranges: Vec<(f32, f32)>, bits: usize) -> Self {
16        Self { ranges, bits }
17    }
18
19    /// 从样本计算各维度的 min/max
20    pub fn from_samples(samples: &[Vec<f32>], bits: usize) -> Self {
21        if samples.is_empty() {
22            return Self { ranges: vec![], bits };
23        }
24
25        let dim = samples[0].len();
26        let mut ranges = Vec::with_capacity(dim);
27
28        for d in 0..dim {
29            let mut min = f32::MAX;
30            let mut max = f32::MIN;
31            for sample in samples {
32                let val = sample[d];
33                if val < min {
34                    min = val;
35                }
36                if val > max {
37                    max = val;
38                }
39            }
40            ranges.push((min, max));
41        }
42
43        Self { ranges, bits }
44    }
45
46    fn quantize_value(&self, value: f32, min: f32, max: f32) -> u8 {
47        let range = max - min;
48        if range < f32::EPSILON {
49            return 0;
50        }
51        let normalized = (value - min) / range;
52        let max_val = (1u32 << self.bits) - 1;
53        (normalized * max_val as f32).round().clamp(0.0, max_val as f32) as u8
54    }
55
56    fn dequantize_value(&self, q: u8, min: f32, max: f32) -> f32 {
57        let max_val = (1u32 << self.bits) - 1;
58        let normalized = q as f32 / max_val as f32;
59        min + normalized * (max - min)
60    }
61}
62
63impl VectorQuantizer for ScalarQuantizer {
64    fn compress(&self, vectors: &[Vec<f32>]) -> Vec<Vec<u8>> {
65        if self.ranges.is_empty() {
66            return vec![];
67        }
68
69        vectors
70            .iter()
71            .map(|v| {
72                v.iter()
73                    .enumerate()
74                    .map(|(d, &val)| {
75                        let (min, max) = self.ranges[d];
76                        self.quantize_value(val, min, max)
77                    })
78                    .collect()
79            })
80            .collect()
81    }
82
83    fn decompress(&self, quantized: &[Vec<u8>]) -> Vec<Vec<f32>> {
84        if self.ranges.is_empty() {
85            return vec![];
86        }
87
88        quantized
89            .iter()
90            .map(|q| {
91                q.iter()
92                    .enumerate()
93                    .map(|(d, &qval)| {
94                        let (min, max) = self.ranges[d];
95                        self.dequantize_value(qval, min, max)
96                    })
97                    .collect()
98            })
99            .collect()
100    }
101}