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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 15:31:58