Pytorch Tensor MNIST Image Ploting

Viewed 184

I'm new to PyTorch and have been struggling a little bit with it. I'm trying to plot some MNIST dataset Images and I'm confused about the tensor indexing

Here's my code:

from torch.utils.data import DataLoader

train_dataloader = DataLoader(training_data,batch_size=100,shuffle=True)
test_dataloader = DataLoader(test_data,batch_size=100,shuffle=True)

examples = enumerate(test_dataloader)
batch_idx, (example_data, example_targets) = next(examples)

examples_idx = np.random.randint(0,high=len(example_data),size=25)
fig = plt.figure(figsize=(10,8))
rows, cols = 5,5

for i,j in enumerate(examples_idx):
  fig.add_subplot(rows,cols,i+1)
  plt.tight_layout()
  plt.imshow(example_data[j][0],cmap='gray')
  plt.title("Label: %g" %example_targets[j])
  plt.axis('off')
plt.show()

I was getting the following error when I tried to plot just example_data[j] with only a single indexing:

Invalid shape (1, 28, 28) for image data

I did some research and found out that apparently imshow expects a 2D array instead of a 3D one, since example_data[j] is a tensor with a [1,28,28] size. With that being said, I found two workarounds, I can either use np.squeeze(example_data[j)) or example_data[j][0] to make it work. However, I'm confused about the second one. What does this "0" stand for? Is it indexing the channel? If so, since example_data has shape [1,28,28], shouldn't I be indexing the 0 on the front like example_data[0,j] instead of doing it after the j index like before?

0 Answers
Related