(self, in_features: List[torch.Tensor])
| 171 | self.res_blocks[i][j] = wrap_module_with_gradient_checkpointing(self.res_blocks[i][j]) |
| 172 | |
| 173 | def forward(self, in_features: List[torch.Tensor]): |
| 174 | out_features = [] |
| 175 | for i in range(len(self.res_blocks)): |
| 176 | feature = self.input_blocks[i](in_features[i]) |
| 177 | if i == 0: |
| 178 | x = feature |
| 179 | elif feature is not None: |
| 180 | x = x + feature |
| 181 | x = self.res_blocks[i](x) |
| 182 | out_features.append(self.output_blocks[i](x)) |
| 183 | if i < len(self.res_blocks) - 1: |
| 184 | x = self.resamplers[i](x) |
| 185 | return out_features |
nothing calls this directly
no outgoing calls
no test coverage detected