| 589 | |
| 590 | |
| 591 | class UpBlock1DNoSkip(nn.Module): |
| 592 | def __init__(self, in_channels: int, out_channels: int, mid_channels: int | None = None): |
| 593 | super().__init__() |
| 594 | mid_channels = in_channels if mid_channels is None else mid_channels |
| 595 | |
| 596 | resnets = [ |
| 597 | ResConvBlock(2 * in_channels, mid_channels, mid_channels), |
| 598 | ResConvBlock(mid_channels, mid_channels, mid_channels), |
| 599 | ResConvBlock(mid_channels, mid_channels, out_channels, is_last=True), |
| 600 | ] |
| 601 | |
| 602 | self.resnets = nn.ModuleList(resnets) |
| 603 | |
| 604 | def forward( |
| 605 | self, |
| 606 | hidden_states: torch.Tensor, |
| 607 | res_hidden_states_tuple: tuple[torch.Tensor, ...], |
| 608 | temb: torch.Tensor | None = None, |
| 609 | ) -> torch.Tensor: |
| 610 | res_hidden_states = res_hidden_states_tuple[-1] |
| 611 | hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) |
| 612 | |
| 613 | for resnet in self.resnets: |
| 614 | hidden_states = resnet(hidden_states) |
| 615 | |
| 616 | return hidden_states |
| 617 | |
| 618 | |
| 619 | DownBlockType = DownResnetBlock1D | DownBlock1D | AttnDownBlock1D | DownBlock1DNoSkip |
no outgoing calls
no test coverage detected
searching dependent graphs…