CoolFace
Apppublic

meng2003/music2dance

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
act_norm.py185 linesDownload Raw Back to flowplusplus
1import torch2import torch.nn as nn3 4from models.util import mean_dim5 6 7class _BaseNorm(nn.Module):8    """Base class for ActNorm (Glow) and PixNorm (Flow++).9 10    The mean and inv_std get initialized using the mean and variance of the11    first mini-batch. After the init, mean and inv_std are trainable parameters.12 13    Adapted from:14        > https://github.com/openai/glow15    """16    def __init__(self, num_channels, height, width):17        super(_BaseNorm, self).__init__()18 19        # Input gets concatenated along channel axis20        #num_channels *= 221 22        self.register_buffer('is_initialized', torch.zeros(1))23        self.mean = nn.Parameter(torch.zeros(1, num_channels, height, width))24        self.inv_std = nn.Parameter(torch.zeros(1, num_channels, height, width))25        self.eps = 1e-626 27    def initialize_parameters(self, x):28        if not self.training:29            return30 31        with torch.no_grad():32            mean, inv_std = self._get_moments(x)33            self.mean.data.copy_(mean.data)34            self.inv_std.data.copy_(inv_std.data)35            self.is_initialized += 1.36 37    def _center(self, x, reverse=False):38        if reverse:39            return x + self.mean40        else:41            return x - self.mean42 43    def _get_moments(self, x):44        raise NotImplementedError('Subclass of _BaseNorm must implement _get_moments')45 46    def _scale(self, x, sldj, reverse=False):47        raise NotImplementedError('Subclass of _BaseNorm must implement _scale')48 49    def forward(self, x, cond, ldj=None, reverse=False):50        #import pdb;pdb.set_trace()51        x = torch.cat(x, dim=1)52        # import pdb;pdb.set_trace()53        if not self.is_initialized:54            print("Initializing norm Layer!")55            self.initialize_parameters(x)56 57        if reverse:58            x, ldj = self._scale(x, ldj, reverse)59            x = self._center(x, reverse)60        else:61            x = self._center(x, reverse)62            x, ldj = self._scale(x, ldj, reverse)63        x = x.chunk(2, dim=1)64 65        return x, ldj66 67 68class BatchNorm(nn.Module):69    def __init__(self, num_channels, momentum=0.1):70        super(BatchNorm, self).__init__()71        self.gamma = nn.Parameter(torch.ones(1, num_channels, 1, 1))72        self.beta = nn.Parameter(torch.zeros(1, num_channels, 1, 1))73        self.register_buffer('running_mean', torch.zeros(1, num_channels, 1, 1))74        self.register_buffer('running_var', torch.ones(1, num_channels, 1, 1))75        self.eps = 1e-576        self.momentum = momentum77        self.inv_std = None78        self.register_buffer('is_initialized', torch.zeros(1))79 80    def _get_moments(self, x):81        mean = mean_dim(x.clone(), dim=[0, 2, 3], keepdims=True).detach()82        var = mean_dim((x.clone() - mean) ** 2, dim=[0, 2, 3], keepdims=True).detach()83        # inv_std = 1. / (var.sqrt() + self.eps)84        if not self.is_initialized:85            self.running_mean.data.copy_(mean.data)86            self.running_var.data.copy_(var.data)87            self.is_initialized += 1.88        else:89            if self.momentum < 1.0:90                self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * mean91                self.running_var = (1 - self.momentum) * self.running_var + self.momentum * var92            else:93                self.running_mean.data.copy_(mean.data)94                self.running_var.data.copy_(var.data)95 96    def forward(self, x, cond, ldj=None, reverse=False):97        # import pdb;pdb.set_trace()98        x = torch.cat(x, dim=1)99        if self.training:100            # print("HI")101            self._get_moments(x)102        # print(self.running_var[0])103        inv_std = 1. / (self.running_var.sqrt() + self.eps)104        if reverse:105            x = self._center(x, self.beta, reverse)106            x, ldj = self._scale(x, ldj, self.gamma, reverse)107            x, ldj = self._scale(x, ldj, inv_std, reverse)108            x = self._center(x, self.running_mean, reverse)109        else:110            x = self._center(x, self.running_mean, reverse)111            x, ldj = self._scale(x, ldj, inv_std, reverse)112            x, ldj = self._scale(x, ldj, self.gamma, reverse)113            x = self._center(x, self.beta, reverse)114        x = x.chunk(2, dim=1)115 116        return x, ldj117 118    def _center(self, x, centerer, reverse=False):119        if reverse:120            return x + centerer121        else:122            return x - centerer123 124    def _scale(self, x, sldj, scaler, reverse=False):125        if reverse:126            x = x / scaler127            sldj = sldj - scaler.log().sum() * x.size(2) * x.size(3)128        else:129            x = x * scaler130            sldj = sldj + scaler.log().sum() * x.size(2) * x.size(3)131 132        return x, sldj133 134class ActNorm(_BaseNorm):135    """Activation Normalization used in Glow136 137    The mean and inv_std get initialized using the mean and variance of the138    first mini-batch. After the init, mean and inv_std are trainable parameters.139    """140    def __init__(self, num_channels):141        super(ActNorm, self).__init__(num_channels, 1, 1)142 143    def _get_moments(self, x):144        mean = mean_dim(x.clone(), dim=[0, 2, 3], keepdims=True)145        var = mean_dim((x.clone() - mean) ** 2, dim=[0, 2, 3], keepdims=True)146        inv_std = 1. / (var.sqrt() + self.eps)147 148        return mean, inv_std149 150    def _scale(self, x, sldj, reverse=False):151        if reverse:152            x = x / self.inv_std153            sldj = sldj - self.inv_std.log().sum() * x.size(2) * x.size(3)154        else:155            x = x * self.inv_std156            sldj = sldj + self.inv_std.log().sum() * x.size(2) * x.size(3)157 158        return x, sldj159 160 161class PixNorm(_BaseNorm):162    """Pixel-wise Activation Normalization used in Flow++163 164    Normalizes every activation independently (note this differs from the variant165    used in in Glow, where they normalize each channel). The mean and stddev get166    initialized using the mean and stddev of the first mini-batch. After the167    initialization, `mean` and `inv_std` become trainable parameters.168    """169    def _get_moments(self, x):170        mean = torch.mean(x.clone(), dim=0, keepdim=True)171        var = torch.mean((x.clone() - mean) ** 2, dim=0, keepdim=True)172        inv_std = 1. / (var.sqrt() + self.eps)173 174        return mean, inv_std175 176    def _scale(self, x, sldj, reverse=False):177        if reverse:178            x = x / self.inv_std179            sldj = sldj - self.inv_std.log().sum()180        else:181            x = x * self.inv_std182            sldj = sldj + self.inv_std.log().sum()183 184        return x, sldj185