1pub mod alibi;
19pub mod block;
20pub mod conv;
21pub mod layer_scale;
22pub mod qk_norm;
23pub mod quantizer;
24
25use anyhow::{Context, Result};
26use burn::tensor::backend::Backend;
27use burn::tensor::Tensor;
28use safetensors::SafeTensors;
29
30use crate::tts::config::CodecDecoderConfig;
31use block::CodecTransformerLayer;
32use conv::{CausalConv1d, CausalConvTranspose1d};
33use quantizer::{Fsq, VqCodebook};
34
35struct TransformerGroup<B: Backend> {
37 layers: Vec<CodecTransformerLayer<B>>,
38}
39
40impl<B: Backend> TransformerGroup<B> {
41 fn forward(&self, mut x: Tensor<B, 3>) -> Tensor<B, 3> {
42 for layer in &self.layers {
43 x = layer.forward(x);
44 }
45 x
46 }
47}
48
49pub struct CodecDecoder<B: Backend> {
57 input_conv: CausalConv1d<B>,
59 transformer_groups: Vec<TransformerGroup<B>>,
61 upsample_convs: Vec<CausalConvTranspose1d<B>>,
63 output_conv: CausalConv1d<B>,
65 vq_codebook: VqCodebook<B>,
67}
68
69impl<B: Backend> CodecDecoder<B> {
70 pub fn from_components(
74 input_conv: CausalConv1d<B>,
75 transformer_group_layers: Vec<Vec<CodecTransformerLayer<B>>>,
76 upsample_convs: Vec<CausalConvTranspose1d<B>>,
77 output_conv: CausalConv1d<B>,
78 vq_codebook: VqCodebook<B>,
79 ) -> Self {
80 let transformer_groups = transformer_group_layers
81 .into_iter()
82 .map(|layers| TransformerGroup { layers })
83 .collect();
84
85 Self {
86 input_conv,
87 transformer_groups,
88 upsample_convs,
89 output_conv,
90 vq_codebook,
91 }
92 }
93
94 pub fn from_safetensors(
101 safetensors: &SafeTensors,
102 config: &CodecDecoderConfig,
103 device: &B::Device,
104 ) -> Result<Self> {
105 let prefix = "audio_tokenizer";
106
107 let input_conv = CausalConv1d::from_safetensors(
109 safetensors,
110 &format!("{prefix}.decoder_blocks.0.conv"),
111 1, device,
113 )
114 .context("Loading input conv (block 0)")?;
115
116 let transformer_block_indices = [1, 3, 5, 7];
118 let mut transformer_groups = Vec::with_capacity(4);
119 for (group_idx, &block_idx) in transformer_block_indices.iter().enumerate() {
120 let sliding_window = config.sliding_windows[group_idx];
121 let mut layers = Vec::with_capacity(config.layers_per_block);
122 for layer_idx in 0..config.layers_per_block {
123 let layer_prefix =
124 format!("{prefix}.decoder_blocks.{block_idx}.layers.{layer_idx}");
125 let layer = CodecTransformerLayer::from_safetensors(
126 safetensors,
127 &layer_prefix,
128 config.n_heads,
129 config.head_dim,
130 sliding_window,
131 config.norm_eps,
132 device,
133 )
134 .with_context(|| {
135 format!("Loading transformer block {block_idx} layer {layer_idx}")
136 })?;
137 layers.push(layer);
138 }
139 transformer_groups.push(TransformerGroup { layers });
140 }
141
142 let upsample_block_indices = [2, 4, 6];
144 let mut upsample_convs = Vec::with_capacity(3);
145 for &block_idx in &upsample_block_indices {
146 let conv_t = CausalConvTranspose1d::from_safetensors(
147 safetensors,
148 &format!("{prefix}.decoder_blocks.{block_idx}.conv"),
149 2, device,
151 )
152 .with_context(|| format!("Loading upsample conv (block {block_idx})"))?;
153 upsample_convs.push(conv_t);
154 }
155
156 let output_conv = CausalConv1d::from_safetensors(
158 safetensors,
159 &format!("{prefix}.output_proj.conv"),
160 1, device,
162 )
163 .context("Loading output conv")?;
164
165 let vq_codebook =
167 VqCodebook::from_safetensors(safetensors, device).context("Loading VQ codebook")?;
168
169 Ok(Self {
170 input_conv,
171 transformer_groups,
172 upsample_convs,
173 output_conv,
174 vq_codebook,
175 })
176 }
177
178 pub fn decode(
187 &self,
188 semantic_indices: &[usize],
189 acoustic_indices: Tensor<B, 2>,
190 ) -> Tensor<B, 2> {
191 let n_frames = semantic_indices.len();
192
193 let semantic_embeds = self.vq_codebook.dequantize(semantic_indices);
196
197 let acoustic_values = Fsq::dequantize(acoustic_indices);
199
200 let features = Tensor::cat(vec![semantic_embeds, acoustic_values], 1);
202 let input_channels = features.dims()[1];
203
204 let x = features
206 .reshape([1, n_frames, input_channels])
207 .swap_dims(1, 2);
208
209 let x = self.input_conv.forward(x);
211
212 let mut x = x;
218 for (i, group) in self.transformer_groups.iter().enumerate() {
219 let x_seq = x.swap_dims(1, 2); let x_seq = group.forward(x_seq);
222 x = x_seq.swap_dims(1, 2); if i < self.upsample_convs.len() {
226 x = self.upsample_convs[i].forward(x);
227 }
228 }
229
230 let x = self.output_conv.forward(x);
232
233 let [_batch, patch_size, n_patches] = x.dims();
236 let x = x.swap_dims(1, 2); x.reshape([1, n_patches * patch_size])
238 }
239
240 pub fn vq_codebook(&self) -> &VqCodebook<B> {
242 &self.vq_codebook
243 }
244}
245
246#[cfg(test)]
247mod tests {
248 use super::*;
249 use burn::backend::Wgpu;
250 use burn::tensor::TensorData;
251
252 type TestBackend = Wgpu;
253
254 fn make_test_decoder(device: &<TestBackend as Backend>::Device) -> CodecDecoder<TestBackend> {
256 let input_channels = 6; let dim = 8; let n_heads = 2;
260 let head_dim = 4; let ffn_dim = 32;
262 let output_patch_size = 4; let g = Tensor::<TestBackend, 3>::ones([dim, 1, 1], device);
266 let v = Tensor::<TestBackend, 3>::ones([dim, input_channels, 3], device);
267 let input_conv = CausalConv1d::from_weight_norm(g, v, 1, device);
268
269 let windows = [2, 4, 8, 16];
271 let mut transformer_groups = Vec::new();
272 for &window in &windows {
273 let mut layers = Vec::new();
274 for _ in 0..2 {
275 use block::CodecAttention;
276 use burn::nn::LinearConfig;
277
278 let wq = LinearConfig::new(dim, dim).with_bias(false).init(device);
279 let wk = LinearConfig::new(dim, dim).with_bias(false).init(device);
280 let wv = LinearConfig::new(dim, dim).with_bias(false).init(device);
281 let wo = LinearConfig::new(dim, dim).with_bias(false).init(device);
282
283 let q_weight = Tensor::<TestBackend, 1>::ones([dim], device);
284 let k_weight = Tensor::<TestBackend, 1>::ones([dim], device);
285 let qk_norm = qk_norm::QkNorm::new(q_weight, k_weight, n_heads, head_dim);
286
287 let attention =
288 CodecAttention::new(wq, wk, wv, wo, qk_norm, n_heads, head_dim, window);
289
290 let attn_scale = layer_scale::LayerScale::new(Tensor::ones([dim], device) * 0.01);
291 let ffn_scale = layer_scale::LayerScale::new(Tensor::ones([dim], device) * 0.01);
292
293 use crate::models::layers::{RmsNorm, SwiGLUConfig};
294 use burn::module::{Param, ParamId};
295
296 let attention_norm = RmsNorm {
297 weight: burn::nn::RmsNorm {
298 gamma: Param::initialized(
299 ParamId::new(),
300 Tensor::<TestBackend, 1>::ones([dim], device),
301 ),
302 epsilon: 1e-5,
303 },
304 };
305 let ffn_norm = RmsNorm {
306 weight: burn::nn::RmsNorm {
307 gamma: Param::initialized(
308 ParamId::new(),
309 Tensor::<TestBackend, 1>::ones([dim], device),
310 ),
311 epsilon: 1e-5,
312 },
313 };
314 let ffn = SwiGLUConfig::new(dim, ffn_dim)
315 .with_bias(false)
316 .init(device);
317
318 let layer = CodecTransformerLayer::new(
319 attention_norm,
320 attention,
321 attn_scale,
322 ffn_norm,
323 ffn,
324 ffn_scale,
325 );
326 layers.push(layer);
327 }
328 transformer_groups.push(TransformerGroup { layers });
329 }
330
331 let mut upsample_convs = Vec::new();
333 for _ in 0..3 {
334 let g = Tensor::<TestBackend, 3>::ones([dim, 1, 1], device);
335 let v = Tensor::<TestBackend, 3>::ones([dim, dim, 4], device);
336 upsample_convs.push(CausalConvTranspose1d::from_weight_norm(g, v, 2, device));
337 }
338
339 let g = Tensor::<TestBackend, 3>::ones([output_patch_size, 1, 1], device);
341 let v = Tensor::<TestBackend, 3>::ones([output_patch_size, dim, 7], device);
342 let output_conv = CausalConv1d::from_weight_norm(g, v, 1, device);
343
344 let embed_sum = Tensor::<TestBackend, 2>::ones([16, 4], device);
346 let usage = Tensor::<TestBackend, 1>::ones([16], device);
347 let cpu_norm =
348 VqCodebook::<TestBackend>::precompute_normalized(&vec![1.0; 64], &vec![1.0; 16], 16, 4);
349 let vq_codebook = VqCodebook::new(embed_sum, usage, cpu_norm);
350
351 CodecDecoder {
352 input_conv,
353 transformer_groups,
354 upsample_convs,
355 output_conv,
356 vq_codebook,
357 }
358 }
359
360 #[test]
361 fn test_codec_decoder_output_shape() {
362 let device = Default::default();
363 let decoder = make_test_decoder(&device);
364
365 let n_frames = 4;
366 let semantic_indices = vec![0usize; n_frames];
367 let acoustic_indices = Tensor::<TestBackend, 2>::zeros([n_frames, 2], &device);
370
371 let output = decoder.decode(&semantic_indices, acoustic_indices);
372
373 assert_eq!(output.dims()[0], 1);
376 assert_eq!(output.dims()[1], 32 * 4);
377 }
378
379 #[test]
380 fn test_codec_decoder_single_frame() {
381 let device = Default::default();
382 let decoder = make_test_decoder(&device);
383
384 let semantic_indices = vec![0usize; 1];
385 let acoustic_indices = Tensor::<TestBackend, 2>::zeros([1, 2], &device);
386
387 let output = decoder.decode(&semantic_indices, acoustic_indices);
388
389 assert_eq!(output.dims(), [1, 32]);
391 }
392
393 #[test]
394 fn test_codec_decoder_output_not_all_zeros() {
395 let device = Default::default();
397 let decoder = make_test_decoder(&device);
398
399 let n_frames = 2;
400 let semantic_indices = vec![0usize; n_frames];
401 let acoustic_indices = Tensor::<TestBackend, 2>::ones([n_frames, 2], &device) * 10.0; let output = decoder.decode(&semantic_indices, acoustic_indices);
404
405 let data = output.to_data();
406 let vals = data.as_slice::<f32>().unwrap();
407
408 let has_nonzero = vals.iter().any(|&v| v.abs() > 1e-6);
410 assert!(has_nonzero, "Decoder output should not be all zeros");
411 }
412
413 #[test]
414 fn test_codec_decoder_upsampling_ratio() {
415 let device = Default::default();
417 let decoder = make_test_decoder(&device);
418
419 for n_frames in [2, 5, 10] {
420 let semantic_indices = vec![0usize; n_frames];
421 let acoustic_indices = Tensor::<TestBackend, 2>::zeros([n_frames, 2], &device);
422
423 let output = decoder.decode(&semantic_indices, acoustic_indices);
424 let total_samples = output.dims()[1];
425 let expected_patches = n_frames * 8; let expected_samples = expected_patches * 4; assert_eq!(
429 total_samples, expected_samples,
430 "For {} frames: expected {} samples, got {}",
431 n_frames, expected_samples, total_samples
432 );
433 }
434 }
435
436 #[test]
437 fn test_fsq_dequantize_integration() {
438 let device: <TestBackend as Backend>::Device = Default::default();
440
441 let indices = Tensor::<TestBackend, 2>::from_data(
443 TensorData::new(vec![0.0f32, 10.0, 20.0, 0.0, 10.0, 20.0], [2, 3]),
444 &device,
445 );
446 let values = Fsq::dequantize(indices);
447
448 let data = values.to_data();
449 let vals = data.as_slice::<f32>().unwrap();
450
451 assert!((vals[0] - (-1.0)).abs() < 1e-5);
453 assert!((vals[1] - 0.0).abs() < 1e-5);
454 assert!((vals[2] - 1.0).abs() < 1e-5);
455 }
456}