如何在ML Engine批量预测中用自定义TensorFlow模型获取Keys
适配自定义tf.contrib.learn Estimator保留批量预测的Key
我之前在做自定义Estimator的ML Engine批量预测时也踩过这个坑,刚好可以给你分享下适配的方法——核心就是要让Key从输入环节一直传递到最终的预测输出里,具体分这几步来做:
1. 在输入函数中保留Key字段
不管是训练还是预测用的input_fn,都要把原始输入里的Key作为特征的一部分加载进来,不能只提取模型训练需要的特征。比如你用CSV作为输入源:
def predict_input_fn(file_pattern): # 定义CSV列名,包含你的Key列(比如叫"user_id")和其他特征列 column_names = ["user_id", "feature1", "feature2", ...] # 读取CSV文件 dataset = tf.contrib.data.make_csv_dataset( file_pattern, batch_size=32, column_names=column_names, label_name=None # 预测时不需要标签 ) # 将数据集转换成特征字典,这里user_id会被保留在features里 features = dataset.make_one_shot_iterator().get_next() return features
如果是JSON格式输入,也要确保解析时把Key字段加入到特征字典中。
2. 修改Model_fn,在预测输出中携带Key
在自定义Estimator的model_fn里,当模式为PREDICT时,要把传入的Key和预测结果一起打包到返回的predictions字典中:
def custom_model_fn(features, labels, mode, params): # 模型结构定义(这里省略你的网络层代码) logits = your_model_layers(features) predicted_values = tf.sigmoid(logits) # 举个例子,根据你的任务调整 if mode == tf.estimator.ModeKeys.PREDICT: # 关键:把Key加入到预测结果字典里 predictions = { "key": features["user_id"], # 这里对应你输入中的Key字段名 "predicted_value": predicted_values } return tf.estimator.EstimatorSpec( mode=mode, predictions=predictions ) # 训练和评估模式的代码(略) ...
这里要注意,features["user_id"]的键名要和你输入函数中保留的Key字段名完全一致。
3. 导出模型时包含Key的输入签名
导出SavedModel给ML Engine用的时候,要通过serving_input_receiver_fn明确声明输入包含Key,确保ML Engine能正确解析输入并传递Key:
def serving_input_receiver_fn(): # 定义和输入格式匹配的占位符,Key一般是字符串类型 key_placeholder = tf.placeholder(tf.string, shape=[None], name="user_id") feature1_placeholder = tf.placeholder(tf.float32, shape=[None], name="feature1") feature2_placeholder = tf.placeholder(tf.float32, shape=[None], name="feature2") # 构建特征字典,和输入函数、model_fn中的字段对应 features = { "user_id": key_placeholder, "feature1": feature1_placeholder, "feature2": feature2_placeholder } return tf.estimator.export.ServingInputReceiver(features, features) # 导出模型 estimator = tf.contrib.learn.Estimator(model_fn=custom_model_fn, params=...) estimator.export_savedmodel(export_dir="your_export_path", serving_input_receiver_fn=serving_input_receiver_fn)
最后验证
部署到ML Engine之前,可以先本地跑一次预测测试:
predictions = estimator.predict(input_fn=lambda: predict_input_fn("test_data.csv")) for pred in predictions: print(f"Key: {pred['key']}, Prediction: {pred['predicted_value']}")
如果本地输出能正常显示Key和对应预测结果,那部署到ML Engine后,批量预测的输出文件里就会包含Key字段,用来匹配原始输入了。
内容的提问来源于stack exchange,提问作者dnovai
相关产品推荐
相关产品推荐

