如何在tensorflow.rb中运行Python导出的SavedModel模型?
解决TensorFlow Ruby加载SavedModel时的ArgumentError问题
你遇到的问题是用Python的DNNClassifier训练导出模型后,在Ruby中使用tensorflow.rb加载预测时触发了无详细信息的ArgumentError,下面我们一步步分析并解决这个问题:
问题根源定位
从你提供的saved_model_cli输出和Ruby代码来看,核心问题是输入张量的形状不匹配:
- 模型期望的输入
Float_input_1和Float_input_2形状是(-1),也就是任意长度的一维张量(支持批量输入)。 - 但你在Ruby中创建的是单个值的标量张量(形状为0维),和模型要求的输入形状不兼容,导致TensorFlow内部抛出参数错误。
- 另外你提供的日志被截断,但结合代码来看,形状不匹配是最可能的触发原因。
修复方案
1. 调整输入张量的形状
在Ruby中创建符合模型期望的一维张量(将单个值包装成数组):
saved_model = Tensorflow::Savedmodel.new saved_model.LoadSavedModel('saved_model_pb', ['serve'], nil) input = [0.97, 1.00] # 创建形状为[1]的一维张量,匹配模型要求的(-1)形状 feature1_tensor = Tensorflow::Tensor.new([input[0]], dtype: Tensorflow::Float) feature2_tensor = Tensorflow::Tensor.new([input[1]], dtype: Tensorflow::Float) feature1_output = saved_model.graph.operation('Placeholder').output(0) feature2_output = saved_model.graph.operation('Placeholder_1').output(0) classes = saved_model.graph.operation('dnn/head/predictions/str_classes').output(0) feeds_tensor_to_output_hash = {feature1_output => feature1_tensor, feature2_output => feature2_tensor} out_tensor = saved_model.session.run(feeds_tensor_to_output_hash, [classes], []) puts out_tensor.first.value
2. 更可靠的方式:使用SignatureDef
直接硬编码节点名称容易出错(比如节点名称可能因导出方式变化),推荐使用模型的SignatureDef来获取输入输出:
saved_model = Tensorflow::Savedmodel.new saved_model.LoadSavedModel('saved_model_pb', ['serve'], nil) # 获取predict签名定义 predict_signature = saved_model.GetSignatureDef('predict') # 获取输入输出的张量信息 input1_info = predict_signature.inputs['Float_input_1'] input2_info = predict_signature.inputs['Float_input_2'] classes_info = predict_signature.outputs['classes'] # 解析张量名称(格式为"节点名:输出索引") def get_tensor(graph, tensor_name) op_name, output_idx = tensor_name.split(':') graph.operation(op_name).output(output_idx.to_i) end input1_tensor = get_tensor(saved_model.graph, input1_info.name) input2_tensor = get_tensor(saved_model.graph, input2_info.name) classes_tensor = get_tensor(saved_model.graph, classes_info.name) # 构造符合形状要求的输入 feed_dict = { input1_tensor => Tensorflow::Tensor.new([0.97], dtype: Tensorflow::Float), input2_tensor => Tensorflow::Tensor.new([1.00], dtype: Tensorflow::Float) } # 执行预测 result = saved_model.session.run(feed_dict, [classes_tensor], []) puts "预测类别:#{result.first.value}"
3. 验证Python端模型导出的正确性
确保你导出模型的serving_input_receiver_fn配置正确:
# 保持现有配置即可,它允许模型接受任意批量大小的输入 features = {'Float_input_1': tf.placeholder(tf.float32, shape=[None]), 'Float_input_2': tf.placeholder(tf.float32, shape=[None])} serving_input_receiver_fn = tf.estimator.export.build_raw_serving_input_receiver_fn(features, default_batch_size=None) classifier.export_savedmodel(SAVED_MODEL_FOLDER + '_pb', serving_input_receiver_fn, strip_default_attrs=True)
额外调试技巧
- 使用
debug_print查看张量的详细信息,确认形状是否匹配:
puts saved_model.graph.debug_print
- 确保tensorflow.rb的版本与Python中TensorFlow的版本尽量接近,版本差异可能导致兼容性问题。
内容的提问来源于stack exchange,提问作者Wiwiweb
相关产品推荐
相关产品推荐

