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

Method forward

src/models/resnet.py:51–83  ·  view source on GitHub ↗
(self, hidden_states, output_size=None)

Source from the content-addressed store, hash-verified

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
86class Downsample3D(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected