(self, config, dim: int, dilations: List[int])
| 233 | """ |
| 234 | |
| 235 | def __init__(self, config, dim: int, dilations: List[int]): |
| 236 | super().__init__() |
| 237 | kernel_sizes = (config.residual_kernel_size, 1) |
| 238 | if len(kernel_sizes) != len(dilations): |
| 239 | raise ValueError("Number of kernel sizes should match number of dilations") |
| 240 | |
| 241 | hidden = dim // config.compress |
| 242 | block = [] |
| 243 | for i, (kernel_size, dilation) in enumerate(zip(kernel_sizes, dilations)): |
| 244 | in_chs = dim if i == 0 else hidden |
| 245 | out_chs = dim if i == len(kernel_sizes) - 1 else hidden |
| 246 | block += [nn.ELU()] |
| 247 | block += [ |
| 248 | EncodecConv1d(config, in_chs, out_chs, kernel_size, dilation=dilation) |
| 249 | ] |
| 250 | self.block = block |
| 251 | |
| 252 | if getattr(config, "use_conv_shortcut", True): |
| 253 | self.shortcut = EncodecConv1d(config, dim, dim, kernel_size=1) |
| 254 | else: |
| 255 | self.shortcut = nn.Identity() |
| 256 | |
| 257 | def __call__(self, hidden_states): |
| 258 | residual = hidden_states |
nothing calls this directly
no test coverage detected