| 71 | |
| 72 | |
| 73 | class Upsample(nn.Module): |
| 74 | def __init__(self, channels, pad_type='repl', filt_size=4, stride=2): |
| 75 | super(Upsample, self).__init__() |
| 76 | self.filt_size = filt_size |
| 77 | self.filt_odd = np.mod(filt_size, 2) == 1 |
| 78 | self.pad_size = int((filt_size - 1) / 2) |
| 79 | self.stride = stride |
| 80 | self.off = int((self.stride - 1) / 2.) |
| 81 | self.channels = channels |
| 82 | |
| 83 | filt = get_filter(filt_size=self.filt_size) * (stride**2) |
| 84 | self.register_buffer('filt', filt[None, None, :, :].repeat((self.channels, 1, 1, 1))) |
| 85 | |
| 86 | self.pad = get_pad_layer(pad_type)([1, 1, 1, 1]) |
| 87 | |
| 88 | def forward(self, inp): |
| 89 | ret_val = F.conv_transpose2d(self.pad(inp), self.filt, stride=self.stride, padding=1 + self.pad_size, groups=inp.shape[1])[:, :, 1:, 1:] |
| 90 | if(self.filt_odd): |
| 91 | return ret_val |
| 92 | else: |
| 93 | return ret_val[:, :, :-1, :-1] |
| 94 | |
| 95 | |
| 96 | def get_pad_layer(pad_type): |