(self, sample, timestep, encoder_hidden_states, class_labels=None, *args, cross_attention_kwargs: dict, **kwargs)
| 194 | return getattr(self.unet, name) |
| 195 | |
| 196 | def forward(self, sample, timestep, encoder_hidden_states, class_labels=None, *args, cross_attention_kwargs: dict, **kwargs): |
| 197 | cross_attention_kwargs = dict(cross_attention_kwargs) |
| 198 | control_depth = cross_attention_kwargs.pop('control_depth') |
| 199 | down_block_res_samples, mid_block_res_sample = self.controlnet( |
| 200 | sample, |
| 201 | timestep, |
| 202 | encoder_hidden_states=encoder_hidden_states, |
| 203 | controlnet_cond=control_depth, |
| 204 | conditioning_scale=self.conditioning_scale, |
| 205 | return_dict=False, |
| 206 | ) |
| 207 | return self.unet( |
| 208 | sample, |
| 209 | timestep, |
| 210 | encoder_hidden_states=encoder_hidden_states, |
| 211 | down_block_res_samples=down_block_res_samples, |
| 212 | mid_block_res_sample=mid_block_res_sample, |
| 213 | cross_attention_kwargs=cross_attention_kwargs |
| 214 | ) |
| 215 | |
| 216 | |
| 217 | class ModuleListDict(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected