img = torch.randn(2, 28, 28, 3) # N, H, W, C (from PIL / OpenCV) x = img.permute(0, 3, 1, 2) # N, C, H, W (what nn.Conv2d wants) print(x.shape) # torch.Size([2, 3, 28, 28])