导出tf.estimator.DNNClassifier模型报错,如何正确保存?
问题分析与解决
你遇到的错误Feature sentence is not in features dictionary核心原因是模型训练时使用的特征名称和导出模型时定义的输入特征名称不匹配,再加上输入数据类型设置错误,导致模型找不到对应的输入特征。
具体问题点
- 特征名不匹配:从错误提示可以看出,你的模型在训练时依赖的特征名为
sentence(对应你创建的embedded_text_feature_column的原始特征列),但你在serving_input_receiver_fn里定义的输入键是x,模型加载时找不到名为sentence的特征,因此报错。 - 数据类型错误:文本分类任务的输入是原始字符串,你却设置了
tf.float32类型,这和模型训练时接收的输入类型不匹配。
修正后的代码
# 修正后的serving_input_receiver_fn def serving_input_receiver_fn(): """Build the serving inputs.""" # 特征名要和训练时使用的特征名完全一致(这里是"sentence") # 文本输入为字符串类型,所以dtype用tf.string inputs = {"sentence": tf.placeholder(shape=[None], dtype=tf.string)} return tf.estimator.export.ServingInputReceiver(inputs, inputs) # 注意变量名要和训练时的estimator一致,不要写成classifier export_dir = estimator.export_savedmodel( export_dir_base="/home/suhail/tensorflow-stubs/", serving_input_receiver_fn=serving_input_receiver_fn)
额外注意事项
- 如果你不确定训练时用的特征名,可以回头查看
train_input_fn里返回的特征字典的键,必须和serving_input_receiver_fn里的输入键完全一致。 shape=[None]表示支持批量输入,单条预测时只需要把文本放在列表里传入即可。
内容的提问来源于stack exchange,提问作者John Aristo
相关产品推荐
相关产品推荐

