【问题标题】:How can I solve the error: TypeError: Invalid shape (60, 60, 8) for image data如何解决错误:TypeError: Invalid shape (60, 60, 8) for image data
【发布时间】:2022-10-02 23:38:12
【问题描述】:

我是 pytorch 的新手。我正在尝试创建一个 DCGAN 项目。我使用了整个官方 pytorch tutorial 作为基础。

我有一个 numpy 数组,它是八个数组的组合,它给出了一个形状 (60,60,8) 这个形状很特别

lista2 = [0, 60, 120, 180, 240, 300, 360, 420]
total = []
for i in lista2:
   N1 = intesity[0:60, i:i+60]
   total.append(N1)
   N2 = intesity[60:120, i:i+60]
   total.append(N2)
   N3 = intesity[120:180, i:i+60]
   total.append(N3)
   N4 = intesity[180:240, i:i+60]
   total.append(N4)
   N5 = intesity[240:300, i:i+60]
   total.append(N5)
   N6 = intesity[300:360, i:i+60]
   total.append(N6)
   N7 = intesity[360:420, i:i+60]
   total.append(N7)
   N8 = intesity[420:480, i:i+60]
   total.append(N8)

total = np.reshape(total, (64, 60,60,8))
total  -= total.min()
total  /= total.max()
total = np.asarray(total)
print(np.shape(total)
(64, 60, 60, 8)

如您所见,该数组中有 64 个元素,有 64 个训练图像(目前很少),该数组先转换为张量,然后再转换为 pytorch 数据集

tensor_c = torch.tensor(total)

创建数据集和数据加载器在尝试绘制此 DCGAN 的训练图像时出现以下错误

dataset = TensorDataset(tensor_c) # create your datset
dataloader = DataLoader(dataset) # create your dataloader

real_batch = next(iter(dataloader))
plt.figure(figsize=(16,16))
plt.axis(\"off\")
plt.title(\"Training Images\")
plt.imshow(np.transpose(vutils.make_grid(real_batch[0].to(device)[:64], padding=0, normalize=True).cpu(),(1,2,0)))
dataset_size = len(dataloader.dataset)
dataset_size
---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
<ipython-input-42-5ba2d666ef25> in <module>()
     10 plt.axis(\"off\")
     11 plt.title(\"Training Images\")
---> 12 plt.imshow(np.transpose(vutils.make_grid(real_batch[0].to(device)[:64], padding=0, normalize=True).cpu(),(1,2,0)))
     13 dataset_size = len(dataloader.dataset)
     14 dataset_size

5 frames
/usr/local/lib/python3.7/dist-packages/matplotlib/image.py in set_data(self, A)
    697                 or self._A.ndim == 3 and self._A.shape[-1] in [3, 4]):
    698             raise TypeError(\"Invalid shape {} for image data\"
--> 699                             .format(self._A.shape))
    700 
    701         if self._A.ndim == 3:

TypeError: Invalid shape (60, 60, 8) for image data

我对 Pytorch 太陌生了,我想知道如何解决这个问题

    标签: pytorch dataset shapes generative-adversarial-network dataloader


    【解决方案1】:

    通常,图像应存储为height x width x n_channels 形式的数组,其中n_channels 对于标准RGB 图像是3,或者在某些情况下对于RGBA 图像是4。 matplotlib 对如何绘制具有 8 个通道的图像没有内置的理解,就像您的图像数据当前所具有的那样。

    还要注意维度的顺序,因为pytorch 需要batch_idx x channel x height x width 形式的图像,这便于应用 2D 卷积,因为它们可以跨越最后 2 个维度。在尝试以matplotlib 形式绘制图像后,请小心转换为pytorch 形式。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2020-10-16
      • 1970-01-01
      • 2016-12-07
      • 2020-09-14
      • 2017-08-18
      • 1970-01-01
      • 2022-01-22
      • 2017-06-01
      相关资源
      最近更新 更多