MCPcopy Create free account
hub / github.com/Francis-Rings/StableAnimator / forward

Method forward

animation/modules/unet_3d_blocks.py:619–639  ·  view source on GitHub ↗
(
        self,
        hidden_states: torch.Tensor,
        temb: Optional[torch.Tensor] = None,
        num_frames: int = 1,
    )

Source from the content-addressed store, hash-verified

617 self.gradient_checkpointing = False
618
619 def forward(
620 self,
621 hidden_states: torch.Tensor,
622 temb: Optional[torch.Tensor] = None,
623 num_frames: int = 1,
624 ) -> Union[torch.Tensor, Tuple[torch.Tensor, ...]]:
625 output_states = ()
626
627 for resnet, temp_conv in zip(self.resnets, self.temp_convs):
628 hidden_states = resnet(hidden_states, temb)
629 hidden_states = temp_conv(hidden_states, num_frames=num_frames)
630
631 output_states += (hidden_states,)
632
633 if self.downsamplers is not None:
634 for downsampler in self.downsamplers:
635 hidden_states = downsampler(hidden_states)
636
637 output_states += (hidden_states,)
638
639 return hidden_states, output_states
640
641
642class CrossAttnUpBlock3D(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected