(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None, encoder_hidden_states=None, compute_motion=True,)
| 831 | self.gradient_checkpointing = False |
| 832 | |
| 833 | def forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None, encoder_hidden_states=None, compute_motion=True,): |
| 834 | for resnet, motion_module in zip(self.resnets, self.motion_modules): |
| 835 | # pop res hidden states |
| 836 | res_hidden_states = res_hidden_states_tuple[-1] |
| 837 | res_hidden_states_tuple = res_hidden_states_tuple[:-1] |
| 838 | hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) |
| 839 | |
| 840 | if self.training and self.gradient_checkpointing: |
| 841 | def create_custom_forward(module): |
| 842 | def custom_forward(*inputs): |
| 843 | return module(*inputs) |
| 844 | |
| 845 | return custom_forward |
| 846 | |
| 847 | hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(resnet), hidden_states, temb) |
| 848 | if motion_module is not None: |
| 849 | hidden_states = torch.utils.checkpoint.checkpoint(create_custom_forward(motion_module), hidden_states.requires_grad_(), temb, encoder_hidden_states) |
| 850 | else: |
| 851 | hidden_states = resnet(hidden_states, temb) |
| 852 | if motion_module is not None and compute_motion: |
| 853 | hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) |
| 854 | else: |
| 855 | hidden_states = hidden_states |
| 856 | |
| 857 | if self.upsamplers is not None: |
| 858 | for upsampler in self.upsamplers: |
| 859 | hidden_states = upsampler(hidden_states, upsample_size) |
| 860 | |
| 861 | return hidden_states |
nothing calls this directly
no outgoing calls
no test coverage detected