OpenMotionLab/MotionGPT
118
1import torch2import torch.nn as nn3import torch.nn.functional as F4 5class AdaptiveInstanceNorm1d(nn.Module):6 def __init__(self, num_features, eps=1e-5, momentum=0.1):7 super(AdaptiveInstanceNorm1d, self).__init__()8 self.num_features = num_features9 self.eps = eps10 self.momentum = momentum11 self.weight = None12 self.bias = None13 self.register_buffer('running_mean', torch.zeros(num_features))14 self.register_buffer('running_var', torch.ones(num_features))15 16 def forward(self, x, direct_weighting=False, no_std=False):17 assert self.weight is not None and \18 self.bias is not None, "Please assign AdaIN weight first"19 # (bs, nfeats, nframe) <= (nframe, bs, nfeats)20 x = x.permute(1,2,0) 21 22 b, c = x.size(0), x.size(1) # batch size & channels23 running_mean = self.running_mean.repeat(b)24 running_var = self.running_var.repeat(b)25 # self.weight = torch.ones_like(self.weight)26 27 if direct_weighting:28 x_reshaped = x.contiguous().view(b * c)29 if no_std:30 out = x_reshaped + self.bias31 else:32 out = x_reshaped.mul(self.weight) + self.bias33 out = out.view(b, c, *x.size()[2:])34 else:35 x_reshaped = x.contiguous().view(1, b * c, *x.size()[2:]) 36 out = F.batch_norm(37 x_reshaped, running_mean, running_var, self.weight, self.bias,38 True, self.momentum, self.eps)39 out = out.view(b, c, *x.size()[2:])40 41 # (nframe, bs, nfeats) <= (bs, nfeats, nframe)42 out = out.permute(2,0,1) 43 return out44 45 def __repr__(self):46 return self.__class__.__name__ + '(' + str(self.num_features) + ')'47 48def assign_adain_params(adain_params, model):49 # assign the adain_params to the AdaIN layers in model50 for m in model.modules():51 if m.__class__.__name__ == "AdaptiveInstanceNorm1d":52 mean = adain_params[: , : m.num_features]53 std = adain_params[: , m.num_features: 2 * m.num_features]54 m.bias = mean.contiguous().view(-1)55 m.weight = std.contiguous().view(-1)56 if adain_params.size(1) > 2 * m.num_features:57 adain_params = adain_params[: , 2 * m.num_features:]58 59 60def get_num_adain_params(model):61 # return the number of AdaIN parameters needed by the model62 num_adain_params = 063 for m in model.modules():64 if m.__class__.__name__ == "AdaptiveInstanceNorm1d":65 num_adain_params += 2 * m.num_features66 return num_adain_params67 