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

Python生成器接入Keras model.fit_generator报错求助及实现咨询

解决Keras fit_generator输入维度错误及生成器实现合理性分析

一、先搞定输入维度错误问题

你碰到的ValueError核心原因很直接:Keras的卷积层要求输入是4维张量,格式为(batch_size, height, width, channels),但你的生成器每次只返回单张3维的图片(28,28,1),少了关键的批量维度。

这里有两种解决思路,推荐第二种更贴合训练逻辑的方式:

  1. 快速修复:给单张图加批量维度
    在生成器里用np.expand_dims给图片和标签扩展维度,修改你的gen方法:

    import numpy as np  # 记得导入numpy
    
    def gen(self, feat, labels):
        i=0
        total_samples = len(feat)
        while True:
            # 循环到末尾重置索引,避免越界报错
            if i >= total_samples:
                i = 0
            im = cv2.imread(feat[i],0)
            im = im.reshape(28,28,1)
            # 增加批量维度,变成(1,28,28,1)
            im = np.expand_dims(im, axis=0)
            label = np.expand_dims(labels[i], axis=0)
            yield im, label
            i+=1
    
  2. 更优方案:生成批量样本
    每次返回一个批次的图片和标签(比如常用的batch_size=32),这才是Keras训练的常规操作:

    def gen(self, feat, labels, batch_size=32):
        total_samples = len(feat)
        while True:
            # 随机打乱样本顺序,提升模型泛化能力
            indices = np.random.permutation(total_samples)
            for start in range(0, total_samples, batch_size):
                end = min(start + batch_size, total_samples)
                batch_indices = indices[start:end]
                batch_imgs = []
                batch_labels = []
                for idx in batch_indices:
                    im = cv2.imread(feat[idx], 0)
                    if im is None:
                        # 处理图片读取失败的情况
                        print(f"警告:无法读取图片 {feat[idx]}")
                        continue
                    im = im.reshape(28,28,1)
                    batch_imgs.append(im)
                    batch_labels.append(labels[idx])
            yield np.array(batch_imgs), np.array(batch_labels)
    

    同时要修改fit_generator的steps_per_epoch,设置为总样本数//batch_size(如果有余数就加1,确保所有样本都被训练到)。


二、你的生成器实现合理吗?

实话实说,原实现有不少可以优化的点:

  • 类封装冗余:完全没必要专门写Generator类,直接用生成器函数更简洁直观
  • 无边界处理:原代码里i会一直递增,超出feat长度后会直接抛出IndexError
  • 缺少数据打乱:固定顺序训练样本,会影响模型的泛化能力
  • 过时API使用:pd.get_dummies(...).as_matrix()已经被废弃,应该改用.values或者.to_numpy()
  • 无错误处理:没考虑图片读取失败的情况(比如文件损坏、路径错误)

三、优化后的完整代码

from keras.utils import to_categorical
from keras.models import Sequential
from keras.layers import Dense, Conv2D, Flatten
import pandas as pd
import os
import cv2
import numpy as np

def data_generator(feat, labels, batch_size=32):
    total_samples = len(feat)
    while True:
        # 随机打乱样本索引
        indices = np.random.permutation(total_samples)
        for start in range(0, total_samples, batch_size):
            end = min(start + batch_size, total_samples)
            batch_indices = indices[start:end]
            batch_imgs = []
            batch_labels = []
            for idx in batch_indices:
                img_path = feat[idx]
                im = cv2.imread(img_path, 0)
                if im is None:
                    print(f"Warning: Failed to read image {img_path}")
                    continue
                # 强制调整图片尺寸为28x28,避免尺寸不一致问题
                im = cv2.resize(im, (28,28))
                im = im.reshape(28,28,1)
                # 这里可以添加数据增强操作,比如随机翻转、旋转等
                batch_imgs.append(im)
                batch_labels.append(labels[idx])
            yield np.array(batch_imgs), np.array(batch_labels)

if __name__ == "__main__":
    input_dir = './mnist'
    output_file = 'dataset.csv'
    filename = []
    label = []
    for root,dirs,files in os.walk(input_dir):
        for file in files:
            full_path = os.path.join(root,file)
            filename.append(full_path)
            label.append(os.path.basename(os.path.dirname(full_path)))
    data = pd.DataFrame(data={'filename': filename, 'label':label})
    data.to_csv(output_file,index=False)
    
    feat = data['filename'].values
    # 替换过时的as_matrix()
    labels = pd.get_dummies(data['label']).to_numpy()
    
    batch_size = 32
    steps_per_epoch = len(feat) // batch_size
    # 处理样本数不能被batch_size整除的情况
    if len(feat) % batch_size != 0:
        steps_per_epoch +=1
    
    # 创建模型
    model = Sequential()
    model.add(Conv2D(64, kernel_size=3, activation="relu", input_shape=(28,28,1)))
    model.add(Conv2D(32, kernel_size=3, activation="relu"))
    model.add(Flatten())
    model.add(Dense(2, activation="softmax"))
    model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
    
    # 使用生成器训练
    model.fit_generator(data_generator(feat, labels, batch_size), 
                        steps_per_epoch=steps_per_epoch,
                        epochs=5, 
                        verbose=1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 23:13:11