使用galeone/tfgo加载TensorFlow目标检测模型报错:找不到detection_boxes
问题根源:Go代码调用TF2导出的目标检测模型方式错误
结论
模型本身无问题(Python推理代码已验证正常运行),问题出在Go代码的模型调用逻辑与输入处理上。
原因分析
- TF2 SavedModel调用逻辑差异
TF2通过exporter_main_v2.py导出的目标检测模型,是基于签名(Signature)封装的服务化模型。Python代码中detect_fn(input_tensor)是直接调用模型的默认推断签名,你看到的detection_boxes、detection_scores等是签名输出的键名,并非计算图中的Operation名称——这就是遍历模型Graph的Operations找不到这些名称的核心原因。
你的Go代码错误地直接通过Operation名称查找输出,不符合TF2 SavedModel的调用规范。
- 输入格式错误
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
相关产品推荐
相关产品推荐

