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

使用Estimator结合from_generator训练TensorFlow模型的问题

使用生成器为TensorFlow Estimator提供批量样本

嘿,我来帮你搞定这个问题!先梳理下你代码里的几个细节,再给你完整的实现方案。

首先,你的_generator()每次yield的其实已经是一个包含4个样本的批次了(feats是(4,2),labels是(4,1)),如果这时候再调用dataset.batch(4),会把4个这样的批次合并成更大的批次(最终feats形状变成(16,2)),这大概率不是你想要的“每次迭代输入一批样本”的效果。我分两种场景给你调整:


场景1:生成器返回单个样本,灵活控制批次大小

这种方式下,生成器每次只生成一个样本,后续用dataset.batch()来凑成你想要的批次大小,好处是你可以随时调整批次大小,不用修改生成器代码:

import numpy as np
import tensorflow as tf

def _generator():
    # 每次生成1个样本,特征形状(2,),标签形状(1,)
    for i in range(100):
        feats = np.random.rand(2)
        labels = np.random.rand(1)
        yield feats, labels

def input_func_gen():
    # 定义输出的类型和单个样本的形状
    output_types = (tf.float32, tf.float32)
    output_shapes = (tf.TensorShape([2]), tf.TensorShape([1]))
    
    # 从生成器创建Dataset
    dataset = tf.data.Dataset.from_generator(
        generator=_generator,
        output_types=output_types,
        output_shapes=output_shapes
    )
    
    # 设置批次大小为4,每次迭代返回4个样本
    dataset = dataset.batch(4)
    # 如果需要重复多轮训练,加上repeat(这里重复20轮)
    dataset = dataset.repeat(20)
    
    # 新版本TensorFlow(1.14+)的Estimator可以直接返回Dataset
    return dataset

如果你的TensorFlow版本比较旧,需要手动创建迭代器并返回特征字典和标签:

def input_func_gen():
    shapes = ((2,), (1,))
    dataset = tf.data.Dataset.from_generator(
        generator=_generator,
        output_types=(tf.float32, tf.float32),
        output_shapes=shapes
    )
    dataset = dataset.batch(4).repeat(20)
    
    # 创建一次性迭代器并获取下一批数据
    iterator = dataset.make_one_shot_iterator()
    features, labels = iterator.get_next()
    
    # Estimator需要的格式:特征字典(键名要和模型输入对应) + 标签张量
    return {'input_features': features}, labels

场景2:生成器直接返回批次数据

如果你确实需要生成器每次就返回一个完整的批次(比如从文件批量读取数据的场景),那直接去掉dataset.batch(4)即可:

def _generator():
    # 每次直接生成4个样本的批次
    for i in range(100):
        feats = np.random.rand(4,2)
        labels = np.random.rand(4,1)
        yield feats, labels

def input_func_gen():
    shapes = ((4,2), (4,1))
    dataset = tf.data.Dataset.from_generator(
        generator=_generator,
        output_types=(tf.float32, tf.float32),
        output_shapes=shapes
    )
    # 重复训练轮次
    dataset = dataset.repeat(20)
    
    # 新版本直接返回Dataset
    return dataset
    # 旧版本写法:
    # iterator = dataset.make_one_shot_iterator()
    # features, labels = iterator.get_next()
    # return {'input_features': features}, labels

最后提醒一下:要确保你的Estimator模型的输入层形状和批次后的特征形状匹配,比如批次后特征是(4,2),模型输入层要接受(None, 2)的张量(None表示批次大小可变)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:07:31