| 84 | |
| 85 | |
| 86 | class UpResnetBlock1D(nn.Module): |
| 87 | def __init__( |
| 88 | self, |
| 89 | in_channels: int, |
| 90 | out_channels: int | None = None, |
| 91 | num_layers: int = 1, |
| 92 | temb_channels: int = 32, |
| 93 | groups: int = 32, |
| 94 | groups_out: int | None = None, |
| 95 | non_linearity: str | None = None, |
| 96 | time_embedding_norm: str = "default", |
| 97 | output_scale_factor: float = 1.0, |
| 98 | add_upsample: bool = True, |
| 99 | ): |
| 100 | super().__init__() |
| 101 | self.in_channels = in_channels |
| 102 | out_channels = in_channels if out_channels is None else out_channels |
| 103 | self.out_channels = out_channels |
| 104 | self.time_embedding_norm = time_embedding_norm |
| 105 | self.add_upsample = add_upsample |
| 106 | self.output_scale_factor = output_scale_factor |
| 107 | |
| 108 | if groups_out is None: |
| 109 | groups_out = groups |
| 110 | |
| 111 | # there will always be at least one resnet |
| 112 | resnets = [ResidualTemporalBlock1D(2 * in_channels, out_channels, embed_dim=temb_channels)] |
| 113 | |
| 114 | for _ in range(num_layers): |
| 115 | resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels)) |
| 116 | |
| 117 | self.resnets = nn.ModuleList(resnets) |
| 118 | |
| 119 | if non_linearity is None: |
| 120 | self.nonlinearity = None |
| 121 | else: |
| 122 | self.nonlinearity = get_activation(non_linearity) |
| 123 | |
| 124 | self.upsample = None |
| 125 | if add_upsample: |
| 126 | self.upsample = Upsample1D(out_channels, use_conv_transpose=True) |
| 127 | |
| 128 | def forward( |
| 129 | self, |
| 130 | hidden_states: torch.Tensor, |
| 131 | res_hidden_states_tuple: tuple[torch.Tensor, ...] | None = None, |
| 132 | temb: torch.Tensor | None = None, |
| 133 | ) -> torch.Tensor: |
| 134 | if res_hidden_states_tuple is not None: |
| 135 | res_hidden_states = res_hidden_states_tuple[-1] |
| 136 | hidden_states = torch.cat((hidden_states, res_hidden_states), dim=1) |
| 137 | |
| 138 | hidden_states = self.resnets[0](hidden_states, temb) |
| 139 | for resnet in self.resnets[1:]: |
| 140 | hidden_states = resnet(hidden_states, temb) |
| 141 | |
| 142 | if self.nonlinearity is not None: |
| 143 | hidden_states = self.nonlinearity(hidden_states) |
no outgoing calls
no test coverage detected
searching dependent graphs…