layer = torch.nn.Conv2d(3, 768, kernel_size=(16, 16), stride=(16, 16)) image = torch.rand(batch_size, 3, 48, 48) projected_patches = layer(image) print(projected_patches.flatten(-2).transpose(-1, -2).shape) # This prints # torch.Size([1, 9, 768])