TensorFlow 1.15.0训练图中移除PyFunc节点并导出SavedModel方法
解决TensorFlow 1.15中PyFunc节点导出问题的方法
方法一:用原生TensorFlow操作重写自定义逻辑
虽然你提到原生无对应操作,但可以尝试用现有基础op组合实现相同功能:
- 自定义的数值转换、特征处理逻辑,多数能通过
tf.math、tf.strings、tf.reshape等基础op拼接实现 - 复杂条件分支可改用
tf.cond、tf.where替代Python的if/else;循环逻辑用tf.while_loop实现 - 重写后直接替换训练代码中的
tf.py_func节点,训练完成后直接导出SavedModel即可,无需额外处理
方法二:训练后替换PyFunc节点为原生op子图
如果训练时必须保留PyFunc(比如调试便利),可在导出阶段手动替换节点:
- 加载训练好的图:
import tensorflow as tf sess = tf.compat.v1.Session() saver = tf.compat.v1.train.import_meta_graph('./model.meta') saver.restore(sess, './model') graph = tf.compat.v1.get_default_graph() - 定位PyFunc节点和其输入输出:
找到PyFunc节点的名称(比如PyFunc),以及对应的输入节点(如input_tensor)和输出节点(如pyfunc_output) - 用原生op实现等效逻辑:
比如原PyFunc是做x * 2 + 1的转换,就用原生op实现:input_tensor = graph.get_tensor_by_name('input_tensor:0') new_output = tf.add(tf.multiply(input_tensor, 2), 1, name='new_output') - 重新构建导出图并保存:
将原来依赖PyFunc输出的节点,改为依赖新的new_output,再用tf.compat.v1.saved_model.simple_save导出:tf.compat.v1.saved_model.simple_save( sess, './exported_model', inputs={'input': input_tensor}, outputs={'output': new_output} )
方法三:固化PyFunc的计算逻辑为常量(仅适用于无参数的固定转换)
如果PyFunc的转换逻辑完全固定(不依赖训练参数,仅对输入做固定处理),可直接将其计算结果固化为常量节点:
- 训练完成后,用固定输入喂给PyFunc节点,得到输出值
- 用
tf.constant创建常量节点,替换原来的PyFunc节点 - 重新保存模型,此时图中不再包含PyFunc节点
注意:该方法仅适用于转换逻辑不随输入变化的场景,动态计算逻辑不适用。
内容的提问来源于stack exchange,提问作者谭明超
相关产品推荐
相关产品推荐

