Python实现CNN时批量读取图像触发ValueError:无法将(28,28,3)数组广播为(28,28)
我在训练卷积神经网络(CNN)时,为了避免内存不足,编写了readImages函数分批次读取图像和标签,而非一次性加载所有数据。函数代码如下:
def readImages(strSet = 'Train', nIni = 1, nFin = 20): if strSet not in ('Train','Test'): return None # 初始化输出数组:图像和标签 arrImages = [] arrLabels = [] # 遍历所选数据集下的所有目录 for strDir in os.listdir(data_dir+'/' + strSet + '/'): # 获取当前处理的类别名称 strClass = strDir[strDir.find('-')+1:] # 获取目录下的文件名列表和文件数量 arrNameFiles = os.listdir(data_dir+'/' + strSet + '/'+strDir) nFiles = len(os.listdir(data_dir+'/' + strSet + '/'+strDir)) # 根据nIni和nFin选择要读取的文件 if (nIni == -1): # 如果nIni为-1,读取目录下所有图像 listChosenFiles = arrNameFiles else: if (nImagesClase(strSet, strClass)<nFin): # 如果类别下的图像总数小于nFin,随机采样 listChosenFiles = random.sample(arrNameFiles, min(nFiles, nFin-nIni)) else: # 否则读取nIni到nFin范围内的文件 listChosenFiles = arrNameFiles[nIni-1:min(nFin,nImagesClase(strSet, strClass))-1] # 遍历选中的文件,读取并处理图像 for file in listChosenFiles: # 读取图像文件 image = plt.imread(data_dir+'/'+strSet+'/'+strDir+'/'+file) # 调整图像尺寸 image = cv2.resize(image, (crop_width, crop_height), interpolation=cv2.INTER_NEAREST) arrImages.append(image) # 创建并添加标签 arrLabel = np.zeros(n_classes) arrLabel[array_classes.index(strClass)] = 1 arrLabels.append(arrLabel) # 将列表转换为numpy数组 y = np.array(arrLabels) X = np.array(arrImages, dtype=np.uint8) # 打乱数据顺序 arrIndexes = np.arange(X.shape[0]) np.random.shuffle(arrIndexes) X = X[arrIndexes] y = y[arrIndexes] # 返回处理好的数据 return X, y
执行测试代码X, y = readImages(strSet = 'Train', nIni = 1, nFin = 5)时运行正常,但设置nIni = 101, nFin = 105时触发如下错误:
ValueError Traceback (most recent call last) <ipython-input-125-8a690256a1fc> in <module> ----> 1 X, y = readImages(strSet = 'Train', nIni = 101, nFin = 105) <ipython-input-123-9e9ebc660c33> in readImages(strSet, nIni, nFin) 50 # Recogemos los valores de entrada y salida en arrays. 51 y = np.array(arrLabels) ---> 52 X = np.array(arrImages, dtype=np.uint8) 53 # Una vez terminado el recorrido por todas las imágenes, reordenamos los índices para que no vayan las imágenes en secuendias de la misma clase. 54 arrIndexes = np.arange(X.shape[0]) ValueError: could not broadcast input array from shape (28,28,3) into shape (28,28)
我在图像读取环节添加了打印语句,确认所有读取的图像形状均为(28,28,3),无法理解错误提示中(28,28)形状的来源。请问该问题可能由什么原因导致?是否有人遇到过类似问题?
这个问题我之前也碰到过,核心原因是:你的数据集里存在单通道灰度图像,但你误以为所有图像都是3通道RGB图。小批次(比如1-5)刚好没抽到这类灰度图,所以运行正常;而101-105的批次刚好包含了这类图像,才触发了形状不匹配的错误。
至于你打印时没发现,大概率是打印范围有限,或者在cv2.resize之前的打印漏掉了那些灰度图——plt.imread读取灰度图时会返回(H,W)的形状,经过cv2.resize处理后依然是(28,28),不会自动扩展到3通道。
给你几个具体的解决办法:
1. 强制统一图像通道数
在cv2.resize之后添加代码,把单通道灰度图转换成3通道:
# 调整尺寸后检查通道数,统一转为3通道 if len(image.shape) == 2: # 方法1:用OpenCV转换 image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) # 方法2:用numpy扩展维度 # image = np.expand_dims(image, axis=-1) # image = np.repeat(image, 3, axis=-1)
2. 排查并预处理数据集
写个小脚本遍历所有图像,找出通道数异常的文件,手动处理(转3通道或删除):
import os import cv2 data_dir = "你的数据集路径" strSet = "Train" for strDir in os.listdir(os.path.join(data_dir, strSet)): dir_path = os.path.join(data_dir, strSet, strDir) for file in os.listdir(dir_path): file_path = os.path.join(dir_path, file) img = cv2.imread(file_path) if img is None: print(f"无法读取文件: {file_path}") continue if len(img.shape) != 3 or img.shape[2] != 3: print(f"通道数异常的文件: {file_path},形状: {img.shape}")
3. 优化批量读取逻辑
你代码里还有个潜在问题:nImagesClase函数的返回值如果和实际目录下的文件数不一致,会导致切片错误。建议直接用len(arrNameFiles)代替nImagesClase的调用,避免依赖外部函数的准确性。
另外,推荐用TensorFlow的tf.data.Dataset或者PyTorch的DataLoader来处理批量读取,这些内置工具已经帮你处理了很多数据一致性问题,比自己写函数更可靠高效。
内容的提问来源于stack exchange,提问作者Juan Diego

