(
self,
hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
num_frames: int = 1,
)
| 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 | |
| 642 | class CrossAttnUpBlock3D(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected