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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:42:20