| 373 | |
| 374 | |
| 375 | class ResConvBlock(nn.Module): |
| 376 | def __init__(self, in_channels: int, mid_channels: int, out_channels: int, is_last: bool = False): |
| 377 | super().__init__() |
| 378 | self.is_last = is_last |
| 379 | self.has_conv_skip = in_channels != out_channels |
| 380 | |
| 381 | if self.has_conv_skip: |
| 382 | self.conv_skip = nn.Conv1d(in_channels, out_channels, 1, bias=False) |
| 383 | |
| 384 | self.conv_1 = nn.Conv1d(in_channels, mid_channels, 5, padding=2) |
| 385 | self.group_norm_1 = nn.GroupNorm(1, mid_channels) |
| 386 | self.gelu_1 = nn.GELU() |
| 387 | self.conv_2 = nn.Conv1d(mid_channels, out_channels, 5, padding=2) |
| 388 | |
| 389 | if not self.is_last: |
| 390 | self.group_norm_2 = nn.GroupNorm(1, out_channels) |
| 391 | self.gelu_2 = nn.GELU() |
| 392 | |
| 393 | def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| 394 | residual = self.conv_skip(hidden_states) if self.has_conv_skip else hidden_states |
| 395 | |
| 396 | hidden_states = self.conv_1(hidden_states) |
| 397 | hidden_states = self.group_norm_1(hidden_states) |
| 398 | hidden_states = self.gelu_1(hidden_states) |
| 399 | hidden_states = self.conv_2(hidden_states) |
| 400 | |
| 401 | if not self.is_last: |
| 402 | hidden_states = self.group_norm_2(hidden_states) |
| 403 | hidden_states = self.gelu_2(hidden_states) |
| 404 | |
| 405 | output = hidden_states + residual |
| 406 | return output |
| 407 | |
| 408 | |
| 409 | class UNetMidBlock1D(nn.Module): |