| 68 | |
| 69 | |
| 70 | class DownSampler(torch.nn.Module): |
| 71 | def __init__(self, channels, padding=1, extra_padding=False): |
| 72 | super().__init__() |
| 73 | self.conv = torch.nn.Conv2d(channels, channels, 3, stride=2, padding=padding) |
| 74 | self.extra_padding = extra_padding |
| 75 | |
| 76 | def forward(self, hidden_states, time_emb, text_emb, res_stack, **kwargs): |
| 77 | if self.extra_padding: |
| 78 | hidden_states = torch.nn.functional.pad(hidden_states, (0, 1, 0, 1), mode="constant", value=0) |
| 79 | hidden_states = self.conv(hidden_states) |
| 80 | return hidden_states, time_emb, text_emb, res_stack |
| 81 | |
| 82 | |
| 83 | class UpSampler(torch.nn.Module): |