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

使用galeone/tfgo加载TensorFlow目标检测模型报错:找不到detection_boxes

问题根源:Go代码调用TF2导出的目标检测模型方式错误

结论

模型本身无问题(Python推理代码已验证正常运行),问题出在Go代码的模型调用逻辑与输入处理上。

原因分析

  1. TF2 SavedModel调用逻辑差异
    TF2通过exporter_main_v2.py导出的目标检测模型,是基于签名(Signature)封装的服务化模型。Python代码中detect_fn(input_tensor)是直接调用模型的默认推断签名,你看到的detection_boxes、detection_scores等是签名输出的键名,并非计算图中的Operation名称——这就是遍历模型Graph的Operations找不到这些名称的核心原因。

你的Go代码错误地直接通过Operation名称查找输出,不符合TF2 SavedModel的调用规范。

  1. 输入格式错误
    Python代码会将图片解码为RGB像素数组,再转换为形状为[1, 高度, 宽度, 3]的uint8张量作为输入;而你的Go代码直接将图片文件的原始字节转成张量,输入格式完全不匹配模型要求,即使找到正确输出也会推理失败。

修正后的Go代码示例

以下是适配TF2导出模型的Go推理代码,解决了签名调用与输入格式问题:

package main

import (
	"fmt"
	"image"
	"image/png"
	"os"

	tf "github.com/galeone/tensorflow/tensorflow/go"
	tg "github.com/galeone/tfgo"
)

// 将图片转换为模型要求的输入张量:[1, H, W, 3] uint8
func imageToTensor(img image.Image) (*tf.Tensor, error) {
	bounds := img.Bounds()
	width, height := bounds.Dx(), bounds.Dy()

	// 按RGB顺序提取像素值
	pixels := make([]uint8, 0, height*width*3)
	for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
		for x := bounds.Min.X; x < bounds.Max.X; x++ {
			r, g, b, _ := img.At(x, y).RGBA()
			pixels = append(pixels, uint8(r>>8), uint8(g>>8), uint8(b>>8))
		}
	}

	// 构造模型要求的4维张量形状 [1, height, width, 3]
	img3D := make([][][]uint8, height)
	for i := range img3D {
		rowStart := i * width * 3
		rowEnd := rowStart + width*3
		row := pixels[rowStart:rowEnd]
		row2D := make([][]uint8, width)
		for j := 0; j < width; j++ {
			row2D[j] = row[j*3 : j*3+3]
		}
		img3D[i] = row2D
	}
	return tf.NewTensor([][][][]uint8{{img3D}})
}

func main() {
	// 加载模型,指定服务签名
	model := tg.LoadModel("saved_model", []string{"serve"}, nil)

	// 读取并解码PNG图片
	imgFile, err := os.Open("img.png")
	if err != nil {
		panic(err)
	}
	defer imgFile.Close()

	img, err := png.Decode(imgFile)
	if err != nil {
		panic(err)
	}

	// 转换为模型兼容的输入张量
	inputTensor, err := imageToTensor(img)
	if err != nil {
		panic(err)
	}

	// 通过模型签名对应的操作调用,TF2目标检测模型的核心推断操作是StatefulPartitionedCall
	results, err := model.Exec(
		[]tf.Output{
			model.Op("StatefulPartitionedCall", 0).Output(0), // detection_boxes
			model.Op("StatefulPartitionedCall", 0).Output(1), // detection_scores
			model.Op("StatefulPartitionedCall", 0).Output(2), // detection_classes
			model.Op("StatefulPartitionedCall", 0).Output(3), // num_detections
		},
		map[tf.Output]*tf.Tensor{
			model.Op("serving_default_image_tensor", 0): inputTensor,
		},
	)

	if err != nil {
		panic(err)
	}

	// 解析并输出结果
	detectionBoxes := results[0].Value().([][][]float32)
	detectionScores := results[1].Value().([][]float32)
	detectionClasses := results[2].Value().([][]float32)
	numDetections := results[3].Value().([]float32)

	fmt.Println("检测到目标数量:", int(numDetections[0]))
	for i := 0; i < int(numDetections[0]); i++ {
		if detectionScores[0][i] < 0.6 { // 过滤低置信度结果
			continue
		}
		fmt.Printf("目标%d:置信度%.2f,边界框%v,类别ID%.0f\n",
			i+1,
			detectionScores[0][i],
			detectionBoxes[0][i],
			detectionClasses[0][i],
		)
	}
}

关键说明

  • 可通过saved_model_cli show --dir saved_model --all命令查看模型的详细签名信息,确认输入输出的操作名称与索引。
  • 必须将图片解码为RGB像素数组并转换为[1, H, W, 3]形状的张量,这是TensorFlow目标检测模型的标准输入格式。

内容的提问来源于stack exchange,提问作者mineroot

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 07:45:00