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

如何用tf.data.Dataset基于非生成器函数分批加载大内存训练数据

问题与解决方案

核心需求

  • 现有非生成器类型的样本生成函数get_samples,一次性生成1000个样本会超出内存限制,需分10次调用(每次生成100个)
  • 基于TensorFlow的tf.data.Dataset实现:Dataset的第i个批次等价于get_samples(100, i)的输出
  • 每个epoch重复使用相同的随机种子参数调用函数,同时借助prefetch实现异步加载批次,避免单个样本初始化的高成本

示例代码

样本生成函数

import numpy as np

def get_samples(num_samples: int, random_seed=0):
    np.random.seed(random_seed)
    x = np.random.randint(0, 100, num_samples)
    y = np.random.randint(0, 2, num_samples)
    return np.array(list(zip(x, y)))

原问题逻辑(内存溢出风险)

batch_size = 100
total_num_samples = 1000
batches = []
for i in range(total_num_samples // batch_size):
    batches.append(get_samples(batch_size, i))

该方式会将所有批次数据一次性存入内存,不符合内存限制要求

解决方案实现

1. 定义批次级生成器

生成器每次仅返回一个批次的数据,对应指定的随机种子,不会提前加载所有批次:

def batch_generator(total_batches, batch_size):
    for seed in range(total_batches):
        yield get_samples(batch_size, seed)

2. 构建tf.data.Dataset

将批次生成器转换为TensorFlow Dataset,指定数据签名后添加prefetch实现异步加载:

import tensorflow as tf

total_batches = 10
batch_size = 100

# 根据实际输出定义数据的形状和类型
output_signature = tf.TensorSpec(shape=(batch_size, 2), dtype=tf.int64)

dataset = tf.data.Dataset.from_generator(
    lambda: batch_generator(total_batches, batch_size),
    output_signature=output_signature
)

# 自动适配异步加载,提升处理效率
dataset = dataset.prefetch(tf.data.AUTOTUNE)

3. 效果验证

遍历Dataset时,每个元素对应get_samples(100, i)的输出,且每个epoch重新遍历时,会用相同种子生成一致的批次数据:

# 第一个epoch遍历验证
for idx, batch in enumerate(dataset):
    print(f"批次{idx}的首个样本: {batch[0].numpy()}")

# 第二个epoch验证批次0的一致性
print("\n第二个epoch的批次0样本:")
for idx, batch in enumerate(dataset):
    if idx == 0:
        print(batch[0].numpy())
        break

关键细节

  • 批次级生成器确保每次仅加载单个批次数据,避免内存占用过高
  • from_generator按需调用生成器产生数据,不会提前缓存所有批次
  • prefetch让TensorFlow在处理当前批次时异步加载下一批次,优化训练流程
  • 每个epoch重新遍历Dataset时,生成器会从头执行,保证相同种子参数生成一致的批次数据

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 07:55:20