ResNet encoder or decoder stack. Channel ratios and strides take the default order of from data/io-layer, to the middle of the model.
| 301 | |
| 302 | @si_module |
| 303 | class ResNetStack(nn.Module): |
| 304 | """ |
| 305 | ResNet encoder or decoder stack. Channel ratios |
| 306 | and strides take the default order of from |
| 307 | data/io-layer, to the middle of the model. |
| 308 | """ |
| 309 | class Config: |
| 310 | input_channels: int = 1 |
| 311 | output_channels: int = 1 |
| 312 | encode_channels: int = 32 |
| 313 | decode_channel_multiplier: int = 1 |
| 314 | latent_dim: int = None |
| 315 | kernel_size: int = 7 |
| 316 | bias: bool = True |
| 317 | channel_ratios: Tuple[int, ...] = (2, 4, 8, 16) |
| 318 | strides: Tuple[int, ...] = (3, 4, 5, 5) |
| 319 | mode: Literal['encoder', 'decoder'] = 'encoder' |
| 320 | |
| 321 | def __init__(self, c: Config): |
| 322 | super().__init__() |
| 323 | assert c.mode in ('encoder', 'decoder'), f"Mode ({c.mode}) is not supported!" |
| 324 | |
| 325 | self.mode = c.mode |
| 326 | |
| 327 | assert len(c.channel_ratios) == len(c.strides) |
| 328 | channel_ratios = (1,) + c.channel_ratios |
| 329 | strides = c.strides |
| 330 | self.middle_channels = c.encode_channels * channel_ratios[-1] |
| 331 | if c.mode == 'decoder': |
| 332 | channel_ratios = tuple(reversed(channel_ratios)) |
| 333 | strides = tuple(reversed(strides)) |
| 334 | |
| 335 | self.multiplier = c.decode_channel_multiplier if c.mode == 'decoder' else 1 |
| 336 | res_blocks = [ResNetBlock( |
| 337 | c.encode_channels * channel_ratios[s_idx] * self.multiplier, |
| 338 | c.encode_channels * channel_ratios[s_idx+1] * self.multiplier, |
| 339 | stride, |
| 340 | kernel_size=c.kernel_size, |
| 341 | bias=c.bias, |
| 342 | mode=c.mode, |
| 343 | ) for s_idx, stride in enumerate(strides)] |
| 344 | |
| 345 | data_conv = CausalConv1d( |
| 346 | in_channels=c.input_channels if c.mode == 'encoder' else c.encode_channels * self.multiplier, |
| 347 | out_channels=c.encode_channels if c.mode == 'encoder' else c.output_channels, |
| 348 | kernel_size=c.kernel_size, |
| 349 | stride=1, |
| 350 | bias=False, |
| 351 | ) |
| 352 | |
| 353 | if c.mode == 'encoder': |
| 354 | self.res_stack = nn.Sequential(data_conv, *res_blocks) |
| 355 | elif c.mode == 'decoder': |
| 356 | self.res_stack = nn.Sequential(*res_blocks, data_conv) |
| 357 | |
| 358 | if c.latent_dim is not None: |
| 359 | self.latent_proj = Conv1d1x1(self.middle_channels, c.latent_dim, bias=c.bias) if c.mode == 'encoder' else Conv1d1x1(c.latent_dim, self.middle_channels, bias=c.bias) |
| 360 | if self.multiplier != 1: |
nothing calls this directly
no outgoing calls
no test coverage detected