| 169 | |
| 170 | class MoBYMLP(nn.Module): |
| 171 | def __init__(self, in_dim=256, inner_dim=4096, out_dim=256, num_layers=2): |
| 172 | super(MoBYMLP, self).__init__() |
| 173 | |
| 174 | # hidden layers |
| 175 | linear_hidden = [nn.Identity()] |
| 176 | for i in range(num_layers - 1): |
| 177 | linear_hidden.append(nn.Linear(in_dim if i == 0 else inner_dim, inner_dim)) |
| 178 | linear_hidden.append(nn.BatchNorm1d(inner_dim)) |
| 179 | linear_hidden.append(nn.ReLU(inplace=True)) |
| 180 | self.linear_hidden = nn.Sequential(*linear_hidden) |
| 181 | |
| 182 | self.linear_out = nn.Linear(in_dim if num_layers == 1 else inner_dim, out_dim) if num_layers >= 1 else nn.Identity() |
| 183 | |
| 184 | def forward(self, x): |
| 185 | x = self.linear_hidden(x) |