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

Python实现CNN时批量读取图像触发ValueError:无法将(28,28,3)数组广播为(28,28)

问题:CNN训练分批次读取图像时出现形状不匹配的ValueError

我在训练卷积神经网络(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 09:22:46