MCPcopy Create free account
hub / github.com/360CVGroup/FancyVideo / forward

Method forward

fancyvideo/models/unet_blocks.py:833–861  ·  view source on GitHub ↗
(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None, encoder_hidden_states=None, compute_motion=True,)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected