1use crate::model_config;
6use super::traits::*;
7use anyhow::Result;
8use serde::{Serialize, Deserialize};
9
10model_config!(EncodecConfig {
11 vocab_size: usize = 1024, hidden_size: usize = 128,
13 num_hidden_layers: usize = 4, 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>, 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 Ok(model)
132 }
133
134 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
135 match inputs {
136 ModelInputs::Audio { input_features, .. } => {
137 let encoded = self.encoder.forward(input_features)?;
139
140 let (quantized, codes) = self.quantizer.forward(&encoded)?;
142
143 let decoded = self.decoder.forward(&quantized)?;
145
146 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 let conv_in = EncodecConv1d::new(config.audio_channels, channels, 7, 1, 3, device)?;
179
180 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 let conv_out = EncodecConv1d::new(channels, config.hidden_size, 7, 1, 3, device)?;
192
193 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
217fn 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 let conv_in = EncodecConv1d::new(config.hidden_size, channels, 7, 1, 3, device)?;
230
231 let lstm = Some(EncodecLSTM::new(channels, 2, device)?);
233
234 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 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 for _ in 0..num_residual {
280 residual_layers.push(EncodecResidualUnit::new(in_channels, device)?);
281 }
282
283 conv_layers.push(EncodecConv1d::new(in_channels, in_channels, 3, 1, 1, device)?);
285
286 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 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 conv_layers.push(EncodecConv1d::new(out_channels, out_channels, 3, 1, 1, device)?);
317
318 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 let input_candle = input.to_candle()?;
386 let weight_candle = self.weight.to_candle()?;
387
388 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 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 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 let (code_indices, quantized_part) = self.quantize_residual(&residual, codebook)?;
453
454 quantized = ops_fn::add(&quantized, &quantized_part)?;
456
457 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 let input_flat = input_candle.transpose(1, 2)?.reshape(&[batch * seq_len, channels])?;
475
476 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 let indices = distances.argmin(1)?;
489
490 let quantized_flat = codebook_candle.index_select(&indices, 0)?;
492
493 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}