(self, in_channels, out_channels, *args, **kwargs)
| 702 | |
| 703 | class SimpleDecoder(nn.Module): |
| 704 | def __init__(self, in_channels, out_channels, *args, **kwargs): |
| 705 | super().__init__() |
| 706 | self.model = nn.ModuleList( |
| 707 | [ |
| 708 | nn.Conv2d(in_channels, in_channels, 1), |
| 709 | ResnetBlock( |
| 710 | in_channels=in_channels, |
| 711 | out_channels=2 * in_channels, |
| 712 | temb_channels=0, |
| 713 | dropout=0.0, |
| 714 | ), |
| 715 | ResnetBlock( |
| 716 | in_channels=2 * in_channels, |
| 717 | out_channels=4 * in_channels, |
| 718 | temb_channels=0, |
| 719 | dropout=0.0, |
| 720 | ), |
| 721 | ResnetBlock( |
| 722 | in_channels=4 * in_channels, |
| 723 | out_channels=2 * in_channels, |
| 724 | temb_channels=0, |
| 725 | dropout=0.0, |
| 726 | ), |
| 727 | nn.Conv2d(2 * in_channels, in_channels, 1), |
| 728 | Upsample(in_channels, with_conv=True), |
| 729 | ] |
| 730 | ) |
| 731 | # end |
| 732 | self.norm_out = Normalize(in_channels) |
| 733 | self.conv_out = torch.nn.Conv2d( |
| 734 | in_channels, out_channels, kernel_size=3, stride=1, padding=1 |
| 735 | ) |
| 736 | |
| 737 | def forward(self, x): |
| 738 | for i, layer in enumerate(self.model): |
nothing calls this directly
no test coverage detected