如何在TF2.5 Keras SavedModel SignatureDef中自定义输出tensor_info.name
解决办法
方案1:Python侧预先固定输出张量名(最优,TF 2.5兼容)
你可以在定义serve函数的返回值时,用tf.identity给输出张量指定固定名称,完全避免自动生成的StatefulPartitionedCall:0这类不确定名称,修改后的代码如下:
@tf.function def serve(*args, **kwargs): outputs = model(*args, **kwargs) # 用tf.identity指定输出张量的固定名称 named_output = tf.identity(outputs, name='prediction_default_outputs') return {'outputs': named_output}
按上述代码保存模型后,输出的tensor_info.name会固定为prediction_default_outputs:0,和你之前拼接输入名的规则一致,不需要再额外查询。
方案2:不改保存逻辑,Python侧直接读取输出名(半可靠方案)
不需要借助C++ API,你可以直接在Python环境中加载保存好的SavedModel,从签名对象中直接读取输出张量的真实名称,代码如下:
import tensorflow as tf # 加载保存的模型 loaded_model = tf.saved_model.load("<你的模型保存目录路径>") # 取对应签名 pred_sig = loaded_model.signatures['prediction_default'] # 直接打印所有输出的名称和对应key for output_key, output_tensor in pred_sig.structured_outputs.items(): print(f"输出key {output_key} 对应的tensor_info.name为:{output_tensor.name}")
这个方法读取的结果和saved_model_cli的输出完全一致,而且不需要手动解析命令行输出,稳定性更高,完全适配你当前的TF 2.5环境。
补充:输入名也可通过同款方法读取
你之前拼接输入名的方式可以替换为上述读取签名的方法,避免拼接规则出错:
for input_key, input_tensor in pred_sig.structured_inputs.items(): print(f"输入key {input_key} 对应的tensor_info.name为:{input_tensor.name}")
内容的提问来源于stack exchange,提问作者glinka
相关产品推荐
相关产品推荐

