OpenMotionLab/MotionGPT
118
1import torch2import torch.nn as nn3import torch.nn.functional as F4from mGPT.models.notused import AdaptiveInstanceNorm1d5 6 7class MLP(nn.Module):8 9 def __init__(self, cfg, out_dim, is_init):10 super(MLP, self).__init__()11 dims = cfg.MODEL.MOTION_DECODER.MLP_DIM12 n_blk = len(dims)13 norm = 'none'14 acti = 'lrelu'15 16 layers = []17 for i in range(n_blk - 1):18 layers += LinearBlock(dims[i], dims[i + 1], norm=norm, acti=acti)19 layers += LinearBlock(dims[-1], out_dim, norm='none', acti='none')20 self.model = nn.Sequential(*layers)21 22 if is_init:23 for m in self.modules():24 if isinstance(m, nn.Linear):25 #nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')26 nn.init.constant_(m.weight, 1)27 elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):28 nn.init.constant_(m.weight, 1)29 nn.init.constant_(m.bias, 0)30 31 def forward(self, x):32 return self.model(x.view(x.size(0), -1))33 34 35def ZeroPad1d(sizes):36 return nn.ConstantPad1d(sizes, 0)37 38 39def get_acti_layer(acti='relu', inplace=True):40 41 if acti == 'relu':42 return [nn.ReLU(inplace=inplace)]43 elif acti == 'lrelu':44 return [nn.LeakyReLU(0.2, inplace=inplace)]45 elif acti == 'tanh':46 return [nn.Tanh()]47 elif acti == 'none':48 return []49 else:50 assert 0, "Unsupported activation: {}".format(acti)51 52 53def get_norm_layer(norm='none', norm_dim=None):54 55 if norm == 'bn':56 return [nn.BatchNorm1d(norm_dim)]57 elif norm == 'in':58 # return [nn.InstanceNorm1d(norm_dim, affine=False)] # for rt42!59 return [nn.InstanceNorm1d(norm_dim, affine=True)]60 elif norm == 'adain':61 return [AdaptiveInstanceNorm1d(norm_dim)]62 elif norm == 'none':63 return []64 else:65 assert 0, "Unsupported normalization: {}".format(norm)66 67 68def get_dropout_layer(dropout=None):69 if dropout is not None:70 return [nn.Dropout(p=dropout)]71 else:72 return []73 74 75def ConvLayers(kernel_size,76 in_channels,77 out_channels,78 stride=1,79 pad_type='reflect',80 use_bias=True):81 """82 returns a list of [pad, conv] => should be += to some list, then apply sequential83 """84 85 if pad_type == 'reflect':86 pad = nn.ReflectionPad1d87 elif pad_type == 'replicate':88 pad = nn.ReplicationPad1d89 elif pad_type == 'zero':90 pad = ZeroPad1d91 else:92 assert 0, "Unsupported padding type: {}".format(pad_type)93 94 pad_l = (kernel_size - 1) // 295 pad_r = kernel_size - 1 - pad_l96 return [97 pad((pad_l, pad_r)),98 nn.Conv1d(in_channels,99 out_channels,100 kernel_size=kernel_size,101 stride=stride,102 bias=use_bias)103 ]104 105 106def ConvBlock(kernel_size,107 in_channels,108 out_channels,109 stride=1,110 pad_type='reflect',111 dropout=None,112 norm='none',113 acti='lrelu',114 acti_first=False,115 use_bias=True,116 inplace=True):117 """118 returns a list of [pad, conv, norm, acti] or [acti, pad, conv, norm]119 """120 121 layers = ConvLayers(kernel_size,122 in_channels,123 out_channels,124 stride=stride,125 pad_type=pad_type,126 use_bias=use_bias)127 layers += get_dropout_layer(dropout)128 layers += get_norm_layer(norm, norm_dim=out_channels)129 acti_layers = get_acti_layer(acti, inplace=inplace)130 131 if acti_first:132 return acti_layers + layers133 else:134 return layers + acti_layers135 136 137def LinearBlock(in_dim, out_dim, dropout=None, norm='none', acti='relu'):138 139 use_bias = True140 layers = []141 layers.append(nn.Linear(in_dim, out_dim, bias=use_bias))142 layers += get_dropout_layer(dropout)143 layers += get_norm_layer(norm, norm_dim=out_dim)144 layers += get_acti_layer(acti)145 146 return layers147 