meng2003/music2dance
0
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 