(self, x, timesteps, low_res=None, **kwargs)
| 685 | super().__init__(image_size, in_channels * 2, *args, **kwargs) |
| 686 | |
| 687 | def forward(self, x, timesteps, low_res=None, **kwargs): |
| 688 | _, _, new_height, new_width = x.shape |
| 689 | upsampled = F.interpolate(low_res, (new_height, new_width), mode="bilinear") |
| 690 | x = th.cat([x, upsampled], dim=1) |
| 691 | return super().forward(x, timesteps, **kwargs) |
| 692 | |
| 693 | |
| 694 | class EncoderUNetModel(nn.Module): |