fix torch UserWarning

This commit is contained in:
BobRandomNumber
2025-08-13 00:59:03 -04:00
committed by GitHub
parent ae26ee5a1f
commit a0d17a1438
@@ -977,20 +977,21 @@ class WaveletTransform3D(torch.nn.Module):
# Flatten the tensor along the specified dimension for Conv1d
if dim == 0: # frames
b, c, f, h, w = x.shape
x = x.permute(0, 1, 3, 4, 2).reshape(b * c * h * w, f) # Flatten (b, c, h, w) -> (B, frames)
x = F.conv1d(x.unsqueeze(1), filter, padding=0).squeeze(1) # Apply Conv1d
x = x.view(b, c, h, w, -1).permute(0, 1, 4, 2, 3) # Unflatten back to original shape
x_flat = x.permute(0, 1, 3, 4, 2).reshape(b * c * h * w, f) # Flatten (b, c, h, w) -> (B, frames)
x_conv = F.conv1d(x_flat.unsqueeze(1), filter, padding=0, stride=1).squeeze(1) # Apply Conv1d
x = x_conv.view(b, c, h, w, -1).permute(0, 1, 4, 2, 3) # Unflatten back to original shape
elif dim == 1: # height
b, c, f, h, w = x.shape
x = x.permute(0, 1, 2, 4, 3).reshape(b * c * f * w, h) # Flatten (b, c, f, w) -> (B, height)
x = F.conv1d(x.unsqueeze(1), filter, padding="same",stride = 1).squeeze(1) # Apply Conv1d
x = x.view(b, c, f, w, -1).permute(0, 1, 2, 4, 3) # Unflatten back to original shape
x_flat = x.permute(0, 1, 2, 4, 3).reshape(b * c * f * w, h) # Flatten (b, c, f, w) -> (B, height)
x_padded = F.pad(x_flat.unsqueeze(1), (1, 0), mode='replicate')
x_conv = F.conv1d(x_padded, filter, padding=0, stride=1).squeeze(1)
x = x_conv.view(b, c, f, w, -1).permute(0, 1, 2, 4, 3) # Unflatten back to original shape
elif dim == 2: # width
b, c, f, h, w = x.shape
x = x.permute(0, 1, 2, 3, 4).reshape(b * c * f * h, w) # Flatten (b, c, f, h) -> (B, width)
x = F.conv1d(x.unsqueeze(1), filter, padding="same",stride = 1).squeeze(1) # Apply Conv1d
x = x.view(b, c, f, h, -1) # Unflatten back to original shape
x_flat = x.permute(0, 1, 2, 3, 4).reshape(b * c * f * h, w) # Flatten (b, c, f, h) -> (B, width)
x_padded = F.pad(x_flat.unsqueeze(1), (1, 0), mode='replicate')
x_conv = F.conv1d(x_padded, filter, padding=0, stride=1).squeeze(1)
x = x_conv.view(b, c, f, h, -1) # Unflatten back to original shape
return x
def forward(self, x):
@@ -1307,7 +1308,7 @@ class FlowDecoderWithResidual(nn.Module):
curr_res = resolution // 2**(self.num_resolutions-1) # 64
self.z_shape = (1,z_channels,curr_res,curr_res) # (1,3,64,64)
#print("Working with z of shape {} = {} dimensions.".format(
self.z_shape, np.prod(self.z_shape)))
# self.z_shape, np.prod(self.z_shape)))
# z to block_in
self.conv_in = torch.nn.Conv2d(z_channels,
@@ -1514,8 +1515,8 @@ class FlowDecoderWithResidual(nn.Module):
block_in = int(ch*ch_mult[self.num_resolutions-1]) # 512
curr_res = resolution // 2**(self.num_resolutions-1) # 64
self.z_shape = (1,z_channels,curr_res,curr_res) # (1,3,64,64)
print("Working with z of shape {} = {} dimensions.".format(
self.z_shape, np.prod(self.z_shape)))
#print("Working with z of shape {} = {} dimensions.".format(
# self.z_shape, np.prod(self.z_shape)))
# z to block_in
self.conv_in = torch.nn.Conv3d(z_channels,
@@ -1744,5 +1745,4 @@ class FlowDecoderWithResidual(nn.Module):
imgs = torch.cat((img0, img1), 1)
flow, flow_list = self.flownet(imgs)
flow, c0, c1 = self.refinenet(feats,flow)
return flow
return flow