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

TensorFlow数据集加载引发内存泄漏,训练出现OOM错误求助

解决Xception训练数据集重采样导致的OOM问题

问题根源分析

你的代码中,load_rebalanced_dataset通过对同一个原始数据集多次执行filter来拆分不同类别的数据,TensorFlow会为每个filter分支保留独立的数据副本,内存占用随类别数量线性增长;再加上对无限重复的平衡数据集直接执行cache(),会持续缓存新生成的批次,最终导致内存耗尽。


具体解决办法

1. 避免重复过滤,直接按类别拆分原始数据集

不再基于同一个原始数据集多次filter,而是一次性遍历原始数据集,把每个类别的数据单独收集后创建专属数据集,从根源减少内存副本:

def load_rebalanced_dataset(
        height: int,
        width: int,
        path: str,
        kind: str,
        batch_size=32,
        do_repeat: bool = True) -> (tf.data.Dataset, list[str]):
    # 加载无batch的原始数据集,方便按单样本拆分
    raw_dataset = load_dataset(height, width, path, kind, batch_size=None)
    classes = raw_dataset.class_names
    num_classes = len(classes)
    
    # 按类别初始化数据容器
    class_data = [[] for _ in range(num_classes)]
    for x, y in raw_dataset:
        # 获取当前样本的类别索引
        class_idx = tf.argmax(y, axis=-1).numpy()
        class_data[class_idx].append(x)
    
    # 为每个类别创建独立数据集
    class_datasets = []
    for idx in range(num_classes):
        if not class_data[idx]:
            continue
        # 合并同类别样本并生成对应标签
        class_tensor = tf.concat(class_data[idx], axis=0)
        class_labels = tf.one_hot([idx] * len(class_data[idx]), num_classes)
        # 构建带shuffle和batch的类数据集
        class_ds = tf.data.Dataset.from_tensor_slices((class_tensor, class_labels))
        class_ds = class_ds.shuffle(len(class_data[idx])).batch(batch_size).cache()
        class_datasets.append(class_ds)
    
    # 采样生成平衡数据集
    balanced_ds = tf.data.Dataset.sample_from_datasets(class_datasets, [1.0/num_classes]*num_classes)
    
    if do_repeat:
        balanced_ds = balanced_ds.repeat()
    
    balanced_ds = balanced_ds.prefetch(tf.data.AUTOTUNE)
    
    return balanced_ds, classes

2. 优化缓存策略

  • 不要对无限重复的平衡数据集执行cache(),改为缓存每个类别的独立数据集(如上代码所示),避免持续累积缓存内容。
  • 如果内存仍紧张,可将缓存写入磁盘而非内存,把cache()改为cache(f"./cache_{kind}_{idx}"),指定磁盘缓存路径。

3. 降低单批次内存占用

  • 适当调小batch_size,减少每批数据的内存消耗。
  • 若图像尺寸非必须,缩小height和width参数,降低单张图片的内存占用。

4. 限制TensorFlow内存占用(GPU场景)

如果使用GPU训练,开启内存增长模式,避免TensorFlow一次性占用全部GPU内存:

physical_devices = tf.config.list_physical_devices('GPU')
if physical_devices:
    tf.config.experimental.set_memory_growth(physical_devices[0], True)

内容的提问来源于stack exchange,提问作者Marek M.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 13:52:20