| 228 | |
| 229 | |
| 230 | class ResNetBlock(nn.Module): |
| 231 | def __init__(self, |
| 232 | in_channels, |
| 233 | out_channels, |
| 234 | stride, |
| 235 | kernel_size=7, |
| 236 | dilations=(1, 3, 9), |
| 237 | bias=True, |
| 238 | mode='encoder', |
| 239 | ): |
| 240 | super().__init__() |
| 241 | assert mode in ('encoder', 'decoder'), f"Mode ({mode}) is not supported!" |
| 242 | |
| 243 | self.mode = mode |
| 244 | self.stride = stride |
| 245 | |
| 246 | ConvUnit = CausalConv1d if mode == 'encoder' else CausalConvTranspose1d |
| 247 | |
| 248 | res_channels = in_channels if mode == 'encoder' else out_channels |
| 249 | |
| 250 | res_units = [CausalResUnit( |
| 251 | res_channels, |
| 252 | res_channels, |
| 253 | kernel_size=kernel_size, |
| 254 | dilation=dilation, |
| 255 | ) for dilation in dilations] |
| 256 | |
| 257 | if in_channels == out_channels: |
| 258 | if mode == 'encoder': |
| 259 | self.pool = nn.AvgPool1d(kernel_size=stride, stride=stride) |
| 260 | if mode == 'decoder': |
| 261 | self.upsample = nn.Upsample(scale_factor=stride, mode='nearest') |
| 262 | conv_unit = nn.Conv1d( |
| 263 | in_channels=in_channels, |
| 264 | out_channels=out_channels, |
| 265 | kernel_size=1, |
| 266 | bias=bias, |
| 267 | ) if in_channels != out_channels else nn.Identity() |
| 268 | else: |
| 269 | conv_unit = ConvUnit( |
| 270 | in_channels=in_channels, |
| 271 | out_channels=out_channels, |
| 272 | kernel_size=(2 * stride), |
| 273 | stride=stride, |
| 274 | bias=bias, |
| 275 | ) |
| 276 | |
| 277 | if mode == 'encoder': |
| 278 | if in_channels == out_channels: |
| 279 | self.res_block = nn.Sequential(*res_units, self.pool, conv_unit) |
| 280 | else: |
| 281 | self.res_block = nn.Sequential(*res_units, conv_unit) |
| 282 | elif mode == 'decoder': |
| 283 | if in_channels == out_channels: |
| 284 | self.res_block = nn.Sequential(self.upsample, conv_unit, *res_units) |
| 285 | else: |
| 286 | self.res_block = nn.Sequential(conv_unit, *res_units) |
| 287 | |