如何实现接收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
相关产品推荐
相关产品推荐

