TensorFlow:如何让Estimator调用的input_fn使用GPU
让TensorFlow r1.8 Estimator的input_fn运行在GPU上
我之前在TensorFlow r1.8版本里折腾过类似的需求,确实Estimator的_call_input_fn方法硬编码了with ops.device('/cpu:0'):的约束,强制让input_fn默认跑在CPU上。不过有几个可行的办法能绕开这个限制,实现input_fn在GPU上执行,给你参考:
方法1:在input_fn内部显式指定GPU设备(推荐)
这是最稳妥且无需修改框架源码的方式。TensorFlow的设备上下文支持嵌套覆盖,你只需要在input_fn的逻辑外层,手动用tf.device指定GPU设备,就能覆盖Estimator外层的CPU约束。
举个实际的代码例子:
def custom_input_fn(): # 显式指定使用第0块GPU,根据你的硬件情况修改device_index with tf.device('/gpu:0'): # 你的输入流水线逻辑:读取数据、预处理、生成batch等 raw_dataset = tf.data.TFRecordDataset("train.tfrecords") parsed_dataset = raw_dataset.map(parse_example_fn) batched_dataset = parsed_dataset.batch(64).shuffle(1000) iterator = batched_dataset.make_one_shot_iterator() features, labels = iterator.get_next() return features, labels
这样即使Estimator在调用input_fn时套了CPU的设备上下文,input_fn内部的所有操作都会优先分配到你指定的GPU上。
方法2:修改Estimator源码(快速但不推荐)
如果你只是在本地环境临时测试,不想改input_fn代码,可以直接修改TensorFlow的Estimator源码:
- 找到你Python环境中TensorFlow的安装路径,定位到
estimator.py文件(通常在site-packages/tensorflow/python/estimator/estimator.py) - 找到
_call_input_fn方法,删掉或者注释掉with ops.device('/cpu:0'):这一行,或者直接改成你想要的GPU设备(比如/gpu:0)
不过这个方法的弊端很明显:一旦你更新TensorFlow或者切换到其他环境,修改会丢失,而且会影响所有使用Estimator的代码,容易引发其他问题。
方法3:使用动态设备函数(适配多GPU场景)
如果你的环境有多个GPU,或者不确定GPU的编号,可以用动态设备函数来指定GPU:
def gpu_device_fn(op): # 强制分配到GPU,可根据需求修改device_index return tf.DeviceSpec(device_type='GPU', device_index=0) def custom_input_fn(): with tf.device(gpu_device_fn): # 输入流水线逻辑 ...
这种方式会让input_fn里的所有操作都优先分配到指定的GPU上,灵活性更高。
注意事项
- 确保你的GPU有足够的显存:输入流水线的预处理(比如图像解码、数据增强)也会占用显存,要避免和模型训练的显存需求冲突,必要时可以调整batch size或者预处理逻辑。
- 官方建议CPU跑输入流水线是为了让GPU专注于模型计算,但如果你的GPU有闲置资源,这么做完全没问题。
内容的提问来源于stack exchange,提问作者Vipin Pillai
相关产品推荐
相关产品推荐

