TensorFlow Dataset API中dataset.batch(n).prefetch(m)预取m个批次还是样本?
关于TensorFlow Dataset中
prefetch(m)的预取单位问题 嘿,这个问题确实是刚接触TF Dataset API的同学常踩的小坑,我来给你掰扯清楚~
当你写dataset.batch(n).prefetch(m)的时候,prefetch(m)预取的是m个批次,不是m个样本。
具体来说:
- 首先
batch(n)会把原始数据集里的n个样本打包成一个独立的批次,这时候数据集的每个元素都是一个包含n个样本的批次对象。 - 紧接着
prefetch(m)是作用在这个已经分好批次的数据集上的,它会提前在后台准备好m个这样的批次。打个比方,如果n=32(每个批次32个样本),m=2,那预取的就是2个批次,总共64个样本,但预取的单位是批次,不是单个样本。
这么设计的原因也很实际:训练时模型是按批次接收数据的,按批次预取能最大化CPU和GPU的并行效率——CPU在后台忙着准备下几个批次的同时,GPU在专心训练当前批次,不会出现GPU等数据“摸鱼”的情况,能大幅提升训练速度。
内容的提问来源于stack exchange,提问作者lynn
相关产品推荐
相关产品推荐

