| 319 | return torch.cat(sampled_latents_pyramid, dim=-1) |
| 320 | |
| 321 | class FeedForward(nn.Module): |
| 322 | def __init__(self, dim, dropout = 0.): |
| 323 | super().__init__() |
| 324 | self.net = nn.Sequential( |
| 325 | nn.Linear(dim, dim), |
| 326 | nn.GELU(), |
| 327 | nn.Dropout(dropout), |
| 328 | nn.Linear(dim, dim), |
| 329 | nn.Dropout(dropout) |
| 330 | ) |
| 331 | def forward(self, x): |
| 332 | x = self.net(x) |
| 333 | return x |
| 334 | |
| 335 | class MLP(nn.Module): |
| 336 | def __init__(self, in_dim=22, out_dim=1, innter_dim=96, depth=5): |
nothing calls this directly
no outgoing calls
no test coverage detected