| 77 | |
| 78 | |
| 79 | class Down(nn.Module): |
| 80 | def __init__(self, in_channels, out_channels, emb_dim=256): |
| 81 | super().__init__() |
| 82 | self.maxpool_conv = nn.Sequential( |
| 83 | nn.MaxPool2d(2), |
| 84 | DoubleConv(in_channels, in_channels, residual=True), |
| 85 | DoubleConv(in_channels, out_channels), |
| 86 | ) |
| 87 | |
| 88 | self.emb_layer = nn.Sequential( |
| 89 | nn.SiLU(), |
| 90 | nn.Linear( |
| 91 | emb_dim, |
| 92 | out_channels |
| 93 | ), |
| 94 | ) |
| 95 | |
| 96 | def forward(self, x, t): |
| 97 | x = self.maxpool_conv(x) |
| 98 | emb = self.emb_layer(t)[:, :, None, None].repeat(1, 1, x.shape[-2], x.shape[-1]) |
| 99 | return x + emb |
| 100 | |
| 101 | |
| 102 | class Up(nn.Module): |