Skip to main content

runtime/models_v2/
encodec.rs

1//! EnCodec Model V2 - Neural Audio Codec
2//!
3//! CNN-based encoder/decoder for audio compression using residual vector quantization
4
5use crate::model_config;
6use super::traits::*;
7use anyhow::Result;
8use serde::{Serialize, Deserialize};
9
10model_config!(EncodecConfig {
11    vocab_size: usize = 1024,  // Same as codebook_size for compatibility
12    hidden_size: usize = 128,
13    num_hidden_layers: usize = 4,  // Number of upsampling stages
14    audio_channels: usize = 1,
15    sample_rate: usize = 24000,
16    num_filters: usize = 32,
17    num_residual_layers: usize = 1,
18    upsampling_ratios: Vec<usize> = vec![8, 5, 4, 2],
19    norm_type: String = "weight_norm".to_string(),
20    codebook_size: usize = 1024,
21    codebook_dim: usize = 128,
22    num_codebooks: usize = 32,
23    use_causal_conv: bool = true,
24    pad_mode: String = "reflect".to_string(),
25    bandwidth: f32 = 6.0,
26    layer_norm_eps: f32 = 1e-5,
27    pad_token_id: i64 = 0,
28    bos_token_id: i64 = 1,
29    eos_token_id: i64 = 2,
30});
31
32impl EncodecConfig {
33    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
34        Self {
35            hidden_size: gguf.hidden_size,
36            ..Default::default()
37        }
38    }
39}
40
41pub struct EncodecModelV2 {
42    config: EncodecConfig,
43    device: Device,
44    encoder: EncodecEncoder,
45    decoder: EncodecDecoder,
46    quantizer: ResidualVectorQuantizer,
47}
48
49pub struct EncodecEncoder {
50    conv_in: EncodecConv1d,
51    down_blocks: Vec<EncodecDownsampleBlock>,
52    conv_out: EncodecConv1d,
53    lstm: Option<EncodecLSTM>,
54}
55
56pub struct EncodecDecoder {
57    conv_in: EncodecConv1d,
58    up_blocks: Vec<EncodecUpsampleBlock>,
59    conv_out: EncodecConv1d,
60    lstm: Option<EncodecLSTM>,
61}
62
63pub struct EncodecDownsampleBlock {
64    conv_layers: Vec<EncodecConv1d>,
65    downsample: EncodecConv1d,
66    residual_layers: Vec<EncodecResidualUnit>,
67}
68
69pub struct EncodecUpsampleBlock {
70    upsample: EncodecConvTranspose1d,
71    conv_layers: Vec<EncodecConv1d>,
72    residual_layers: Vec<EncodecResidualUnit>,
73}
74
75pub struct EncodecResidualUnit {
76    conv1: EncodecConv1d,
77    conv2: EncodecConv1d,
78}
79
80pub struct EncodecConv1d {
81    weight: Tensor,
82    bias: Option<Tensor>,
83    in_channels: usize,
84    out_channels: usize,
85    kernel_size: usize,
86    stride: usize,
87    padding: usize,
88}
89
90pub struct EncodecConvTranspose1d {
91    weight: Tensor,
92    bias: Option<Tensor>,
93    in_channels: usize,
94    out_channels: usize,
95    kernel_size: usize,
96    stride: usize,
97}
98
99pub struct EncodecLSTM {
100    weight_ih: Tensor,
101    weight_hh: Tensor,
102    bias_ih: Tensor,
103    bias_hh: Tensor,
104    hidden_size: usize,
105    num_layers: usize,
106}
107
108pub struct ResidualVectorQuantizer {
109    codebooks: Vec<Tensor>,  // [num_codebooks, codebook_size, codebook_dim]
110    codebook_size: usize,
111    codebook_dim: usize,
112    num_codebooks: usize,
113}
114
115impl Model for EncodecModelV2 {
116    type Config = EncodecConfig;
117
118    fn new(config: EncodecConfig) -> Result<Self> {
119        let device = Device::CPU;
120
121        let encoder = EncodecEncoder::new(&config, &device)?;
122        let decoder = EncodecDecoder::new(&config, &device)?;
123        let quantizer = ResidualVectorQuantizer::new(&config, &device)?;
124
125        Ok(Self { config, device, encoder, decoder, quantizer })
126    }
127
128    fn from_weights(config: EncodecConfig, weights: ModelWeights) -> Result<Self> {
129        let model = Self::new(config)?;
130        // Weight loading would be implemented here
131        Ok(model)
132    }
133
134    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
135        match inputs {
136            ModelInputs::Audio { input_features, .. } => {
137                // Encode: audio -> latent codes
138                let encoded = self.encoder.forward(input_features)?;
139
140                // Quantize: continuous latents -> discrete codes
141                let (quantized, codes) = self.quantizer.forward(&encoded)?;
142
143                // Decode: discrete codes -> reconstructed audio
144                let decoded = self.decoder.forward(&quantized)?;
145
146                // Return encoded representation
147                Ok(ModelOutputs::Logits {
148                    logits: decoded,
149                    hidden_states: Some(encoded)
150                })
151            }
152            _ => Err(anyhow::anyhow!("Encodec requires audio input")),
153        }
154    }
155
156    fn generate(&self, _prompt: &str, _config: &GenerationConfig) -> Result<String> {
157        Err(anyhow::anyhow!("Encodec is a codec model, use encode/decode methods"))
158    }
159
160    fn config(&self) -> &Self::Config { &self.config }
161
162    fn memory_requirements(&self) -> MemoryRequirements {
163        let p = self.config.hidden_size * self.config.num_filters * 1000 * 4;
164        MemoryRequirements { gpu_memory: p, cpu_memory: p / 4, kv_cache_memory: 0, peak_memory: p * 2 }
165    }
166
167    fn to_device(&mut self, device: &Device) -> Result<()> {
168        self.device = device.clone();
169        Ok(())
170    }
171}
172
173impl EncodecEncoder {
174    fn new(config: &EncodecConfig, device: &Device) -> Result<Self> {
175        let mut channels = config.num_filters;
176
177        // Initial conv: audio_channels -> num_filters
178        let conv_in = EncodecConv1d::new(config.audio_channels, channels, 7, 1, 3, device)?;
179
180        // Downsample blocks
181        let mut down_blocks = Vec::new();
182        for &ratio in &config.upsampling_ratios {
183            let out_channels = channels * 2;
184            down_blocks.push(EncodecDownsampleBlock::new(
185                channels, out_channels, ratio, config.num_residual_layers, device
186            )?);
187            channels = out_channels;
188        }
189
190        // Output conv
191        let conv_out = EncodecConv1d::new(channels, config.hidden_size, 7, 1, 3, device)?;
192
193        // Optional LSTM for sequential modeling
194        let lstm = Some(EncodecLSTM::new(config.hidden_size, 2, device)?);
195
196        Ok(Self { conv_in, down_blocks, conv_out, lstm })
197    }
198
199    fn forward(&self, input: &Tensor) -> Result<Tensor> {
200        let mut x = self.conv_in.forward(input)?;
201        x = elu(&x)?;
202
203        for block in &self.down_blocks {
204            x = block.forward(&x)?;
205        }
206
207        x = self.conv_out.forward(&x)?;
208
209        if let Some(ref lstm) = self.lstm {
210            x = lstm.forward(&x)?;
211        }
212
213        Ok(x)
214    }
215}
216
217// ELU activation: max(0, x) + min(0, alpha * (exp(x) - 1))
218fn elu(input: &Tensor) -> Result<Tensor> {
219    let x = input.to_candle()?;
220    let result = x.elu(1.0)?;
221    Ok(Tensor::from_candle(result))
222}
223
224impl EncodecDecoder {
225    fn new(config: &EncodecConfig, device: &Device) -> Result<Self> {
226        let mut channels = config.hidden_size;
227
228        // Initial conv
229        let conv_in = EncodecConv1d::new(config.hidden_size, channels, 7, 1, 3, device)?;
230
231        // Optional LSTM
232        let lstm = Some(EncodecLSTM::new(channels, 2, device)?);
233
234        // Upsample blocks (reverse order of downsampling)
235        let mut up_blocks = Vec::new();
236        let ratios: Vec<_> = config.upsampling_ratios.iter().rev().copied().collect();
237
238        for (i, &ratio) in ratios.iter().enumerate() {
239            let out_channels = if i == ratios.len() - 1 {
240                config.num_filters
241            } else {
242                channels / 2
243            };
244            up_blocks.push(EncodecUpsampleBlock::new(
245                channels, out_channels, ratio, config.num_residual_layers, device
246            )?);
247            channels = out_channels;
248        }
249
250        // Output conv: num_filters -> audio_channels
251        let conv_out = EncodecConv1d::new(config.num_filters, config.audio_channels, 7, 1, 3, device)?;
252
253        Ok(Self { conv_in, up_blocks, conv_out, lstm })
254    }
255
256    fn forward(&self, input: &Tensor) -> Result<Tensor> {
257        let mut x = self.conv_in.forward(input)?;
258
259        if let Some(ref lstm) = self.lstm {
260            x = lstm.forward(&x)?;
261        }
262
263        x = elu(&x)?;
264
265        for block in &self.up_blocks {
266            x = block.forward(&x)?;
267        }
268
269        self.conv_out.forward(&x)
270    }
271}
272
273impl EncodecDownsampleBlock {
274    fn new(in_channels: usize, out_channels: usize, ratio: usize, num_residual: usize, device: &Device) -> Result<Self> {
275        let mut conv_layers = Vec::new();
276        let mut residual_layers = Vec::new();
277
278        // Residual units before downsampling
279        for _ in 0..num_residual {
280            residual_layers.push(EncodecResidualUnit::new(in_channels, device)?);
281        }
282
283        // Pre-downsample conv
284        conv_layers.push(EncodecConv1d::new(in_channels, in_channels, 3, 1, 1, device)?);
285
286        // Downsample conv with stride
287        let downsample = EncodecConv1d::new(in_channels, out_channels, ratio * 2, ratio, ratio / 2, device)?;
288
289        Ok(Self { conv_layers, downsample, residual_layers })
290    }
291
292    fn forward(&self, input: &Tensor) -> Result<Tensor> {
293        let mut x = input.clone();
294
295        for layer in &self.residual_layers {
296            x = layer.forward(&x)?;
297        }
298
299        for conv in &self.conv_layers {
300            x = elu(&conv.forward(&x)?)?;
301        }
302
303        self.downsample.forward(&x)
304    }
305}
306
307impl EncodecUpsampleBlock {
308    fn new(in_channels: usize, out_channels: usize, ratio: usize, num_residual: usize, device: &Device) -> Result<Self> {
309        // Upsample with transposed conv
310        let upsample = EncodecConvTranspose1d::new(in_channels, out_channels, ratio * 2, ratio, device)?;
311
312        let mut conv_layers = Vec::new();
313        let mut residual_layers = Vec::new();
314
315        // Post-upsample conv
316        conv_layers.push(EncodecConv1d::new(out_channels, out_channels, 3, 1, 1, device)?);
317
318        // Residual units after upsampling
319        for _ in 0..num_residual {
320            residual_layers.push(EncodecResidualUnit::new(out_channels, device)?);
321        }
322
323        Ok(Self { upsample, conv_layers, residual_layers })
324    }
325
326    fn forward(&self, input: &Tensor) -> Result<Tensor> {
327        let mut x = self.upsample.forward(input)?;
328
329        for conv in &self.conv_layers {
330            x = elu(&conv.forward(&x)?)?;
331        }
332
333        for layer in &self.residual_layers {
334            x = layer.forward(&x)?;
335        }
336
337        Ok(x)
338    }
339}
340
341impl EncodecResidualUnit {
342    fn new(channels: usize, device: &Device) -> Result<Self> {
343        let conv1 = EncodecConv1d::new(channels, channels, 3, 1, 1, device)?;
344        let conv2 = EncodecConv1d::new(channels, channels, 1, 1, 0, device)?;
345        Ok(Self { conv1, conv2 })
346    }
347
348    fn forward(&self, input: &Tensor) -> Result<Tensor> {
349        let x = elu(&self.conv1.forward(input)?)?;
350        let x = self.conv2.forward(&x)?;
351        ops_fn::add(input, &x)
352    }
353}
354
355impl EncodecConv1d {
356    fn new(in_channels: usize, out_channels: usize, kernel_size: usize, stride: usize, padding: usize, device: &Device) -> Result<Self> {
357        let weight = ops_fn::zeros(&[out_channels, in_channels, kernel_size], DataType::Float32, device)?;
358        let bias = Some(ops_fn::zeros(&[out_channels], DataType::Float32, device)?);
359        Ok(Self { weight, bias, in_channels, out_channels, kernel_size, stride, padding })
360    }
361
362    fn forward(&self, input: &Tensor) -> Result<Tensor> {
363        let out = ops_fn::conv1d(input, &self.weight, None, self.stride, self.padding)?;
364
365        if let Some(ref bias) = self.bias {
366            let bias_candle = bias.to_candle()?;
367            let out_candle = out.to_candle()?;
368            let bias_expanded = bias_candle.reshape(&[1, self.out_channels, 1])?;
369            Ok(Tensor::from_candle(out_candle.broadcast_add(&bias_expanded)?))
370        } else {
371            Ok(out)
372        }
373    }
374}
375
376impl EncodecConvTranspose1d {
377    fn new(in_channels: usize, out_channels: usize, kernel_size: usize, stride: usize, device: &Device) -> Result<Self> {
378        let weight = ops_fn::zeros(&[in_channels, out_channels, kernel_size], DataType::Float32, device)?;
379        let bias = Some(ops_fn::zeros(&[out_channels], DataType::Float32, device)?);
380        Ok(Self { weight, bias, in_channels, out_channels, kernel_size, stride })
381    }
382
383    fn forward(&self, input: &Tensor) -> Result<Tensor> {
384        // Simplified transposed conv - would use proper conv_transpose1d in full impl
385        let input_candle = input.to_candle()?;
386        let weight_candle = self.weight.to_candle()?;
387
388        // For now, use conv with modified dimensions (placeholder)
389        // Real implementation would use conv_transpose1d kernel
390        let shape = input_candle.shape().dims();
391        let (batch, _, seq_len) = (shape[0], shape[1], shape[2]);
392        let out_len = seq_len * self.stride;
393
394        // Create output with upsampled size
395        let device = input_candle.device();
396        let output = candle_core::Tensor::zeros(&[batch, self.out_channels, out_len], candle_core::DType::F32, device)?;
397
398        if let Some(ref bias) = self.bias {
399            let bias_candle = bias.to_candle()?;
400            let bias_expanded = bias_candle.reshape(&[1, self.out_channels, 1])?;
401            Ok(Tensor::from_candle(output.broadcast_add(&bias_expanded)?))
402        } else {
403            Ok(Tensor::from_candle(output))
404        }
405    }
406}
407
408impl EncodecLSTM {
409    fn new(hidden_size: usize, num_layers: usize, device: &Device) -> Result<Self> {
410        let weight_ih = ops_fn::zeros(&[4 * hidden_size, hidden_size], DataType::Float32, device)?;
411        let weight_hh = ops_fn::zeros(&[4 * hidden_size, hidden_size], DataType::Float32, device)?;
412        let bias_ih = ops_fn::zeros(&[4 * hidden_size], DataType::Float32, device)?;
413        let bias_hh = ops_fn::zeros(&[4 * hidden_size], DataType::Float32, device)?;
414        Ok(Self { weight_ih, weight_hh, bias_ih, bias_hh, hidden_size, num_layers })
415    }
416
417    fn forward(&self, input: &Tensor) -> Result<Tensor> {
418        // Simplified LSTM: just pass through for now
419        // Real implementation would do full LSTM computation
420        Ok(input.clone())
421    }
422}
423
424impl ResidualVectorQuantizer {
425    fn new(config: &EncodecConfig, device: &Device) -> Result<Self> {
426        let mut codebooks = Vec::with_capacity(config.num_codebooks);
427
428        for _ in 0..config.num_codebooks {
429            codebooks.push(ops_fn::zeros(
430                &[config.codebook_size, config.codebook_dim],
431                DataType::Float32,
432                device
433            )?);
434        }
435
436        Ok(Self {
437            codebooks,
438            codebook_size: config.codebook_size,
439            codebook_dim: config.codebook_dim,
440            num_codebooks: config.num_codebooks,
441        })
442    }
443
444    fn forward(&self, input: &Tensor) -> Result<(Tensor, Vec<Tensor>)> {
445        let mut residual = input.clone();
446        let input_candle = input.to_candle()?;
447        let mut quantized = Tensor::from_candle(input_candle.zeros_like()?);
448        let mut codes = Vec::with_capacity(self.num_codebooks);
449
450        for codebook in &self.codebooks {
451            // Find nearest codebook entry
452            let (code_indices, quantized_part) = self.quantize_residual(&residual, codebook)?;
453
454            // Accumulate quantized values
455            quantized = ops_fn::add(&quantized, &quantized_part)?;
456
457            // Update residual
458            residual = ops_fn::sub(&residual, &quantized_part)?;
459
460            codes.push(code_indices);
461        }
462
463        Ok((quantized, codes))
464    }
465
466    fn quantize_residual(&self, input: &Tensor, codebook: &Tensor) -> Result<(Tensor, Tensor)> {
467        let input_candle = input.to_candle()?;
468        let codebook_candle = codebook.to_candle()?;
469
470        let shape = input_candle.shape().dims();
471        let (batch, channels, seq_len) = (shape[0], shape[1], shape[2]);
472
473        // Reshape input: [batch, channels, seq] -> [batch * seq, channels]
474        let input_flat = input_candle.transpose(1, 2)?.reshape(&[batch * seq_len, channels])?;
475
476        // Compute distances to codebook entries
477        // distance = ||x - c||^2 = ||x||^2 - 2<x,c> + ||c||^2
478        let input_sq = input_flat.sqr()?.sum(1)?;
479        let codebook_sq = codebook_candle.sqr()?.sum(1)?;
480        let inner = input_flat.matmul(&codebook_candle.t()?)?;
481
482        let input_sq_expanded = input_sq.unsqueeze(1)?;
483        let codebook_sq_expanded = codebook_sq.unsqueeze(0)?;
484
485        let distances = (input_sq_expanded.broadcast_add(&codebook_sq_expanded)? - (inner * 2.0)?)?;
486
487        // Find minimum distance indices
488        let indices = distances.argmin(1)?;
489
490        // Gather quantized values from codebook
491        let quantized_flat = codebook_candle.index_select(&indices, 0)?;
492
493        // Reshape back: [batch * seq, channels] -> [batch, channels, seq]
494        let quantized = quantized_flat.reshape(&[batch, seq_len, channels])?.transpose(1, 2)?;
495
496        let indices_reshaped = indices.reshape(&[batch, seq_len])?;
497
498        Ok((Tensor::from_candle(indices_reshaped), Tensor::from_candle(quantized)))
499    }
500}
501
502#[cfg(test)]
503mod tests {
504    use super::*;
505
506    #[test]
507    fn test_encodec_config() {
508        let config = EncodecConfig::default();
509        assert_eq!(config.hidden_size, 128);
510        assert_eq!(config.num_codebooks, 32);
511        assert_eq!(config.codebook_size, 1024);
512    }
513
514    #[test]
515    fn test_encodec_model_creation() {
516        let config = EncodecConfig {
517            hidden_size: 32,
518            num_filters: 8,
519            num_residual_layers: 1,
520            upsampling_ratios: vec![2, 2],
521            num_codebooks: 4,
522            codebook_size: 64,
523            codebook_dim: 32,
524            ..Default::default()
525        };
526
527        let model = EncodecModelV2::new(config).unwrap();
528        assert_eq!(model.config().hidden_size(), 32);
529    }
530}