如何用Golang移植TensorFlow Python预测代码?求相关教程
使用Golang实现TensorFlow模型预测(基于github.com/galeone/tensorflow/tensorflow/go)
前置准备
确保你已将训练好的Python TensorFlow模型导出为SavedModel格式(Python中可通过model.save("path/to/saved_model")完成导出)。
完整Go实现代码
package main import ( "fmt" "log" "math" tf "github.com/galeone/tensorflow/tensorflow/go" ) // 实现Sigmoid激活函数,对应Python的tf.nn.sigmoid func sigmoid(x float32) float32 { return 1.0 / (1.0 + float32(math.Exp(float64(-x)))) } func main() { // 1. 加载SavedModel模型文件 modelPath := "path/to/your/saved_model" graph := tf.NewGraph() session, err := tf.NewSession(graph, nil) if err != nil { log.Fatalf("创建会话失败: %v", err) } defer session.Close() // 加载SavedModel到计算图中 opts := &tf.SavedModelOptions{Tags: []string{"serve"}} if err := tf.LoadSavedModel(modelPath, opts.Tags, graph, session); err != nil { log.Fatalf("加载模型失败: %v", err) } // 2. 构造输入数据,对应Python代码中的sample字典 sample := map[string]int{ "b": 200, "c": 10, "d": 1, } // 3. 转换输入为张量,对应Python的tf.convert_to_tensor([value]) inputTensors := make(map[tf.Output]*tf.Tensor) for name, value := range sample { // 获取模型中对应的输入节点 inputOp := graph.Operation(name) if inputOp == nil { log.Fatalf("计算图中未找到输入节点: %s", name) } inputOutput := inputOp.Output(0) // 创建形状为[1]的张量(单样本输入) tensor, err := tf.NewTensor([]float32{float32(value)}) if err != nil { log.Fatalf("创建张量失败(%s): %v", name, err) } inputTensors[inputOutput] = tensor } // 4. 指定模型输出节点(需与Python模型的输出节点名称一致) // 可通过Python代码`print(reloaded_model.output_names)`获取输出节点名 outputOp := graph.Operation("your_model_output_node_name") // 替换为实际输出节点名 if outputOp == nil { log.Fatalf("计算图中未找到输出节点") } outputOutput := outputOp.Output(0) // 5. 运行会话获取预测结果,对应Python的reloaded_model.predict outputs, err := session.Run(inputTensors, []tf.Output{outputOutput}, nil) if err != nil { log.Fatalf("运行会话失败: %v", err) } // 6. 解析结果并计算Sigmoid概率(若模型输出为logits则需要此步骤) prediction := outputs[0].Value().([][]float32)[0][0] prob := sigmoid(prediction) fmt.Printf("预测概率: %f\n", prob) }
关键注意事项
- 节点名称匹配:输入、输出节点名称必须与Python模型中定义的完全一致,可通过Python代码
print(reloaded_model.input_names)和print(reloaded_model.output_names)查看。 - 张量形状:Go中创建的张量形状要与模型输入要求匹配(示例为单样本输入,形状为
[1])。 - Sigmoid计算:如果模型输出已经经过Sigmoid激活(直接输出概率),可跳过手动计算步骤。
内容的提问来源于stack exchange,提问作者Helmut Januschka
相关产品推荐
相关产品推荐

