Skip to main content

voxtral_micro/tts/codec/
quantizer.rs

1//! FSQ (Finite Scalar Quantization) and VQ (Vector Quantization) for the codec decoder.
2//!
3//! FSQ: Maps continuous 36-dim vectors to 21 discrete levels per dimension.
4//! VQ: Dequantizes semantic token indices via EMA codebook lookup.
5
6use anyhow::{Context, Result};
7use burn::module::{Param, ParamId};
8use burn::tensor::backend::Backend;
9use burn::tensor::Tensor;
10use safetensors::SafeTensors;
11
12use crate::models::weights::load_tensor;
13
14/// Number of FSQ dimensions (acoustic codebooks).
15pub const FSQ_DIM: usize = 36;
16
17/// Number of discrete levels per dimension.
18pub const FSQ_LEVELS: usize = 21;
19
20/// Finite Scalar Quantization.
21///
22/// Each of 36 dimensions is quantized to one of 21 uniformly-spaced levels
23/// in [-1, 1]. The level spacing is `2 / (FSQ_LEVELS - 1) = 0.1`.
24pub struct Fsq;
25
26impl Fsq {
27    /// Quantize continuous values to nearest FSQ level indices.
28    ///
29    /// # Arguments
30    /// * `x` - Continuous tensor with last dim = 36, values ideally in [-1, 1]
31    ///
32    /// # Returns
33    /// Integer indices in [0, 20] per dimension as f32 tensor (same shape as input).
34    pub fn quantize<B: Backend, const D: usize>(x: Tensor<B, D>) -> Tensor<B, D> {
35        // Clamp to [-1, 1], then map to [0, FSQ_LEVELS-1]
36        let clamped = x.clamp(-1.0, 1.0);
37        // Map [-1, 1] -> [0, 20]: idx = round((x + 1) / 2 * 20)
38        let half_levels = (FSQ_LEVELS - 1) as f32;
39        let indices = ((clamped + 1.0) * (half_levels / 2.0)).round();
40        indices.clamp(0.0, half_levels)
41    }
42
43    /// Dequantize FSQ level indices back to continuous values.
44    ///
45    /// # Arguments
46    /// * `indices` - Integer indices in [0, 20] (as f32 tensor)
47    ///
48    /// # Returns
49    /// Continuous values in [-1, 1] at level centers.
50    pub fn dequantize<B: Backend, const D: usize>(indices: Tensor<B, D>) -> Tensor<B, D> {
51        // Map [0, 20] -> [-1, 1]: x = indices / 10 - 1
52        let half_levels = (FSQ_LEVELS - 1) as f32;
53        indices * (2.0 / half_levels) - 1.0
54    }
55
56    /// Generate the 21 uniformly-spaced level values in [-1, 1].
57    pub fn levels<B: Backend>(device: &B::Device) -> Tensor<B, 1> {
58        let half_levels = (FSQ_LEVELS - 1) as f32;
59        let data: Vec<f32> = (0..FSQ_LEVELS)
60            .map(|i| i as f32 * 2.0 / half_levels - 1.0)
61            .collect();
62        Tensor::from_floats(data.as_slice(), device)
63    }
64}
65
66/// Number of semantic VQ codebook entries.
67pub const VQ_CODEBOOK_SIZE: usize = 8192;
68
69/// Embedding dimension for semantic VQ codebook.
70pub const VQ_EMBED_DIM: usize = 256;
71
72/// VQ Semantic Codebook for dequantizing semantic token indices.
73///
74/// Uses EMA (Exponential Moving Average) codebook: the embedding for each
75/// entry is `embedding_sum / cluster_usage`. This normalizes accumulated
76/// embeddings by how often each codebook entry was used during training.
77///
78/// Pre-normalizes embeddings to CPU cache at construction to avoid GPU
79/// readback during dequantize() (required for WASM compatibility).
80#[derive(burn::module::Module, Debug)]
81pub struct VqCodebook<B: Backend> {
82    /// Accumulated embeddings [8192, 256].
83    embedding_sum: Param<Tensor<B, 2>>,
84    /// Per-entry usage counts [8192].
85    cluster_usage: Param<Tensor<B, 1>>,
86    /// CPU-cached normalized embeddings (embedding_sum / cluster_usage).
87    /// Populated at construction to avoid GPU readback during inference.
88    #[module(skip)]
89    cpu_normalized: Vec<f32>,
90    /// Embedding dimension.
91    #[module(skip)]
92    embed_dim: usize,
93}
94
95impl<B: Backend> VqCodebook<B> {
96    /// Create VQ codebook from loaded tensors and pre-computed CPU cache.
97    ///
98    /// Use [`Self::precompute_normalized`] to build the CPU cache from raw
99    /// f32 slices (before uploading to GPU) to avoid GPU readback.
100    pub fn new(
101        embedding_sum: Tensor<B, 2>,
102        cluster_usage: Tensor<B, 1>,
103        cpu_normalized: Vec<f32>,
104    ) -> Self {
105        let embed_dim = embedding_sum.dims()[1];
106        Self {
107            embedding_sum: Param::initialized(ParamId::new(), embedding_sum),
108            cluster_usage: Param::initialized(ParamId::new(), cluster_usage),
109            cpu_normalized,
110            embed_dim,
111        }
112    }
113
114    /// Pre-compute normalized embeddings from raw f32 slices.
115    ///
116    /// Call this on CPU-side data BEFORE constructing tensors, to avoid
117    /// any GPU readback (required for WASM compatibility).
118    pub fn precompute_normalized(
119        embed_vals: &[f32],
120        usage_vals: &[f32],
121        n_entries: usize,
122        embed_dim: usize,
123    ) -> Vec<f32> {
124        let mut normalized = vec![0.0f32; n_entries * embed_dim];
125        for (idx, &usage) in usage_vals.iter().enumerate().take(n_entries) {
126            if usage > 0.0 {
127                let start = idx * embed_dim;
128                for j in 0..embed_dim {
129                    normalized[start + j] = embed_vals[start + j] / usage;
130                }
131            }
132        }
133        normalized
134    }
135
136    /// Load VQ codebook from SafeTensors.
137    ///
138    /// Expects:
139    /// - `audio_tokenizer.quantizer.semantic_codebook.embedding_sum` [8192, 256]
140    /// - `audio_tokenizer.quantizer.semantic_codebook.cluster_usage` [8192]
141    pub fn from_safetensors(safetensors: &SafeTensors, device: &B::Device) -> Result<Self> {
142        let embedding_sum: Tensor<B, 2> = load_tensor(
143            safetensors,
144            "audio_tokenizer.quantizer.semantic_codebook.embedding_sum",
145            device,
146        )
147        .context("Loading VQ embedding_sum")?;
148
149        let cluster_usage: Tensor<B, 1> = load_tensor(
150            safetensors,
151            "audio_tokenizer.quantizer.semantic_codebook.cluster_usage",
152            device,
153        )
154        .context("Loading VQ cluster_usage")?;
155
156        // Pre-compute normalized embeddings on CPU (sync readback OK on native)
157        let embed_data = embedding_sum.to_data();
158        let usage_data = cluster_usage.to_data();
159        let embed_vals = embed_data.as_slice::<f32>().unwrap();
160        let usage_vals = usage_data.as_slice::<f32>().unwrap();
161        let [n_entries, embed_dim] = embedding_sum.dims();
162        let cpu_normalized =
163            Self::precompute_normalized(embed_vals, usage_vals, n_entries, embed_dim);
164
165        Ok(Self::new(embedding_sum, cluster_usage, cpu_normalized))
166    }
167
168    /// Dequantize a batch of semantic token indices to embedding vectors.
169    ///
170    /// # Arguments
171    /// * `indices` - Semantic token indices, each in [0, 8191]. Shape: [N]
172    ///
173    /// # Returns
174    /// Embedding vectors [N, 256], normalized by cluster usage.
175    pub fn dequantize(&self, indices: &[usize]) -> Tensor<B, 2> {
176        let device = self.embedding_sum.device();
177        let n = indices.len();
178
179        // Use pre-normalized CPU cache (no GPU readback needed — WASM-safe)
180        let mut result = Vec::with_capacity(n * self.embed_dim);
181        for &idx in indices {
182            let start = idx * self.embed_dim;
183            let end = start + self.embed_dim;
184            result.extend_from_slice(&self.cpu_normalized[start..end]);
185        }
186
187        let data = burn::tensor::TensorData::new(result, [n, self.embed_dim]);
188        Tensor::from_data(data, &device)
189    }
190
191    /// Dequantize a single semantic token index.
192    ///
193    /// # Arguments
194    /// * `index` - Semantic token index in [0, 8191]
195    ///
196    /// # Returns
197    /// Embedding vector [1, 256], normalized by cluster usage.
198    pub fn dequantize_one(&self, index: usize) -> Tensor<B, 2> {
199        self.dequantize(&[index])
200    }
201}
202
203#[cfg(test)]
204mod tests {
205    use super::*;
206    use burn::backend::Wgpu;
207    use burn::tensor::TensorData;
208
209    type TestBackend = Wgpu;
210
211    #[test]
212    fn test_levels_values() {
213        let device = Default::default();
214        let levels = Fsq::levels::<TestBackend>(&device);
215
216        assert_eq!(levels.dims(), [FSQ_LEVELS]);
217
218        let data = levels.to_data();
219        let vals = data.as_slice::<f32>().unwrap();
220
221        // First level should be -1.0, last should be 1.0
222        assert!((vals[0] - (-1.0)).abs() < 1e-6, "First level: {}", vals[0]);
223        assert!((vals[20] - 1.0).abs() < 1e-6, "Last level: {}", vals[20]);
224
225        // Middle level should be 0.0
226        assert!((vals[10] - 0.0).abs() < 1e-6, "Middle level: {}", vals[10]);
227
228        // Spacing should be uniform (0.1)
229        for i in 1..FSQ_LEVELS {
230            let diff = vals[i] - vals[i - 1];
231            assert!(
232                (diff - 0.1).abs() < 1e-6,
233                "Non-uniform spacing at {}: {}",
234                i,
235                diff
236            );
237        }
238    }
239
240    #[test]
241    fn test_quantize_at_level_centers() {
242        let device = Default::default();
243
244        // Values exactly at level centers should quantize to their indices
245        let levels = Fsq::levels::<TestBackend>(&device);
246        let indices = Fsq::quantize(levels);
247
248        let data = indices.to_data();
249        let vals = data.as_slice::<f32>().unwrap();
250
251        for (i, &v) in vals.iter().enumerate() {
252            assert!(
253                (v - i as f32).abs() < 1e-5,
254                "Level {} quantized to {} (expected {})",
255                i,
256                v,
257                i
258            );
259        }
260    }
261
262    #[test]
263    fn test_roundtrip_preserves_level_values() {
264        let device = Default::default();
265
266        // Round-trip: levels -> quantize -> dequantize should recover original levels
267        let levels = Fsq::levels::<TestBackend>(&device);
268        let indices = Fsq::quantize(levels.clone());
269        let recovered = Fsq::dequantize(indices);
270
271        let orig_data = levels.to_data();
272        let recovered_data = recovered.to_data();
273        let orig = orig_data.as_slice::<f32>().unwrap();
274        let recov = recovered_data.as_slice::<f32>().unwrap();
275
276        for i in 0..FSQ_LEVELS {
277            assert!(
278                (orig[i] - recov[i]).abs() < 1e-6,
279                "Roundtrip mismatch at level {}: {} vs {}",
280                i,
281                orig[i],
282                recov[i]
283            );
284        }
285    }
286
287    #[test]
288    fn test_quantize_clamps_out_of_range() {
289        let device = Default::default();
290
291        // Values outside [-1, 1] should be clamped
292        let x = Tensor::<TestBackend, 1>::from_data(
293            TensorData::new(vec![-2.0f32, -1.5, 0.0, 1.5, 2.0], [5]),
294            &device,
295        );
296        let indices = Fsq::quantize(x);
297        let data = indices.to_data();
298        let vals = data.as_slice::<f32>().unwrap();
299
300        assert!((vals[0] - 0.0).abs() < 1e-5, "Clamped -2.0 -> idx 0");
301        assert!((vals[1] - 0.0).abs() < 1e-5, "Clamped -1.5 -> idx 0");
302        assert!((vals[2] - 10.0).abs() < 1e-5, "Center 0.0 -> idx 10");
303        assert!((vals[3] - 20.0).abs() < 1e-5, "Clamped 1.5 -> idx 20");
304        assert!((vals[4] - 20.0).abs() < 1e-5, "Clamped 2.0 -> idx 20");
305    }
306
307    #[test]
308    fn test_quantize_midpoint_snapping() {
309        let device = Default::default();
310
311        // Value between two levels should snap to nearest
312        // Levels at idx 10 = 0.0 and idx 11 = 0.1
313        // 0.04 is closer to 0.0 (idx 10), 0.06 is closer to 0.1 (idx 11)
314        let x =
315            Tensor::<TestBackend, 1>::from_data(TensorData::new(vec![0.04f32, 0.06], [2]), &device);
316        let indices = Fsq::quantize(x);
317        let data = indices.to_data();
318        let vals = data.as_slice::<f32>().unwrap();
319
320        assert!(
321            (vals[0] - 10.0).abs() < 1e-5,
322            "0.04 should snap to idx 10, got {}",
323            vals[0]
324        );
325        assert!(
326            (vals[1] - 11.0).abs() < 1e-5,
327            "0.06 should snap to idx 11, got {}",
328            vals[1]
329        );
330    }
331
332    #[test]
333    fn test_batch_quantize_shape() {
334        let device = Default::default();
335
336        // [batch, seq, 36] input should produce same shape output
337        let x = Tensor::<TestBackend, 3>::zeros([2, 5, FSQ_DIM], &device);
338        let indices = Fsq::quantize(x);
339        assert_eq!(indices.dims(), [2, 5, FSQ_DIM]);
340
341        let recovered = Fsq::dequantize(indices);
342        assert_eq!(recovered.dims(), [2, 5, FSQ_DIM]);
343    }
344
345    // --- VQ Codebook tests ---
346
347    fn make_test_codebook() -> VqCodebook<TestBackend> {
348        let device = Default::default();
349        let n = 16; // small codebook for testing
350        let dim = 4; // small embedding dim
351
352        // embedding_sum: each row i = [i+1, i+1, i+1, i+1] * usage
353        // cluster_usage: [2, 2, 2, ..., 2] (each used 2 times)
354        let mut embed_data = vec![0.0f32; n * dim];
355        let mut usage_data = vec![2.0f32; n];
356
357        for i in 0..n {
358            for d in 0..dim {
359                embed_data[i * dim + d] = (i + 1) as f32 * 2.0; // sum = val * usage
360            }
361        }
362        // Set entry 5 to zero usage
363        usage_data[5] = 0.0;
364
365        let cpu_norm =
366            VqCodebook::<TestBackend>::precompute_normalized(&embed_data, &usage_data, n, dim);
367        let embedding_sum =
368            Tensor::<TestBackend, 2>::from_data(TensorData::new(embed_data, [n, dim]), &device);
369        let cluster_usage =
370            Tensor::<TestBackend, 1>::from_data(TensorData::new(usage_data, [n]), &device);
371
372        VqCodebook::new(embedding_sum, cluster_usage, cpu_norm)
373    }
374
375    #[test]
376    fn test_vq_dequantize_single() {
377        let codebook = make_test_codebook();
378
379        // Index 0: embedding_sum = [2, 2, 2, 2], usage = 2 => result = [1, 1, 1, 1]
380        let result = codebook.dequantize_one(0);
381        assert_eq!(result.dims(), [1, 4]);
382
383        let data = result.to_data();
384        let vals = data.as_slice::<f32>().unwrap();
385        for &v in vals {
386            assert!((v - 1.0).abs() < 1e-6, "Expected 1.0, got {}", v);
387        }
388    }
389
390    #[test]
391    fn test_vq_dequantize_batch() {
392        let codebook = make_test_codebook();
393
394        let result = codebook.dequantize(&[0, 3, 7]);
395        assert_eq!(result.dims(), [3, 4]);
396
397        let data = result.to_data();
398        let vals = data.as_slice::<f32>().unwrap();
399
400        // Index 0: sum=2, usage=2 => 1.0
401        assert!((vals[0] - 1.0).abs() < 1e-6);
402        // Index 3: sum=8, usage=2 => 4.0
403        assert!((vals[4] - 4.0).abs() < 1e-6);
404        // Index 7: sum=16, usage=2 => 8.0
405        assert!((vals[8] - 8.0).abs() < 1e-6);
406    }
407
408    #[test]
409    fn test_vq_dequantize_zero_usage() {
410        let codebook = make_test_codebook();
411
412        // Index 5 has zero usage — should return zero vector
413        let result = codebook.dequantize_one(5);
414        let data = result.to_data();
415        let vals = data.as_slice::<f32>().unwrap();
416
417        for (i, &v) in vals.iter().enumerate() {
418            assert!(
419                v.abs() < 1e-7,
420                "Zero-usage entry should be zero, got val[{}] = {}",
421                i,
422                v
423            );
424        }
425    }
426
427    #[test]
428    fn test_vq_dequantize_output_shape() {
429        let codebook = make_test_codebook();
430
431        let result = codebook.dequantize(&[0, 1, 2, 3, 4]);
432        assert_eq!(result.dims(), [5, 4]);
433    }
434}