Skip to main content

candle_transformers/models/
encodec.rs

1//! EnCodec neural audio codec based on the Encodec implementation.
2//!
3//! See ["High Fidelity Neural Audio Compression"](https://arxiv.org/abs/2210.13438)
4//!
5//! Based on implementation from [huggingface/transformers](https://github.com/huggingface/transformers/blob/main/src/transformers/models/encodec/modeling_encodec.py)
6
7use candle::{DType, IndexOp, Layout, Module, Result, Shape, Tensor, D};
8use candle_nn::{conv1d, Conv1d, ConvTranspose1d, VarBuilder};
9
10// Encodec Model
11// https://github.com/huggingface/transformers/blob/main/src/transformers/models/encodec/modeling_encodec.py
12
13#[derive(Debug, Copy, Clone, PartialEq, Eq, serde::Deserialize)]
14pub enum NormType {
15    WeightNorm,
16    TimeGroupNorm,
17    None,
18}
19
20#[derive(Debug, Copy, Clone, PartialEq, Eq, serde::Deserialize)]
21pub enum PadMode {
22    Constant,
23    Reflect,
24    Replicate,
25}
26
27#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
28pub struct Config {
29    pub target_bandwidths: Vec<f64>,
30    pub sampling_rate: usize,
31    pub audio_channels: usize,
32    pub normalize: bool,
33    pub chunk_length_s: Option<usize>,
34    pub overlap: Option<usize>,
35    pub hidden_size: usize,
36    pub num_filters: usize,
37    pub num_residual_layers: usize,
38    pub upsampling_ratios: Vec<usize>,
39    pub norm_type: NormType,
40    pub kernel_size: usize,
41    pub last_kernel_size: usize,
42    pub residual_kernel_size: usize,
43    pub dilation_growth_rate: usize,
44    pub use_causal_conv: bool,
45    pub pad_mode: PadMode,
46    pub compress: usize,
47    pub num_lstm_layers: usize,
48    pub trim_right_ratio: f64,
49    pub codebook_size: usize,
50    pub codebook_dim: Option<usize>,
51    pub use_conv_shortcut: bool,
52}
53
54impl Default for Config {
55    fn default() -> Self {
56        Self {
57            target_bandwidths: vec![1.5, 3.0, 6.0, 12.0, 24.0],
58            sampling_rate: 24_000,
59            audio_channels: 1,
60            normalize: false,
61            chunk_length_s: None,
62            overlap: None,
63            hidden_size: 128,
64            num_filters: 32,
65            num_residual_layers: 1,
66            upsampling_ratios: vec![8, 5, 4, 2],
67            norm_type: NormType::WeightNorm,
68            kernel_size: 7,
69            last_kernel_size: 7,
70            residual_kernel_size: 3,
71            dilation_growth_rate: 2,
72            use_causal_conv: true,
73            // This should be PadMode::Reflect which is currently unsupported in candle.
74            pad_mode: PadMode::Replicate,
75            compress: 2,
76            num_lstm_layers: 2,
77            trim_right_ratio: 1.0,
78            codebook_size: 1024,
79            codebook_dim: None,
80            use_conv_shortcut: true,
81        }
82    }
83}
84
85impl Config {
86    fn codebook_dim(&self) -> usize {
87        self.codebook_dim.unwrap_or(self.hidden_size)
88    }
89
90    fn frame_rate(&self) -> usize {
91        let hop_length: usize = self.upsampling_ratios.iter().product();
92        self.sampling_rate.div_ceil(hop_length)
93    }
94
95    fn num_quantizers(&self) -> usize {
96        let num = 1000f64
97            * self
98                .target_bandwidths
99                .last()
100                .expect("empty target_bandwidths");
101        (num as usize) / (self.frame_rate() * 10)
102    }
103}
104
105fn get_extra_padding_for_conv1d(
106    xs: &Tensor,
107    k_size: usize,
108    stride: usize,
109    padding_total: usize,
110) -> Result<usize> {
111    let len = xs.dim(D::Minus1)?;
112    let n_frames = (len + padding_total).saturating_sub(k_size) as f64 / stride as f64 + 1.0;
113    let ideal_len =
114        ((n_frames.ceil() as usize - 1) * stride + k_size).saturating_sub(padding_total);
115    Ok(ideal_len.saturating_sub(len))
116}
117
118fn pad1d(xs: &Tensor, pad_l: usize, pad_r: usize, mode: PadMode) -> Result<Tensor> {
119    match mode {
120        PadMode::Constant => xs.pad_with_zeros(D::Minus1, pad_l, pad_r),
121        PadMode::Reflect => candle::bail!("pad-mode 'reflect' is not supported"),
122        PadMode::Replicate => xs.pad_with_same(D::Minus1, pad_l, pad_r),
123    }
124}
125
126// Applies weight norm for inference by recomputing the weight tensor. This
127// does not apply to training.
128// https://pytorch.org/docs/stable/generated/torch.nn.utils.weight_norm.html
129pub fn conv1d_weight_norm(
130    in_c: usize,
131    out_c: usize,
132    kernel_size: usize,
133    config: candle_nn::Conv1dConfig,
134    vb: VarBuilder,
135) -> Result<Conv1d> {
136    let weight_g = vb.get((out_c, 1, 1), "weight_g")?;
137    let weight_v = vb.get((out_c, in_c, kernel_size), "weight_v")?;
138    let norm_v = weight_v.sqr()?.sum_keepdim((1, 2))?.sqrt()?;
139    let weight = weight_v.broadcast_mul(&weight_g)?.broadcast_div(&norm_v)?;
140    let bias = vb.get(out_c, "bias")?;
141    Ok(Conv1d::new(weight, Some(bias), config))
142}
143
144pub fn conv1d_weight_norm_no_bias(
145    in_c: usize,
146    out_c: usize,
147    kernel_size: usize,
148    config: candle_nn::Conv1dConfig,
149    vb: VarBuilder,
150) -> Result<Conv1d> {
151    let weight_g = vb.get((out_c, 1, 1), "weight_g")?;
152    let weight_v = vb.get((out_c, in_c, kernel_size), "weight_v")?;
153    let norm_v = weight_v.sqr()?.sum_keepdim((1, 2))?.sqrt()?;
154    let weight = weight_v.broadcast_mul(&weight_g)?.broadcast_div(&norm_v)?;
155    Ok(Conv1d::new(weight, None, config))
156}
157
158pub fn conv_transpose1d_weight_norm(
159    in_c: usize,
160    out_c: usize,
161    kernel_size: usize,
162    bias: bool,
163    config: candle_nn::ConvTranspose1dConfig,
164    vb: VarBuilder,
165) -> Result<ConvTranspose1d> {
166    let weight_g = vb.get((in_c, 1, 1), "weight_g")?;
167    let weight_v = vb.get((in_c, out_c, kernel_size), "weight_v")?;
168    let norm_v = weight_v.sqr()?.sum_keepdim((1, 2))?.sqrt()?;
169    let weight = weight_v.broadcast_mul(&weight_g)?.broadcast_div(&norm_v)?;
170    let bias = if bias {
171        Some(vb.get(out_c, "bias")?)
172    } else {
173        None
174    };
175    Ok(ConvTranspose1d::new(weight, bias, config))
176}
177
178struct CodebookEncode;
179
180impl candle::CustomOp2 for CodebookEncode {
181    fn name(&self) -> &'static str {
182        "cb"
183    }
184
185    fn cpu_fwd(
186        &self,
187        lhs_storage: &candle::CpuStorage,
188        lhs_layout: &Layout,
189        rhs_storage: &candle::CpuStorage,
190        rhs_layout: &Layout,
191    ) -> Result<(candle::CpuStorage, Shape)> {
192        use rayon::prelude::*;
193
194        let (lhs_dim1, lhs_dim2) = lhs_layout.shape().dims2()?;
195        let (rhs_dim1, rhs_dim2) = rhs_layout.shape().dims2()?;
196        if lhs_dim2 != rhs_dim2 {
197            candle::bail!("CodebookEncode, mismatch on last dim, {lhs_layout:?} {rhs_layout:?}");
198        }
199        if lhs_dim2 == 0 {
200            candle::bail!("CodebookEncode, empty last dim {lhs_layout:?}")
201        }
202        let lhs = match lhs_layout.contiguous_offsets() {
203            None => candle::bail!("CodebookEncode, lhs has to be contiguous, got {lhs_layout:?}"),
204            Some((o1, o2)) => {
205                let slice = lhs_storage.as_slice::<f32>()?;
206                &slice[o1..o2]
207            }
208        };
209        let rhs = match rhs_layout.contiguous_offsets() {
210            None => candle::bail!("CodebookEncode, rhs has to be contiguous, got {rhs_layout:?}"),
211            Some((o1, o2)) => {
212                let slice = rhs_storage.as_slice::<f32>()?;
213                &slice[o1..o2]
214            }
215        };
216        let dst = (0..lhs_dim1)
217            .into_par_iter()
218            .map(|idx1| {
219                let mut where_min = 0;
220                let mut min_dist = f32::INFINITY;
221                let lhs = &lhs[idx1 * lhs_dim2..(idx1 + 1) * lhs_dim2];
222                for idx2 in 0..rhs_dim1 {
223                    let rhs = &rhs[idx2 * rhs_dim2..(idx2 + 1) * rhs_dim2];
224                    let mut dist = 0f32;
225                    for (a, b) in lhs.iter().zip(rhs.iter()) {
226                        dist += (a - b) * (a - b)
227                    }
228                    if dist < min_dist {
229                        min_dist = dist;
230                        where_min = idx2;
231                    }
232                }
233                where_min as u32
234            })
235            .collect();
236        let storage = candle::WithDType::to_cpu_storage_owned(dst);
237        Ok((storage, (lhs_dim1,).into()))
238    }
239}
240
241// https://github.com/huggingface/transformers/blob/abaca9f9432a84cfaa95531de4c72334f38a42f2/src/transformers/models/encodec/modeling_encodec.py#L340
242#[allow(unused)]
243#[derive(Clone, Debug)]
244pub struct EuclideanCodebook {
245    inited: Tensor,
246    cluster_size: Tensor,
247    embed: candle_nn::Embedding,
248    embed_avg: Tensor,
249    c2: Tensor,
250}
251
252impl EuclideanCodebook {
253    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
254        let inited = vb.get(1, "inited")?;
255        let cluster_size = vb.get(cfg.codebook_size, "cluster_size")?;
256        let e_shape = (cfg.codebook_size, cfg.codebook_dim());
257        let embed = vb.get(e_shape, "embed")?;
258        let c2 = ((&embed * &embed)?.sum(D::Minus1)? / 2.0)?;
259        let embed_avg = vb.get(e_shape, "embed_avg")?;
260        Ok(Self {
261            inited,
262            cluster_size,
263            embed: candle_nn::Embedding::new(embed, cfg.codebook_dim()),
264            embed_avg,
265            c2,
266        })
267    }
268
269    pub fn encode_slow(&self, xs: &Tensor) -> Result<Tensor> {
270        let mut target_shape = xs.dims().to_vec();
271        target_shape.pop();
272        let xs = xs.flatten_to(D::Minus2)?;
273        let _ = xs.dims2()?;
274        let dot_prod = xs.matmul(&self.embed.embeddings().t()?)?;
275        let codes = self.c2.broadcast_sub(&dot_prod)?.argmin(D::Minus1)?;
276        codes.reshape(target_shape)
277    }
278
279    pub fn encode(&self, xs: &Tensor) -> Result<Tensor> {
280        let mut target_shape = xs.dims().to_vec();
281        target_shape.pop();
282        let xs = xs.flatten_to(D::Minus2)?;
283        let _ = xs.dims2()?;
284        let codes = Tensor::apply_op2(&xs, self.embed.embeddings(), CodebookEncode)?;
285        codes.reshape(target_shape)
286    }
287
288    pub fn decode(&self, embed_ind: &Tensor) -> Result<Tensor> {
289        let quantize = self.embed.forward(embed_ind)?;
290        Ok(quantize)
291    }
292}
293
294#[derive(Clone, Debug)]
295pub struct VectorQuantization {
296    codebook: EuclideanCodebook,
297}
298
299impl VectorQuantization {
300    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
301        let codebook = EuclideanCodebook::new(cfg, vb.pp("codebook"))?;
302        Ok(Self { codebook })
303    }
304
305    pub fn encode(&self, xs: &Tensor) -> Result<Tensor> {
306        let xs = xs.transpose(1, 2)?;
307        self.codebook.encode_slow(&xs)
308    }
309
310    pub fn decode(&self, embed_ind: &Tensor) -> Result<Tensor> {
311        let quantize = self.codebook.decode(embed_ind)?;
312        let quantize = quantize.transpose(1, 2)?;
313        Ok(quantize)
314    }
315}
316
317#[derive(Clone, Debug)]
318pub struct ResidualVectorQuantizer {
319    layers: Vec<VectorQuantization>,
320    dtype: DType,
321}
322
323impl ResidualVectorQuantizer {
324    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
325        let vb = &vb.pp("layers");
326        let layers = (0..cfg.num_quantizers())
327            .map(|i| VectorQuantization::new(cfg, vb.pp(i)))
328            .collect::<Result<Vec<_>>>()?;
329        Ok(Self {
330            layers,
331            dtype: vb.dtype(),
332        })
333    }
334
335    pub fn encode(&self, xs: &Tensor) -> Result<Tensor> {
336        let mut codes = Vec::with_capacity(self.layers.len());
337        let mut residual = xs.clone();
338        for layer in self.layers.iter() {
339            let indices = layer.encode(&residual)?;
340            let quantized = layer.decode(&indices)?;
341            residual = (residual - quantized)?;
342            codes.push(indices)
343        }
344        Tensor::stack(&codes, 0)
345    }
346
347    pub fn decode(&self, codes: &Tensor) -> Result<Tensor> {
348        let mut quantized_out = Tensor::zeros((), self.dtype, codes.device())?;
349        let ncodes = codes.dim(0)?;
350        if ncodes > self.layers.len() {
351            candle::bail!(
352                "codes shape {:?} does not match the number of quantization layers {}",
353                codes.shape(),
354                self.layers.len()
355            )
356        }
357        for (i, layer) in self.layers.iter().take(ncodes).enumerate() {
358            let quantized = layer.decode(&codes.i(i)?)?;
359            quantized_out = quantized.broadcast_add(&quantized_out)?;
360        }
361        Ok(quantized_out)
362    }
363}
364
365// https://github.com/huggingface/transformers/blob/abaca9f9432a84cfaa95531de4c72334f38a42f2/src/transformers/models/encodec/modeling_encodec.py#L226
366#[derive(Clone, Debug)]
367pub struct EncodecLSTM {
368    layers: Vec<candle_nn::LSTM>,
369}
370
371impl EncodecLSTM {
372    pub fn new(dim: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
373        let vb = &vb.pp("lstm");
374        let mut layers = vec![];
375        for layer_idx in 0..cfg.num_lstm_layers {
376            let config = candle_nn::LSTMConfig {
377                layer_idx,
378                ..Default::default()
379            };
380            let lstm = candle_nn::lstm(dim, dim, config, vb.clone())?;
381            layers.push(lstm)
382        }
383        Ok(Self { layers })
384    }
385}
386
387impl Module for EncodecLSTM {
388    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
389        use candle_nn::RNN;
390        // This is different from the Python transformers version as candle LSTM is batch first.
391        let xs = xs.t()?;
392        let residual = &xs;
393        let mut xs = xs.clone();
394        for layer in self.layers.iter() {
395            let states = layer.seq(&xs)?;
396            xs = layer.states_to_tensor(&states)?;
397        }
398        let xs = (xs + residual)?.t()?;
399        Ok(xs)
400    }
401}
402
403#[derive(Clone, Debug)]
404pub struct EncodecConvTranspose1d {
405    conv: ConvTranspose1d,
406}
407
408impl EncodecConvTranspose1d {
409    fn new(
410        in_c: usize,
411        out_c: usize,
412        k: usize,
413        stride: usize,
414        _cfg: &Config,
415        vb: VarBuilder,
416    ) -> Result<Self> {
417        let cfg = candle_nn::ConvTranspose1dConfig {
418            stride,
419            ..Default::default()
420        };
421        let conv = conv_transpose1d_weight_norm(in_c, out_c, k, true, cfg, vb.pp("conv"))?;
422        Ok(Self { conv })
423    }
424}
425
426impl Module for EncodecConvTranspose1d {
427    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
428        xs.apply(&self.conv)
429    }
430}
431
432#[derive(Clone, Debug)]
433pub struct EncodecConv1d {
434    causal: bool,
435    conv: Conv1d,
436    norm: Option<candle_nn::GroupNorm>,
437    pad_mode: PadMode,
438}
439
440impl EncodecConv1d {
441    pub fn new(
442        in_c: usize,
443        out_c: usize,
444        kernel_size: usize,
445        stride: usize,
446        dilation: usize,
447        cfg: &Config,
448        vb: VarBuilder,
449    ) -> Result<Self> {
450        let conv = match cfg.norm_type {
451            NormType::WeightNorm => conv1d_weight_norm(
452                in_c,
453                out_c,
454                kernel_size,
455                candle_nn::Conv1dConfig {
456                    stride,
457                    dilation,
458                    ..Default::default()
459                },
460                vb.pp("conv"),
461            )?,
462            NormType::None | NormType::TimeGroupNorm => conv1d(
463                in_c,
464                out_c,
465                kernel_size,
466                candle_nn::Conv1dConfig {
467                    padding: 0,
468                    stride,
469                    groups: 1,
470                    dilation: 1,
471                    cudnn_fwd_algo: None,
472                },
473                vb.pp("conv"),
474            )?,
475        };
476        let norm = match cfg.norm_type {
477            NormType::None | NormType::WeightNorm => None,
478            NormType::TimeGroupNorm => {
479                let gn = candle_nn::group_norm(1, out_c, 1e-5, vb.pp("norm"))?;
480                Some(gn)
481            }
482        };
483        Ok(Self {
484            causal: cfg.use_causal_conv,
485            conv,
486            norm,
487            pad_mode: cfg.pad_mode,
488        })
489    }
490}
491
492impl Module for EncodecConv1d {
493    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
494        let (_b, _t, _c) = xs.dims3()?;
495        let k_size = self.conv.weight().dim(D::Minus1)?;
496        let conv_cfg = self.conv.config();
497        // Effective kernel size with dilations.
498        let k_size = (k_size - 1) * conv_cfg.dilation + 1;
499        let padding_total = k_size - conv_cfg.stride;
500        let extra_padding =
501            get_extra_padding_for_conv1d(xs, k_size, conv_cfg.stride, padding_total)?;
502        let xs = if self.causal {
503            pad1d(xs, padding_total, extra_padding, self.pad_mode)?
504        } else {
505            let padding_right = padding_total / 2;
506            let padding_left = padding_total - padding_right;
507            pad1d(
508                xs,
509                padding_left,
510                padding_right + extra_padding,
511                self.pad_mode,
512            )?
513        };
514        let xs = self.conv.forward(&xs)?;
515        match &self.norm {
516            None => Ok(xs),
517            Some(norm) => xs.apply(norm),
518        }
519    }
520}
521
522#[derive(Clone, Debug)]
523pub struct EncodecResnetBlock {
524    block_conv1: EncodecConv1d,
525    block_conv2: EncodecConv1d,
526    shortcut: Option<EncodecConv1d>,
527}
528
529impl EncodecResnetBlock {
530    pub fn new(
531        dim: usize,
532        (dilation1, dilation2): (usize, usize),
533        cfg: &Config,
534        vb: VarBuilder,
535    ) -> Result<Self> {
536        let h = dim / cfg.compress;
537        let mut layer = Layer::new(vb.pp("block"));
538        // TODO: Apply dilations!
539        layer.inc();
540        let block_conv1 = EncodecConv1d::new(
541            dim,
542            h,
543            cfg.residual_kernel_size,
544            1,
545            dilation1,
546            cfg,
547            layer.next(),
548        )?;
549        layer.inc();
550        let block_conv2 = EncodecConv1d::new(h, dim, 1, 1, dilation2, cfg, layer.next())?;
551        let shortcut = if cfg.use_conv_shortcut {
552            let conv = EncodecConv1d::new(dim, dim, 1, 1, 1, cfg, vb.pp("shortcut"))?;
553            Some(conv)
554        } else {
555            None
556        };
557        Ok(Self {
558            block_conv1,
559            block_conv2,
560            shortcut,
561        })
562    }
563}
564
565impl Module for EncodecResnetBlock {
566    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
567        let residual = xs.clone();
568        let xs = xs.elu(1.)?;
569        let xs = self.block_conv1.forward(&xs)?;
570        let xs = xs.elu(1.)?;
571        let xs = self.block_conv2.forward(&xs)?;
572        let xs = match &self.shortcut {
573            None => (xs + residual)?,
574            Some(shortcut) => xs.add(&shortcut.forward(&residual)?)?,
575        };
576        Ok(xs)
577    }
578}
579
580struct Layer<'a> {
581    vb: VarBuilder<'a>,
582    cnt: usize,
583}
584
585impl<'a> Layer<'a> {
586    fn new(vb: VarBuilder<'a>) -> Self {
587        Self { vb, cnt: 0 }
588    }
589
590    fn inc(&mut self) {
591        self.cnt += 1;
592    }
593
594    fn next(&mut self) -> VarBuilder<'_> {
595        let vb = self.vb.pp(self.cnt.to_string());
596        self.cnt += 1;
597        vb
598    }
599}
600
601#[derive(Clone, Debug)]
602pub struct Encoder {
603    init_conv: EncodecConv1d,
604    sampling_layers: Vec<(Vec<EncodecResnetBlock>, EncodecConv1d)>,
605    final_lstm: EncodecLSTM,
606    final_conv: EncodecConv1d,
607}
608
609impl Encoder {
610    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
611        let mut layer = Layer::new(vb.pp("layers"));
612        let init_conv = EncodecConv1d::new(
613            cfg.audio_channels,
614            cfg.num_filters,
615            cfg.kernel_size,
616            1,
617            1,
618            cfg,
619            layer.next(),
620        )?;
621        let mut sampling_layers = vec![];
622        let mut scaling = 1;
623        for &ratio in cfg.upsampling_ratios.iter().rev() {
624            let current_scale = scaling * cfg.num_filters;
625            let mut resnets = vec![];
626            for j in 0..(cfg.num_residual_layers as u32) {
627                let resnet = EncodecResnetBlock::new(
628                    current_scale,
629                    (cfg.dilation_growth_rate.pow(j), 1),
630                    cfg,
631                    layer.next(),
632                )?;
633                resnets.push(resnet)
634            }
635            layer.inc(); // ELU
636            let conv1d = EncodecConv1d::new(
637                current_scale,
638                current_scale * 2,
639                ratio * 2,
640                ratio,
641                1,
642                cfg,
643                layer.next(),
644            )?;
645            sampling_layers.push((resnets, conv1d));
646            scaling *= 2;
647        }
648        let final_lstm = EncodecLSTM::new(cfg.num_filters * scaling, cfg, layer.next())?;
649        layer.inc(); // ELU
650        let final_conv = EncodecConv1d::new(
651            cfg.num_filters * scaling,
652            cfg.hidden_size,
653            cfg.last_kernel_size,
654            1,
655            1,
656            cfg,
657            layer.next(),
658        )?;
659        Ok(Self {
660            init_conv,
661            sampling_layers,
662            final_conv,
663            final_lstm,
664        })
665    }
666}
667
668impl Module for Encoder {
669    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
670        let mut xs = xs.apply(&self.init_conv)?;
671        for (resnets, conv) in self.sampling_layers.iter() {
672            for resnet in resnets.iter() {
673                xs = xs.apply(resnet)?;
674            }
675            xs = xs.elu(1.0)?.apply(conv)?;
676        }
677        xs.apply(&self.final_lstm)?
678            .elu(1.0)?
679            .apply(&self.final_conv)
680    }
681}
682
683#[derive(Clone, Debug)]
684pub struct Decoder {
685    init_conv: EncodecConv1d,
686    init_lstm: EncodecLSTM,
687    sampling_layers: Vec<(EncodecConvTranspose1d, Vec<EncodecResnetBlock>)>,
688    final_conv: EncodecConv1d,
689}
690
691impl Decoder {
692    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
693        let mut layer = Layer::new(vb.pp("layers"));
694        let mut scaling = usize::pow(2, cfg.upsampling_ratios.len() as u32);
695        let init_conv = EncodecConv1d::new(
696            cfg.hidden_size,
697            cfg.num_filters * scaling,
698            cfg.last_kernel_size,
699            1,
700            1,
701            cfg,
702            layer.next(),
703        )?;
704        let init_lstm = EncodecLSTM::new(cfg.num_filters * scaling, cfg, layer.next())?;
705        let mut sampling_layers = vec![];
706        for &ratio in cfg.upsampling_ratios.iter() {
707            let current_scale = scaling * cfg.num_filters;
708            layer.inc(); // ELU
709            let conv1d = EncodecConvTranspose1d::new(
710                current_scale,
711                current_scale / 2,
712                ratio * 2,
713                ratio,
714                cfg,
715                layer.next(),
716            )?;
717            let mut resnets = vec![];
718            for j in 0..(cfg.num_residual_layers as u32) {
719                let resnet = EncodecResnetBlock::new(
720                    current_scale / 2,
721                    (cfg.dilation_growth_rate.pow(j), 1),
722                    cfg,
723                    layer.next(),
724                )?;
725                resnets.push(resnet)
726            }
727            sampling_layers.push((conv1d, resnets));
728            scaling /= 2;
729        }
730        layer.inc(); // ELU
731        let final_conv = EncodecConv1d::new(
732            cfg.num_filters,
733            cfg.audio_channels,
734            cfg.last_kernel_size,
735            1,
736            1,
737            cfg,
738            layer.next(),
739        )?;
740        Ok(Self {
741            init_conv,
742            init_lstm,
743            sampling_layers,
744            final_conv,
745        })
746    }
747}
748
749impl Module for Decoder {
750    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
751        let mut xs = xs.apply(&self.init_conv)?.apply(&self.init_lstm)?;
752        for (conv, resnets) in self.sampling_layers.iter() {
753            xs = xs.elu(1.)?.apply(conv)?;
754            for resnet in resnets.iter() {
755                xs = xs.apply(resnet)?
756            }
757        }
758        xs.elu(1.)?.apply(&self.final_conv)
759    }
760}
761
762#[derive(Debug)]
763pub struct Model {
764    encoder: Encoder,
765    decoder: Decoder,
766    quantizer: ResidualVectorQuantizer,
767}
768
769impl Model {
770    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
771        let encoder = Encoder::new(cfg, vb.pp("encoder"))?;
772        let decoder = Decoder::new(cfg, vb.pp("decoder"))?;
773        let quantizer = ResidualVectorQuantizer::new(cfg, vb.pp("quantizer"))?;
774        Ok(Self {
775            encoder,
776            decoder,
777            quantizer,
778        })
779    }
780
781    pub fn encode(&self, xs: &Tensor) -> Result<Tensor> {
782        let xs = self.encoder.forward(xs)?;
783        let codes = self.quantizer.encode(&xs)?;
784        codes.transpose(0, 1)
785    }
786
787    pub fn decode(&self, codes: &Tensor) -> Result<Tensor> {
788        let (_b_sz, _codebooks, _seqlen) = codes.dims3()?;
789        let codes = codes.transpose(0, 1)?;
790        let embeddings = self.quantizer.decode(&codes)?;
791        let outputs = self.decoder.forward(&embeddings)?;
792        Ok(outputs)
793    }
794}