(self, hidden_states, output_size=None)
| 49 | self.conv = InflatedConv3d(self.channels, self.out_channels, 3, padding=1) |
| 50 | |
| 51 | def forward(self, hidden_states, output_size=None): |
| 52 | assert hidden_states.shape[1] == self.channels |
| 53 | |
| 54 | if self.use_conv_transpose: |
| 55 | raise NotImplementedError |
| 56 | |
| 57 | # Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16 |
| 58 | dtype = hidden_states.dtype |
| 59 | if dtype == torch.bfloat16: |
| 60 | hidden_states = hidden_states.to(torch.float32) |
| 61 | |
| 62 | # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 |
| 63 | if hidden_states.shape[0] >= 64: |
| 64 | hidden_states = hidden_states.contiguous() |
| 65 | |
| 66 | # if `output_size` is passed we force the interpolation output |
| 67 | # size and do not make use of `scale_factor=2` |
| 68 | if output_size is None: |
| 69 | hidden_states = F.interpolate( |
| 70 | hidden_states, scale_factor=[1.0, 2.0, 2.0], mode="nearest" |
| 71 | ) |
| 72 | else: |
| 73 | hidden_states = F.interpolate( |
| 74 | hidden_states, size=output_size, mode="nearest" |
| 75 | ) |
| 76 | |
| 77 | # If the input is bfloat16, we cast back to bfloat16 |
| 78 | if dtype == torch.bfloat16: |
| 79 | hidden_states = hidden_states.to(dtype) |
| 80 | |
| 81 | hidden_states = self.conv(hidden_states) |
| 82 | |
| 83 | return hidden_states |
| 84 | |
| 85 | |
| 86 | class Downsample3D(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected