(self, x)
| 771 | ) |
| 772 | |
| 773 | def forward(self, x): |
| 774 | x = self.conv_in(x) |
| 775 | for block in self.res_block1: |
| 776 | x = block(x, None) |
| 777 | x = torch.nn.functional.interpolate(x, size=(int(round(x.shape[2]*self.factor)), int(round(x.shape[3]*self.factor)))) |
| 778 | x = self.attn(x) |
| 779 | for block in self.res_block2: |
| 780 | x = block(x, None) |
| 781 | x = self.conv_out(x) |
| 782 | return x |
| 783 | |
| 784 | |
| 785 | class MergedRescaleEncoder(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected