如何冻结绑定特定设备的Saved Model?解决部署时设备不匹配问题
我之前也碰到过一模一样的状况——多GPU训练导出的Saved Model带着硬编码的设备绑定,冻结的时候直接报设备找不到的错。咱们从问题根源到解决办法一步步来:
问题根源分析
你的报错信息很典型:
Cannot assign a device for operation NmtModel/transpose/Rank: Operation was explicitly assigned to /device:GPU:4 but available devices are [ /job:localhost/replica:0/task:0/device:CPU:0, /job:localhost/replica:0/task:0/device:GPU:0, /job:localhost/replica:0/task:0/device:GPU:1, /job:localhost/replica:0/task:0/device:XLA_CPU:0, /job:localhost/replica:0/task:0/device:XLA_GPU:0 ]. Make sure the device specification refers to a valid device.
本质原因是多GPU训练时,模型节点被硬绑定到了特定GPU(比如你的GPU:4),但导出Saved Model时没清除这些设备约束,导致后续加载/冻结时,TensorFlow硬要找不存在的设备,直接报错。
解决方案:分两种场景处理
场景1:已有带设备绑定的Saved Model,需要冻结
你原来的代码尝试修改inference_graph_def的节点设备,但问题出在加载模型阶段就因为设备绑定失败了。需要调整顺序,先加载模型,再强制清空所有节点的设备约束,最后再执行冻结:
import tensorflow as tf from tensorflow.python.saved_model import tag_constants from tensorflow.python.tools import freeze_graph import os import logging # 替换成你的路径 saved_model_dir = "your_saved_model_path" frozen_dir = "your_frozen_output_dir" with tf.Session(config=tf.ConfigProto( allow_soft_placement=True, # 允许自动分配到可用设备 log_device_placement=False )) as sess: sess.run(tf.tables_initializer()) # 加载Saved Model tf.saved_model.loader.load(sess, [tag_constants.SERVING], saved_model_dir) # 核心操作:清空图中所有节点的设备绑定 graph = tf.get_default_graph() for op in graph.get_operations(): op._set_device("") # 移除硬编码的设备约束 # 导出修改后的图定义 inference_graph_def = graph.as_graph_def() # 执行冻结 frozen_graph_path = os.path.join(frozen_dir, 'frozen_inference_graph.pb') output_keys = ['ToInt64', 'ToInt32', 'while/Exit_5'] output_node_names = ','.join(["%s/%s" % ('NmtModel', output_key) for output_key in output_keys]) _ = freeze_graph.freeze_graph( input_graph=inference_graph_def, input_saver=None, input_binary=True, input_saved_model_dir=saved_model_dir, input_checkpoint=None, output_node_names=output_node_names, restore_op_name=None, filename_tensor_name=None, output_graph=frozen_graph_path, clear_devices=True, # 双重保险:冻结时再清除一次设备约束 initializer_nodes='' ) logging.info("export frozen_inference_graph.pb success!!!")
场景2:还能重新导出Saved Model(从根源避免问题)
如果你的训练流程还能重新执行,建议在导出Saved Model时就清除设备绑定,这样后续冻结就不会踩坑:
# 导出Saved Model时的代码片段 builder = tf.saved_model.builder.SavedModelBuilder(saved_model_dir) # 假设你已经定义好signature_def_map等必要内容 builder.add_meta_graph_and_variables( sess, [tag_constants.SERVING], signature_def_map=signature_def_map, clear_devices=True # 关键参数:导出时自动清除所有设备绑定 ) builder.save()
关键说明
allow_soft_placement=True只是让TensorFlow在设备不可用时自动 fallback,但如果节点被硬绑定了设备,这个配置不会生效,必须手动清空设备约束。- 直接遍历图中的
Operation对象修改设备,比只修改graph_def的节点更彻底,因为graph_def是静态定义,而加载后的Operation才是会话中实际运行的对象。
内容的提问来源于stack exchange,提问作者duduscript

