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

如何在模型训练中正确使用DataGenerator?遇TypeError问题求助

问题分析与解决方案

核心原因

你遇到的TypeError: expected str, bytes or os.PathLike object, not ndarray错误,是因为自定义的DataGenerator类大概率是按照读取文件路径的逻辑编写的,但你传入的是已经预处理好的numpy图像数组,导致内部代码(比如PIL.Image.open()或tf.keras.utils.load_img())试图把数组当成文件路径解析,触发类型错误。

修正方案:适配numpy数组的DataGenerator实现

使用Keras官方推荐的tf.keras.utils.Sequence基类编写生成器,直接处理传入的numpy数组,无需再读取文件。以下是适配目标检测任务的完整实现:

import numpy as np
from tensorflow.keras.utils import Sequence

class DataGenerator(Sequence):
    def __init__(self, images, targets, batch_size=32, shuffle=True):
        self.images = images  # 传入的预处理后numpy图像数组
        self.targets = targets  # 传入的(xmin,ymin,xmax,ymax)标签数组
        self.batch_size = batch_size
        self.shuffle = shuffle
        self.indexes = np.arange(len(self.images))
        self.on_epoch_end()

    def __len__(self):
        # 计算每个epoch的批次数
        return int(np.ceil(len(self.images) / self.batch_size))

    def __getitem__(self, index):
        # 生成单个批次的图像与标签
        batch_idx = self.indexes[index*self.batch_size : (index+1)*self.batch_size]
        batch_imgs = self.images[batch_idx]
        batch_targets = self.targets[batch_idx]
        return batch_imgs, batch_targets

    def on_epoch_end(self):
        # 每个epoch结束后打乱数据顺序
        if self.shuffle:
            np.random.shuffle(self.indexes)

正确训练代码(替代fit_generator)

fit_generator已被Keras弃用,直接使用fit()即可支持Sequence类型的生成器:

# 创建生成器(根据CPU内存调整batch_size)
training_generator = DataGenerator(train_images, train_targets, batch_size=16)
validation_generator = DataGenerator(val_images, val_targets, batch_size=16)

# 启动训练
resnet_model.fit(
    training_generator,
    epochs=4,
    validation_data=validation_generator,
    workers=4,  # 根据CPU核心数设置,比如4核就设4
    use_multiprocessing=True  # 开启多进程加速数据加载
)

无GPU环境优化建议

  • 数据增强:在__getitem__方法中添加在线数据增强(如随机翻转、缩放),提升模型泛化性且不占用额外内存,示例:
    from tensorflow.keras.layers import RandomFlip, RandomZoom
    
    data_augmentation = tf.keras.Sequential([
        RandomFlip("horizontal"),
        RandomZoom(0.1)
    ])
    
    # 在__getitem__中加入:
    batch_imgs = data_augmentation(batch_imgs, training=True)
    
  • 降低batch_size:无GPU时内存有限,建议将batch_size设为8-16,避免内存溢出
  • 预加载优化:如果内存允许,可提前对所有图像完成预处理并保存为numpy数组(你已经这么做了),减少生成器内的计算开销

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 16:57:33