| 558 | |
| 559 | |
| 560 | class UpBlock1D(nn.Module): |
| 561 | def __init__(self, in_channels: int, out_channels: int, mid_channels: int | None = None): |
| 562 | super().__init__() |
| 563 | mid_channels = in_channels if mid_channels is None else mid_channels |
| 564 | |
| 565 | resnets = [ |
| 566 | ResConvBlock(2 * in_channels, mid_channels, mid_channels), |
| 567 | ResConvBlock(mid_channels, mid_channels, mid_channels), |
| 568 | ResConvBlock(mid_channels, mid_channels, out_channels), |
| 569 | ] |
| 570 | |
| 571 | self.resnets = nn.ModuleList(resnets) |
| 572 | self.up = Upsample1d(kernel="cubic") |
| 573 | |
| 574 | def forward( |
| 575 | self, |
| 576 | hidden_states: torch.Tensor, |
| 577 | res_hidden_states_tuple: tuple[torch.Tensor, ...], |
| 578 | temb: torch.Tensor | None = None, |
| 579 | ) -> torch.Tensor: |
| 580 | res_hidden_states = res_hidden_states_tuple[-1] |
| 581 | hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) |
| 582 | |
| 583 | for resnet in self.resnets: |
| 584 | hidden_states = resnet(hidden_states) |
| 585 | |
| 586 | hidden_states = self.up(hidden_states) |
| 587 | |
| 588 | return hidden_states |
| 589 | |
| 590 | |
| 591 | class UpBlock1DNoSkip(nn.Module): |
no outgoing calls
no test coverage detected
searching dependent graphs…