Fix F.conv3d call: remove invalid kernel_size kwarg (uses weight.shape)

This commit is contained in:
Daniel Maddern 2026-08-19 22:01:12 +07:00
parent 453aa86328
commit 6866263631

View file

@ -213,7 +213,8 @@ class _VisionPatchEmbed(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
target = self.weight.dtype
x = x.view(-1, 3, VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)
return F.conv3d(x.to(target), self.weight, self.bias, kernel_size=(VISION_TEMPORAL, VISION_PATCH, VISION_PATCH), stride=(VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)).view(-1, self.weight.shape[0])
s = (VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)
return F.conv3d(x.to(target), self.weight, self.bias, stride=s).view(-1, self.weight.shape[0])
class _VisionMLP(nn.Module):