You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

CIFAR10数据集CNN训练报错:TypeError标量索引转换问题

CIFAR10数据集CNN训练中可视化图片的TypeError问题解决

在用Python训练基于CIFAR10数据集的卷积神经网络时,编写训练集图片展示代码时,plt.xlabel(class_names[y_train[i]])语句触发以下错误:

TypeError: only integer scalar arrays can be converted to a scalar index

代码参考了适用于Fashion-MNIST数据集的可行代码,仅修改了图像尺寸和颜色通道数;检查y_train的数据类型和部分值均为标量,但仍无法定位问题。相关代码及报错信息如下:

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras.datasets import cifar10

import numpy as np
import matplotlib.pyplot as plt

(X_train_full, y_train_full), (X_test, y_test) = cifar10.load_data()
# reshape dataset to the format suitable for CNN. The array has 4 dimensions: (number of images,
# width of each image (32), height of each image(32), number of color channels(in this case 3))
X_train_full = X_train_full.reshape((X_train_full.shape[0], 32, 32, 3))
X_test = X_test.reshape((X_test.shape[0], 32, 32, 3))

# split the full training data into the training set and the validation set
X_valid, X_train = X_train_full[:10000]/255.0, X_train_full[10000:]/255.0
y_valid, y_train = y_train_full[:10000], y_train_full[10000:]

print("Training set dimensions: ", X_train.shape) # training and test
print("Validate set dimensions: ", X_valid.shape) # validation
print("Training labels dimensions: ", len(y_train))
print("Validate labels dimensions: ", len(y_valid))

plt.figure()
plt.imshow(X_test[0])
plt.colorbar()
plt.grid(False)
plt.show()

class_names = ["Airplane", "Automobile", "Bird", "Cat", "Deer", "Dog", "Frog", "Horse", "Ship", "Truck"]

print(y_train.dtype)

print(y_train[:10])

plt.figure(figsize=(10,10))

for i in range(25):
    plt.subplot(5,5,i+1)
    plt.xticks([])
    plt.yticks([])
    plt.grid(False)
    plt.imshow(X_train[i], cmap=plt.cm.binary)
    plt.xlabel(class_names[y_train[i]])
plt.show()

报错信息:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In[40], line 9
      7     plt.grid(False)
      8     plt.imshow(X_train[i], cmap=plt.cm.binary)
----> 9     plt.xlabel(class_names[y_train[i]])
     10 plt.show()

TypeError: only integer scalar arrays can be converted to a scalar index

错误原因

问题出在CIFAR10数据集的标签格式差异:

  • Fashion-MNIST的load_data()返回的标签是一维数组(如形状为(60000,)),每个标签是单个整数标量
  • CIFAR10的load_data()返回的标签是二维数组(如y_train_full的形状为(50000,1)),y_train[i]取出来的是一个单元素数组(如array([3])),而非单个整数。列表class_names仅支持整数标量作为索引,传入数组就会触发错误。

解决方法

方法1:将标签转换为一维数组

在加载数据后,直接把二维标签数组转为一维,后续代码无需修改:

(X_train_full, y_train_full), (X_test, y_test) = cifar10.load_data()
# 新增:将标签转为一维数组
y_train_full = y_train_full.flatten()
y_test = y_test.flatten()

方法2:索引时提取标量值

如果不想修改标签数组的形状,可在索引class_names时取出数组中的标量值:
把plt.xlabel(class_names[y_train[i]])修改为以下任意一种:

# 方式A:通过索引提取标量
plt.xlabel(class_names[y_train[i][0]])
# 方式B:用item()方法提取标量
plt.xlabel(class_names[y_train[i].item()])

内容的提问来源于stack exchange,提问作者Faito Dayo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.08 06:16:00