Team Ai
Apppublic

codejin/diffsingerkr

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes
Diffusion.py403 linesDownload Raw Back to Modules
1import torch2import math3from argparse import Namespace4from typing import Optional, List, Dict, Union5from tqdm import tqdm6 7from .Layer import Conv1d, Lambda8 9class Diffusion(torch.nn.Module):10    def __init__(11        self,12        hyper_parameters: Namespace13        ):14        super().__init__()15        self.hp = hyper_parameters16 17        if self.hp.Feature_Type == 'Mel':18            self.feature_size = self.hp.Sound.Mel_Dim19        elif self.hp.Feature_Type == 'Spectrogram':20            self.feature_size = self.hp.Sound.N_FFT // 2 + 121 22        self.denoiser = Denoiser(23            hyper_parameters= self.hp24            )25 26        self.timesteps = self.hp.Diffusion.Max_Step27        betas = torch.linspace(1e-4, 0.06, self.timesteps)28        alphas = 1.0 - betas29        alphas_cumprod = torch.cumprod(alphas, axis= 0)30        alphas_cumprod_prev = torch.cat([torch.tensor([1.0]), alphas_cumprod[:-1]])31        32        # calculations for diffusion q(x_t | x_{t-1}) and others33        self.register_buffer('alphas_cumprod', alphas_cumprod)  # [Diffusion_t]34        self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)  # [Diffusion_t]35        self.register_buffer('sqrt_alphas_cumprod', alphas_cumprod.sqrt())36        self.register_buffer('sqrt_one_minus_alphas_cumprod', (1.0 - alphas_cumprod).sqrt())37        self.register_buffer('sqrt_recip_alphas_cumprod', (1.0 / alphas_cumprod).sqrt())38        self.register_buffer('sqrt_recipm1_alphas_cumprod', (1.0 / alphas_cumprod - 1.0).sqrt())39 40        # calculations for posterior q(x_{t-1} | x_t, x_0)41        posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)42 43        # below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain44        self.register_buffer('posterior_log_variance', torch.maximum(posterior_variance, torch.tensor([1e-20])).log())45        self.register_buffer('posterior_mean_coef1', betas * alphas_cumprod_prev.sqrt() / (1.0 - alphas_cumprod))46        self.register_buffer('posterior_mean_coef2', (1.0 - alphas_cumprod_prev) * alphas.sqrt() / (1.0 - alphas_cumprod))47 48    def forward(49        self,50        encodings: torch.Tensor,51        features: torch.Tensor= None52        ):53        '''54        encodings: [Batch, Enc_d, Enc_t]55        features: [Batch, Feature_d, Feature_t]56        feature_lengths: [Batch]57        '''58        if not features is None:    # train59            diffusion_steps = torch.randint(60                low= 0,61                high= self.timesteps,62                size= (encodings.size(0),),63                dtype= torch.long,64                device= encodings.device65                )    # random single step66            67            noises, epsilons = self.Get_Noise_Epsilon_for_Train(68                features= features,69                encodings= encodings,70                diffusion_steps= diffusion_steps,71                )72            return None, noises, epsilons73        else:   # inference74            features = self.Sampling(75                encodings= encodings,76                )77            return features, None, None78 79    def Sampling(80        self,81        encodings: torch.Tensor,82        ):83        features = torch.randn(84            size= (encodings.size(0), self.feature_size, encodings.size(2)),85            device= encodings.device86            )87        for diffusion_step in reversed(range(self.timesteps)):88            features = self.P_Sampling(89                features= features,90                encodings= encodings,91                diffusion_steps= torch.full(92                    size= (encodings.size(0), ),93                    fill_value= diffusion_step,94                    dtype= torch.long,95                    device= encodings.device96                    ),97                )98        99        return features100 101    def P_Sampling(102        self,103        features: torch.Tensor,104        encodings: torch.Tensor,105        diffusion_steps: torch.Tensor,106        ):107        posterior_means, posterior_log_variances = self.Get_Posterior(108            features= features,109            encodings= encodings,110            diffusion_steps= diffusion_steps,111            )112 113        noises = torch.randn_like(features) # [Batch, Feature_d, Feature_d]114        masks = (diffusion_steps > 0).float().unsqueeze(1).unsqueeze(1) #[Batch, 1, 1]115        116        return posterior_means + masks * (0.5 * posterior_log_variances).exp() * noises117 118    def Get_Posterior(119        self,120        features: torch.Tensor,121        encodings: torch.Tensor,122        diffusion_steps: torch.Tensor123        ):124        noised_predictions = self.denoiser(125            features= features,126            encodings= encodings,127            diffusion_steps= diffusion_steps128            )129 130        epsilons = \131            features * self.sqrt_recip_alphas_cumprod[diffusion_steps][:, None, None] - \132            noised_predictions * self.sqrt_recipm1_alphas_cumprod[diffusion_steps][:, None, None]133        epsilons.clamp_(-1.0, 1.0)  # clipped134        135        posterior_means = \136            epsilons * self.posterior_mean_coef1[diffusion_steps][:, None, None] + \137            features * self.posterior_mean_coef2[diffusion_steps][:, None, None]138        posterior_log_variances = \139            self.posterior_log_variance[diffusion_steps][:, None, None]140 141        return posterior_means, posterior_log_variances142 143    def Get_Noise_Epsilon_for_Train(144        self,145        features: torch.Tensor,146        encodings: torch.Tensor,147        diffusion_steps: torch.Tensor,148        ):149        noises = torch.randn_like(features)150 151        noised_features = \152            features * self.sqrt_alphas_cumprod[diffusion_steps][:, None, None] + \153            noises * self.sqrt_one_minus_alphas_cumprod[diffusion_steps][:, None, None]154 155        epsilons = self.denoiser(156            features= noised_features,157            encodings= encodings,158            diffusion_steps= diffusion_steps159            )160        161        return noises, epsilons162 163    def DDIM(164        self,165        encodings: torch.Tensor,166        ddim_steps: int,167        eta: float= 0.0,168        temperature: float= 1.0,169        use_tqdm: bool= False170        ):171        ddim_timesteps = self.Get_DDIM_Steps(172            ddim_steps= ddim_steps173            )174        sigmas, alphas, alphas_prev = self.Get_DDIM_Sampling_Parameters(175            ddim_timesteps= ddim_timesteps,176            eta= eta177            )178        sqrt_one_minus_alphas = (1. - alphas).sqrt()179 180        features = torch.randn(181            size= (encodings.size(0), self.feature_size, encodings.size(2)),182            device= encodings.device183            )184 185        setp_range = reversed(range(ddim_steps))186        if use_tqdm:187            tqdm(188                setp_range,189                desc= '[Diffusion]',190                total= ddim_steps191                )192 193        for diffusion_steps in setp_range:194            noised_predictions = self.denoiser(195                features= features,196                encodings= encodings,197                diffusion_steps= torch.full(198                    size= (encodings.size(0), ),199                    fill_value= diffusion_steps,200                    dtype= torch.long,201                    device= encodings.device202                    )203                )204 205            feature_starts = (features - sqrt_one_minus_alphas[diffusion_steps] * noised_predictions) / alphas[diffusion_steps].sqrt()206            direction_pointings = (1.0 - alphas_prev[diffusion_steps] - sigmas[diffusion_steps].pow(2.0)) * noised_predictions207            noises = sigmas[diffusion_steps] * torch.randn_like(features) * temperature208 209            features = alphas_prev[diffusion_steps].sqrt() * feature_starts + direction_pointings + noises210 211        return features212 213    # https://github.com/CompVis/stable-diffusion/blob/main/ldm/modules/diffusionmodules/util.py214    def Get_DDIM_Steps(215        self,        216        ddim_steps: int,217        ddim_discr_method: str= 'uniform'218        ):219        if ddim_discr_method == 'uniform':            220            ddim_timesteps = torch.arange(0, self.timesteps, self.timesteps // ddim_steps).long()221        elif ddim_discr_method == 'quad':222            ddim_timesteps = torch.linspace(0, (torch.tensor(self.timesteps) * 0.8).sqrt(), ddim_steps).pow(2.0).long()223        else:224            raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')225        226        ddim_timesteps[-1] = self.timesteps - 1227 228        return ddim_timesteps229 230    def Get_DDIM_Sampling_Parameters(self, ddim_timesteps, eta):231        alphas = self.alphas_cumprod[ddim_timesteps]232        alphas_prev = self.alphas_cumprod_prev[ddim_timesteps]233        sigmas = eta * ((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev)).sqrt()234 235        return sigmas, alphas, alphas_prev236 237class Denoiser(torch.nn.Module):238    def __init__(239        self,240        hyper_parameters: Namespace241        ):242        super().__init__()243        self.hp = hyper_parameters244 245        if self.hp.Feature_Type == 'Mel':246            feature_size = self.hp.Sound.Mel_Dim247        elif self.hp.Feature_Type == 'Spectrogram':248            feature_size = self.hp.Sound.N_FFT // 2 + 1249 250        self.prenet = torch.nn.Sequential(251            Conv1d(252                in_channels= feature_size,253                out_channels= self.hp.Diffusion.Size,254                kernel_size= 1,255                w_init_gain= 'relu'256                ),257            torch.nn.Mish()258            )259 260        self.step_ffn = torch.nn.Sequential(261            Diffusion_Embedding(262                channels= self.hp.Diffusion.Size263                ),264            Lambda(lambda x: x.unsqueeze(2)),265            Conv1d(266                in_channels= self.hp.Diffusion.Size,267                out_channels= self.hp.Diffusion.Size * 4,268                kernel_size= 1,269                w_init_gain= 'relu'270                ),271            torch.nn.Mish(),272            Conv1d(273                in_channels= self.hp.Diffusion.Size * 4,274                out_channels= self.hp.Diffusion.Size,275                kernel_size= 1,276                w_init_gain= 'linear'277                )278            )279 280        self.residual_blocks = torch.nn.ModuleList([281            Residual_Block(282                in_channels= self.hp.Diffusion.Size,283                kernel_size= self.hp.Diffusion.Kernel_Size,284                condition_channels= self.hp.Encoder.Size + feature_size285                )286            for _ in range(self.hp.Diffusion.Stack)287            ])288        289        self.projection =  torch.nn.Sequential(290            Conv1d(291                in_channels= self.hp.Diffusion.Size,292                out_channels= self.hp.Diffusion.Size,293                kernel_size= 1,294                w_init_gain= 'relu'295                ),296            torch.nn.ReLU(),297            Conv1d(298                in_channels= self.hp.Diffusion.Size,299                out_channels= feature_size,300                kernel_size= 1301                ),302            )303        torch.nn.init.zeros_(self.projection[-1].weight)    # This is key factor....304 305    def forward(306        self,307        features: torch.Tensor,308        encodings: torch.Tensor,309        diffusion_steps: torch.Tensor310        ):311        '''312        features: [Batch, Feature_d, Feature_t]313        encodings: [Batch, Enc_d, Feature_t]314        diffusion_steps: [Batch]315        '''316        x = self.prenet(features)317        318        diffusion_steps = self.step_ffn(diffusion_steps) # [Batch, Res_d, 1]319        320        skips_list = []321        for residual_block in self.residual_blocks:322            x, skips = residual_block(323                x= x,324                conditions= encodings,325                diffusion_steps= diffusion_steps326                )327            skips_list.append(skips)328 329        x = torch.stack(skips_list, dim= 0).sum(dim= 0) / math.sqrt(self.hp.Diffusion.Stack)330        x = self.projection(x)331 332        return x333 334class Diffusion_Embedding(torch.nn.Module):335    def __init__(336        self,337        channels: int338        ):339        super().__init__()340        self.channels = channels341 342    def forward(self, x: torch.Tensor):343        half_channels = self.channels // 2  # sine and cosine344        embeddings = math.log(10000.0) / (half_channels - 1)345        embeddings = torch.exp(torch.arange(half_channels, device= x.device) * -embeddings)346        embeddings = x.unsqueeze(1) * embeddings.unsqueeze(0)347        embeddings = torch.cat([embeddings.sin(), embeddings.cos()], dim= -1)348 349        return embeddings350 351class Residual_Block(torch.nn.Module):352    def __init__(353        self,354        in_channels: int,355        kernel_size: int,356        condition_channels: int357        ):358        super().__init__()359        self.in_channels = in_channels360        361        self.condition = Conv1d(362            in_channels= condition_channels,363            out_channels= in_channels * 2,364            kernel_size= 1365            )366        self.diffusion_step = Conv1d(367            in_channels= in_channels,368            out_channels= in_channels,369            kernel_size= 1370            )371 372        self.conv = Conv1d(373            in_channels= in_channels,374            out_channels= in_channels * 2,375            kernel_size= kernel_size,376            padding= kernel_size // 2377            )378 379        self.projection = Conv1d(380            in_channels= in_channels,381            out_channels= in_channels * 2,382            kernel_size= 1383            )384 385    def forward(386        self,387        x: torch.Tensor,388        conditions: torch.Tensor,389        diffusion_steps: torch.Tensor390        ):391        residuals = x392 393        conditions = self.condition(conditions)394        diffusion_steps = self.diffusion_step(diffusion_steps)395 396        x = self.conv(x + diffusion_steps) + conditions397        x_a, x_b = x.chunk(chunks= 2, dim= 1)398        x = x_a.sigmoid() * x_b.tanh()399 400        x = self.projection(x)401        x, skips = x.chunk(chunks= 2, dim= 1)402 403        return (x + residuals) / math.sqrt(2.0), skips