1use candle::{DType, IndexOp, Layout, Module, Result, Shape, Tensor, D};
8use candle_nn::{conv1d, Conv1d, ConvTranspose1d, VarBuilder};
9
10#[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 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
126pub 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#[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#[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 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 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 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(); 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(); 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(); 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(); 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}