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

Keras自定义Image DataGenerator极慢:列表转数组瓶颈求助

排查思路与优化建议

看起来你遇到的核心问题是真实加载的掩码列表转numpy数组时耗时过长,但同规格随机数组却很快,这说明开销并非来自np.array()本身,而是加载后的掩码对象的特性导致的。结合你的代码,我整理了几个关键排查方向和优化点:

1. 核心问题:掩码是Tensor而非Numpy数组

你的__load__函数最后将掩码转换为了tf.float32类型的Tensor:

mask = tf.cast(mask, tf.float32)

当你在__getitem__中把这些Tensor组成列表再转np.array()时,每个Tensor都需要先从GPU(如果Tensor在GPU内存上)拷贝到CPU,再执行.numpy()转换——这一步的跨设备数据传输是巨大的性能瓶颈!而你测试的随机数组本身就是CPU上的Numpy数组,自然转换极快。

解决办法:保持掩码为Numpy数组,避免不必要的Tensor转换:

def __load__(self, imgName, maskName):
    img = cv2.imread(os.path.join(self.imagePath,imgName))
    img = img/255.0
    mask = np.load(os.path.join(self.maskPath,maskName))
    mask = mask * self.weights  # 用Numpy乘法替代Tensor运算
    mask = mask.astype(np.float32)  # 转为float32,和TensorFlow默认类型一致
    return (img, mask)

2. 检查掩码形状的一致性

虽然你提到每个掩码都是224×224×4,但实际加载时可能存在形状不一致的情况(比如部分掩码文件损坏、保存时维度错误)。np.array()在处理形状不一致的列表时,会触发广播或创建更高维度的数组,导致耗时激增。

排查方式:在__load__中添加形状断言:

mask = np.load(os.path.join(self.maskPath,maskName))
assert mask.shape == (224, 224, 4), f"Invalid mask shape: {mask.shape} for {maskName}"

3. 避免修改实例的batchSize参数

你在__getitem__中修改了self.batchSize:

if(index+1)*self.batchSize > len(self.imgIds):
    self.batchSize = len(self.imgIds) - index*self.batchSize

这会破坏实例的初始配置,导致后续__len__计算出错(比如最后一轮之后,self.batchSize被改成了余数,下一次迭代时长度计算会错误)。应该用临时变量处理最后一批:

def __getitem__(self, index):
    start = index * self.batchSize
    end = min(start + self.batchSize, len(self.imgIds))
    batchImgs = self.imgIds[start:end]
    batchMasks = self.maskIds[start:end]
    batchfiles = [self.__load__(imgFile, maskFile) for imgFile, maskFile in zip(batchImgs, batchMasks)]
    images, masks = zip(*batchfiles)
    return np.array(images), np.array(masks)

4. 排查掩码文件的加载效率

如果掩码文件是用压缩格式保存的(比如np.savez),np.load的解压过程可能会增加开销。你可以测试直接加载单个掩码文件并计时,看是否是加载步骤本身耗时。如果是,可以考虑将掩码转换为未压缩的.npy格式,或者使用内存映射模式加载:

mask = np.load(os.path.join(self.maskPath,maskName), mmap_mode='r')

5. 批量加载优化(可选)

当前逐个加载图像和掩码的方式可以进一步优化,比如使用多进程/线程并行加载。不过在你的案例中,核心瓶颈应该是Tensor转Numpy的问题,先解决这个再考虑并行加载的优化。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 20:53:11