You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何实现接收JSON输入的tf.contrib.learn Estimator模型导出服务?

导出接收JSON输入的tf.contrib.learn模型

嘿,我之前也遇到过用tf.contrib.learn导出模型时不想碰tf.Example的麻烦!其实完全可以绕过它,直接导出能接收JSON输入的模型,核心就是用build_raw_serving_input_receiver_fn这个工具函数,我给你一步步讲清楚怎么弄:

核心思路

tf.contrib.learn.build_raw_serving_input_receiver_fn可以让你直接定义原始张量输入规格,不需要把数据包装成tf.Example格式。导出后的模型能直接接收键名和输入特征对应的JSON数据,完美解决tf.Example构建复杂的问题。

完整代码示例

假设我们已经有一个训练好的tf.contrib.learn模型(这里用DNN分类器做演示),下面是导出的完整流程:

1. 先定义模型(训练部分省略)

import tensorflow as tf

# 定义模型的输入特征列
feature_columns = [
    tf.contrib.layers.real_valued_column("age", dimension=1),
    tf.contrib.layers.sparse_column_with_keys("gender", keys=["male", "female"])
]

# 创建tf.contrib.learn Estimator
estimator = tf.contrib.learn.DNNClassifier(
    feature_columns=feature_columns,
    hidden_units=[64, 32],
    n_classes=2
)

# --- 这里省略模型训练代码 ---

2. 构建JSON友好的Serving Input Receiver函数

# 定义输入特征的张量规格:键名对应JSON的键,shape和dtype要和训练时一致
feature_spec = {
    # age是数值特征,[None]表示支持批量输入
    "age": tf.FixedLenFeature(shape=[None], dtype=tf.float32),
    # gender是字符串分类特征
    "gender": tf.FixedLenFeature(shape=[None], dtype=tf.string)
}

# 构建原始输入接收函数,跳过tf.Example的转换
serving_input_receiver_fn = tf.contrib.learn.build_raw_serving_input_receiver_fn(feature_spec)

3. 导出模型

# 指定导出路径
export_dir = "./saved_model_json"
# 执行导出
estimator.export_savedmodel(export_dir, serving_input_receiver_fn)

关键说明

  • 导出后的模型可以直接接收JSON格式的请求,比如这样的输入:
    {
        "age": [28.0, 35.0],
        "gender": ["female", "male"]
    }
    
  • feature_spec里的键必须和你模型输入特征的名称完全一致,dtype和shape也要和训练时的特征列匹配,否则模型会报错。
  • shape=[None]表示支持批量输入,如果你只需要单样本输入,可以改成shape=[],但保留[None]会更灵活。

内容的提问来源于stack exchange,提问作者jamborta

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 08:34:15