| 1189 | return x |
| 1190 | |
| 1191 | class Downsample(nn.Module): |
| 1192 | def __init__( |
| 1193 | self, |
| 1194 | in_channels, |
| 1195 | stride, |
| 1196 | out_channels=-1, |
| 1197 | with_conv=False |
| 1198 | ): |
| 1199 | super().__init__() |
| 1200 | self.with_conv = with_conv |
| 1201 | self.stride = stride |
| 1202 | if(out_channels == -1): |
| 1203 | out_channels = in_channels |
| 1204 | if(self.with_conv == "CircularConv2D"): |
| 1205 | k, p = DOWNSAMPLE_STRIDE2KERNEL_DICT[stride], DOWNSAMPLE_STRIDE2PAD_DICT[stride] |
| 1206 | self.conv = CircularConv2D(in_channels, out_channels, kernel_size=k, stride=stride, padding=p) |
| 1207 | elif(self.with_conv == "Conv2D"): |
| 1208 | self.conv = nn.Conv2d( |
| 1209 | in_channels, |
| 1210 | out_channels, |
| 1211 | kernel_size=stride, |
| 1212 | stride=stride, |
| 1213 | padding=0 |
| 1214 | ) |
| 1215 | |
| 1216 | def forward(self, x): |
| 1217 | if self.with_conv: |
| 1218 | x = self.conv(x) |
| 1219 | else: |
| 1220 | x = torch.nn.functional.avg_pool2d(x, kernel_size=self.stride, stride=self.stride) |
| 1221 | return x |
| 1222 | |
| 1223 | class Upsample(nn.Module): |
| 1224 | def __init__( |