要馈送的节点的名称取决于parsing_serving_input_fn 所做的事情,特别是它们应该是parsing_serving_input_fn 返回的Tensor 对象的名称。要获取的节点名称取决于您的预测(如果使用 Python 中的模型,则为 model.predict() 的参数)。
也就是说,TensorFlow 保存的模型格式确实包含模型的“签名”(即,可以馈送或获取的所有张量的名称)作为可以提供提示的元数据。
您可以从 Python 加载保存的模型并使用以下内容列出其签名:
with tf.Session() as sess:
md = tf.saved_model.loader.load(sess, ['serve'], export_dir)
sig = md.signature_def[tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY]
print(sig)
将打印如下内容:
inputs {
key: "inputs"
value {
name: "input_example_tensor:0"
dtype: DT_STRING
tensor_shape {
dim {
size: -1
}
}
}
}
outputs {
key: "scores"
value {
name: "linear/binary_logistic_head/predictions/probabilities:0"
dtype: DT_FLOAT
tensor_shape {
dim {
size: -1
}
dim {
size: 2
}
}
}
}
method_name: "tensorflow/serving/classify"
建议你想用 Java 做的是:
Tensor t = /* Tensor object to be fed */
model.session().runner().feed("input_example_tensor", t).fetch("linear/binary_logistic_head/predictions/probabilities").run()
如果您的程序包含为 TensorFlow 协议缓冲区生成的 Java 代码(打包在 org.tensorflow:proto artifact 中),您也可以纯粹在 Java 中提取此信息,使用如下:
// Same as tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY
// in Python. Perhaps this should be an exported constant in TensorFlow's Java API.
final String DEFAULT_SERVING_SIGNATURE_DEF_KEY = "serving_default";
final SignatureDef sig =
MetaGraphDef.parseFrom(model.metaGraphDef())
.getSignatureDefOrThrow(DEFAULT_SERVING_SIGNATURE_DEF_KEY);
您必须添加:
import org.tensorflow.framework.MetaGraphDef;
import org.tensorflow.framework.SignatureDef;
由于 Java API 和保存模型格式有些新,文档还有很大的改进空间。
希望对您有所帮助。