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.
相关产品推荐
相关产品推荐

