如何基于TFRecordDataset计算model.fit的steps_per_epoch取值
TensorFlow中steps_per_epoch参数取值问题
问题场景
在使用如下方式训练TensorFlow模型时:
model.fit(..., steps_per_epoch=10000, ....)
需要基于给定数据集计算steps_per_epoch的正确取值,对应数据集构造代码如下:
dataset = tf.data.TFRecordDataset([filenames]) dataset = dataset.repeat(1) dataset = dataset.batch(512) total = 0 for i in dataset: total += 1 print("Total is {}".format(total))
上述代码运行后输出结果为:
Total is 393
待确认的疑问:
- 此时
steps_per_epoch的取值是否等于393? - 还是应当按照
steps_per_epoch = 393 / 512的方式计算该参数值?
结论
steps_per_epoch的正确取值就是393,不需要除以512,393/512的计算方式是完全错误的。
原因说明
首先明确两个基础定义:
dataset.batch(512)里传入的512是batch size,也就是单步训练时喂给模型的样本数量steps_per_epoch的实际含义是:跑完一个完整epoch,模型需要执行的梯度下降步数,每一步会消耗1个batch的训练数据
你写的遍历计数逻辑,统计的是数据集按batch size=512切分完成后,总共能产出的batch个数:前392个batch每个包含512条完整样本,最后1个batch是不足512条的剩余样本,393个batch加起来正好覆盖全部训练数据,和steps_per_epoch要求的定义完全匹配。
容易搞混的计算逻辑边界:
- 如果你统计的是未做batch切分的原始样本总条数,才需要用「总样本数 / batch size 向上取整」的方式计算总步数
- 你现在统计的对象已经是batch处理后的数据集,计数结果本身就是总步数,不需要再做除法运算。
额外提示:你当前数据集设置了
repeat(1),也就是全量数据只遍历1次不会自动重复,这种场景下甚至可以不手动传入steps_per_epoch参数,TensorFlow会自动遍历完所有batch判定一个epoch结束,自动识别总步数为393。只有当数据集设置为无限重复(比如repeat()不传参数)时,才必须手动指定steps_per_epoch,告诉模型一个epoch跑多少步后停止进入下一个轮次。
内容的提问来源于stack exchange,提问作者stackbiz
相关产品推荐
相关产品推荐

