TensorFlow Dataset.repeat(count=None)如何工作?无限重复具体指什么?
关于TensorFlow Dataset.repeat()的常见疑问解答
一、默认(count=None/-1)的“无限重复”具体含义
当调用dataset.repeat()不指定count参数(或设为None/-1)时,这个数据集会循环遍历原始数据内容,永远不会终止。也就是说,当你迭代该数据集时,读完一遍原始数据后会立刻从头开始再次遍历,周而复始,没有天然的停止节点。
举个简单例子:如果原始数据集是[1,2,3],调用repeat()后,迭代输出的序列会是1,2,3,1,2,3,1,2,3,...,持续到你主动停止迭代(比如训练时设置固定epochs数、手动中断程序等)。
二、“默认重复次数等于batch运行次数”的说法是否正确?
这个说法不准确,甚至是错误的。
在代码train_univariate.cache().shuffle(BUFFER_SIZE).batch(BATCH_SIZE).repeat()中:
repeat()本身是无限重复整个数据集的逻辑,和batch的运行次数没有直接绑定关系。- 实际训练时,你是通过Keras的
model.fit()等方法指定epochs数来控制训练轮次,此时repeat()的无限性会被fit()的epochs限制——比如设epochs=5,数据集就会重复5次,而batch运行次数(总步数=样本数//batch_size * epochs)是另一个独立概念,和repeat()的默认行为无关。
所谓“等于batch运行次数”的说法,大概率是混淆了数据集迭代的总步数和重复轮次。实际场景中,若不限制epochs,数据集会一直重复直到训练停止;batch运行次数是每轮epoch的batch数量乘以重复轮次,由训练终止条件决定,并非repeat()的默认属性。
内容的提问来源于stack exchange,提问作者Ana G
相关产品推荐
相关产品推荐

