TensorFlow 1.2解码TFRecord报错:MapDataset无prefetch属性求解决方案
首先,你的第一个错误AttributeError: 'MapDataset' object has no attribute 'prefetch'原因很明确:TensorFlow 1.2的Dataset API还没有把prefetch()作为Dataset对象的内置方法,这个方法是在TensorFlow 1.4版本才正式加入的,所以直接调用肯定会报错。而你训练时没遇到这个问题,大概率是训练代码里没有使用prefetch(),或者用了其他等价的队列机制来实现数据预取。
下面给你两个适配TF 1.2的解决方案:
方案1:使用tf.contrib.data.prefetch_to_device实现设备预取
如果你的目标是把数据预取到指定设备(比如GPU/CPU),可以用TF 1.2中contrib.data模块提供的prefetch_to_device方法,替代原来的prefetch():
# 假设你的解码函数是decode_tfrecord dataset = dataset.map(decode_tfrecord) # 替换 dataset.prefetch(batch_size) 为下面的代码 # 第一个参数是目标设备,第二个是预取缓冲区大小 dataset = dataset.apply(tf.contrib.data.prefetch_to_device("/cpu:0", buffer_size=batch_size))
如果需要预取到GPU,把设备路径改成"/gpu:0"即可。
方案2:结合batch操作调整缓冲区实现预取效果
如果你不需要指定设备,只是想提升数据读取的吞吐量,可以在batch操作时通过设置缓冲区参数来模拟预取的效果。TF 1.2中batch()方法的buffer_size参数可以控制内部队列的大小,相当于实现了数据预取:
dataset = dataset.map(decode_tfrecord) # 增大buffer_size来实现类似预取的效果,值可以设为batch_size的2-4倍 dataset = dataset.batch(batch_size, buffer_size=batch_size * 2)
如果你的测试代码中移除prefetch()后出现的新错误和数据读取速度慢、队列空有关,这个方案应该能缓解问题。
额外提示
如果你的训练代码用了旧的队列API(比如tf.train.shuffle_batch、tf.train.batch)而不是Dataset API,那测试代码也可以对齐这种方式:把解码后的张量传入tf.train.batch,通过capacity参数控制预取缓冲区,示例如下:
# 假设decode_tfrecord返回单个样本的张量 sample_tensor = decode_tfrecord(...) # 用旧队列API批量读取,capacity包含了预取的缓冲区 batch_tensor = tf.train.batch([sample_tensor], batch_size=batch_size, capacity=batch_size * 3)
这种方式和训练代码的逻辑一致,也能避免Dataset API版本兼容问题。
内容的提问来源于stack exchange,提问作者csbk

