File size: 3,775 Bytes
7b4c0bf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
import math

import torch
from diffusers import ConfigMixin, ModelMixin
from diffusers.configuration_utils import register_to_config
from diffusers.models.autoencoders.autoencoder_oobleck import (
    AutoencoderOobleckOutput,
    OobleckDecoder,
    OobleckDecoderOutput,
    OobleckDiagonalGaussianDistribution,
    OobleckEncoder,
)
from diffusers.utils.accelerate_utils import apply_forward_hook
from torch import nn
from torch.nn.utils import weight_norm


class YuE2VAE(ModelMixin, ConfigMixin):
    r"""
    The YuE2 VAE: diffusers' Oobleck encoder and decoder, between 48 kHz stereo audio and 64-channel latents at 25
    frames per second.

    [`AutoencoderOobleck`] cannot hold this checkpoint because it ties the encoder's output channels to its width; the
    YuE2 encoder is 64 channels wide and outputs a 64-channel mean and a 64-channel scale.
    """

    _keep_in_fp32_modules = ["encoder", "decoder"]

    @register_to_config
    def __init__(
        self,
        encoder_hidden_size: int = 64,
        downsampling_ratios: list[int] = [2, 2, 4, 4, 5, 6],
        channel_multiples: list[int] = [1, 2, 4, 8, 16, 32],
        decoder_channels: int = 64,
        decoder_input_channels: int = 64,
        audio_channels: int = 2,
        sampling_rate: int = 48000,
    ):
        super().__init__()
        self.hop_length = math.prod(downsampling_ratios)
        self.encoder = OobleckEncoder(
            encoder_hidden_size=encoder_hidden_size,
            audio_channels=audio_channels,
            downsampling_ratios=list(downsampling_ratios),
            channel_multiples=list(channel_multiples),
        )
        # The final projection outputs the posterior's mean and scale rather than `encoder_hidden_size` channels.
        self.encoder.conv2 = weight_norm(
            nn.Conv1d(
                encoder_hidden_size * channel_multiples[-1], 2 * decoder_input_channels, kernel_size=3, padding=1
            )
        )
        self.decoder = OobleckDecoder(
            channels=decoder_channels,
            input_channels=decoder_input_channels,
            audio_channels=audio_channels,
            upsampling_ratios=list(downsampling_ratios)[::-1],
            channel_multiples=list(channel_multiples),
        )

    @apply_forward_hook
    def encode(
        self, x: torch.Tensor, return_dict: bool = True
    ) -> AutoencoderOobleckOutput | tuple[OobleckDiagonalGaussianDistribution]:
        """
        Args:
            x (`torch.Tensor` of shape `(batch_size, audio_channels, num_samples)`):
                48 kHz audio.
            return_dict (`bool`, defaults to `True`):
                Whether to return an [`AutoencoderOobleckOutput`] instead of a plain tuple.

        Returns:
            [`AutoencoderOobleckOutput`] or `tuple`: the latent posterior. Its `mode()` (the mean) is the deterministic
            encoding; `sample(generator)` draws from it.
        """
        posterior = OobleckDiagonalGaussianDistribution(self.encoder(x))
        return AutoencoderOobleckOutput(latent_dist=posterior) if return_dict else (posterior,)

    @apply_forward_hook
    def decode(self, z: torch.Tensor, return_dict: bool = True) -> OobleckDecoderOutput | tuple[torch.Tensor]:
        """
        Args:
            z (`torch.Tensor` of shape `(batch_size, decoder_input_channels, num_frames)`):
                Acoustic latents.
            return_dict (`bool`, defaults to `True`):
                Whether to return an [`OobleckDecoderOutput`] instead of a plain tuple.

        Returns:
            [`OobleckDecoderOutput`] or `tuple`: audio of shape `(batch_size, audio_channels, num_samples)`.
        """
        sample = self.decoder(z)
        return OobleckDecoderOutput(sample=sample) if return_dict else (sample,)