如何通过TensorFlow Serving Predict API获取预测标签?
如何让TensorFlow Serving返回带标签的预测结果?
问题背景
我有一个远程部署的Keras分类模型,通过RESTful API在站点调用进行预测。当前的请求方式如下:
- 请求URL:
POST https://my.site/models/v1:predict - 请求JSON:
{"instances": [["...用于预测标签的有效数据..."]]}
当前返回的结果仅包含预测概率数值(每个数值对应一个标签,但无标签名称):
{ "predictions": [ [ 0.001, 0.832, ..., 0.104 ] ] }
需求
希望调用API时能同时获取每个预测概率对应的标签,比如以下两种格式之一:
格式一(键值对形式)
{ "predictions": [ { "label_1": 0.001, "label_2": 0.832, ..., "label_n": 0.104 } ] }
格式二(概率数组+标签数组形式)
{ "predictions": [ [ 0.001, 0.832, ..., 0.104 ], [ "label_1", "label_2", ..., "label_n" ] ] }
已了解的相关方向
我知道导出模型时指定签名可能解决这个问题,但具体操作不清晰。核心是保存模型时添加signatures参数,示例代码框架如下:
@tf.function() def my_predict(my_prediction_inputs): ... my_signatures = my_predict.get_concrete_function(...) tf.keras.models.save_model(model, path, signatures=my_signatures, ...)
解决方案
要实现带标签的预测返回,你需要修改模型的导出逻辑,让模型输出包含标签信息。以下是两种可行的具体实现方法:
方法1:返回键值对格式的预测结果(推荐)
自定义预测函数,让模型直接输出以标签为键、概率为值的字典。假设你的标签列表是class_names,需与训练时的标签顺序完全一致:
import tensorflow as tf # 替换为你训练好的Keras模型 model = ... # 替换为你的实际标签列表,顺序必须和训练时一致 class_names = ["label_1", "label_2", "label_3"] @tf.function(input_signature=[tf.TensorSpec(shape=(None, 你的输入特征数), dtype=tf.float32)]) def predict_with_labels(inputs): # 获取模型的预测概率 predictions = model(inputs) # 将每个标签与对应位置的概率打包成字典 return {name: predictions[:, i] for i, name in enumerate(class_names)} # 保存模型时指定自定义的服务签名 tf.keras.models.save_model( model, "saved_model_path", signatures={"serving_default": predict_with_labels.get_concrete_function()} )
重新部署模型后,调用API将直接返回键值对格式的结果,与需求中的格式一一致。
方法2:返回概率数组+标签数组的组合
如果希望返回分开的概率数组和标签数组,可以让预测函数返回包含两个元素的列表:
import tensorflow as tf model = ... class_names = ["label_1", "label_2", "label_3"] @tf.function(input_signature=[tf.TensorSpec(shape=(None, 你的输入特征数), dtype=tf.float32)]) def predict_with_labels(inputs): predictions = model(inputs) # 返回预测概率数组,以及常量标签数组 return [predictions, tf.constant(class_names, dtype=tf.string)] # 保存模型 tf.keras.models.save_model( model, "saved_model_path", signatures={"serving_default": predict_with_labels.get_concrete_function()} )
部署后API返回的结果将与需求中的格式二一致。
关键注意事项
input_signature的形状和数据类型必须与模型输入匹配:比如输入是(样本数, 特征数)的二维数组,就设置shape=(None, 特征数),dtype根据实际输入数据调整。- 标签列表
class_names的顺序必须和模型训练时的标签映射完全一致,否则会出现标签与概率不匹配的错误。 - 修改后的模型需要重新部署到TensorFlow Serving服务上,新的签名才会生效。
内容的提问来源于stack exchange,提问作者sound wave
相关产品推荐
相关产品推荐

