(self, c: Config)
| 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: |
| 361 | self.multiplier_proj = Conv1d1x1(self.middle_channels, self.middle_channels * self.multiplier, bias=c.bias) |
| 362 | |
| 363 | def forward(self, x, return_feats=False): |
| 364 | if self.c.latent_dim is not None and self.mode == 'decoder': |
nothing calls this directly
no test coverage detected