1use 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
14pub const FSQ_DIM: usize = 36;
16
17pub const FSQ_LEVELS: usize = 21;
19
20pub struct Fsq;
25
26impl Fsq {
27 pub fn quantize<B: Backend, const D: usize>(x: Tensor<B, D>) -> Tensor<B, D> {
35 let clamped = x.clamp(-1.0, 1.0);
37 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 pub fn dequantize<B: Backend, const D: usize>(indices: Tensor<B, D>) -> Tensor<B, D> {
51 let half_levels = (FSQ_LEVELS - 1) as f32;
53 indices * (2.0 / half_levels) - 1.0
54 }
55
56 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
66pub const VQ_CODEBOOK_SIZE: usize = 8192;
68
69pub const VQ_EMBED_DIM: usize = 256;
71
72#[derive(burn::module::Module, Debug)]
81pub struct VqCodebook<B: Backend> {
82 embedding_sum: Param<Tensor<B, 2>>,
84 cluster_usage: Param<Tensor<B, 1>>,
86 #[module(skip)]
89 cpu_normalized: Vec<f32>,
90 #[module(skip)]
92 embed_dim: usize,
93}
94
95impl<B: Backend> VqCodebook<B> {
96 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 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 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 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 pub fn dequantize(&self, indices: &[usize]) -> Tensor<B, 2> {
176 let device = self.embedding_sum.device();
177 let n = indices.len();
178
179 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 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 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 assert!((vals[10] - 0.0).abs() < 1e-6, "Middle level: {}", vals[10]);
227
228 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 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 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 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 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 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 fn make_test_codebook() -> VqCodebook<TestBackend> {
348 let device = Default::default();
349 let n = 16; let dim = 4; 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; }
361 }
362 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 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 assert!((vals[0] - 1.0).abs() < 1e-6);
402 assert!((vals[4] - 4.0).abs() < 1e-6);
404 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 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}