| 131 | |
| 132 | |
| 133 | class CausalConvTranspose1d(NonCausalConvTranspose1d): |
| 134 | def __init__( |
| 135 | self, |
| 136 | in_channels, |
| 137 | out_channels, |
| 138 | kernel_size, |
| 139 | stride, |
| 140 | bias=True, |
| 141 | pad_buffer=None, |
| 142 | ): |
| 143 | super(CausalConvTranspose1d, self).__init__( |
| 144 | in_channels=in_channels, |
| 145 | out_channels=out_channels, |
| 146 | kernel_size=kernel_size, |
| 147 | stride=stride, |
| 148 | padding=0, |
| 149 | output_padding=0, |
| 150 | bias=bias, |
| 151 | ) |
| 152 | self.stride = stride |
| 153 | self.pad_length = (math.ceil(kernel_size/stride) - 1) |
| 154 | if pad_buffer is None: |
| 155 | pad_buffer = T.zeros(1, in_channels, self.pad_length) |
| 156 | self.register_buffer("pad_buffer", pad_buffer) |
| 157 | |
| 158 | def forward(self, x): |
| 159 | pad = nn.ReplicationPad1d((self.pad_length, 0)) |
| 160 | x = pad(x) |
| 161 | return self.deconv(x)[:, :, self.stride : -self.stride] |
| 162 | |
| 163 | def inference(self, x): |
| 164 | x = T.cat((self.pad_buffer, x), -1) |
| 165 | self.pad_buffer = x[:, :, -self.pad_length:] |
| 166 | return self.deconv(x)[:, :, self.stride : -self.stride] |
| 167 | |
| 168 | def reset_buffer(self): |
| 169 | self.pad_buffer.zero_() |
| 170 | |
| 171 | |
| 172 | class NonCausalResUnit(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected