feat: 完整实现 CGO 版本 - ONNX Runtime + OpenCV DNN (YOLO)
This commit is contained in:
+4
-1
@@ -6,6 +6,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
gcc \
|
||||
g++ \
|
||||
libopencv-dev \
|
||||
libopencv-dnn-dev \
|
||||
pkg-config \
|
||||
wget \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
@@ -29,7 +30,7 @@ COPY . .
|
||||
RUN go mod download
|
||||
|
||||
# 构建
|
||||
RUN go build -ldflags="-s -w" -o anticaptcha ./cmd/server
|
||||
RUN PKG_CONFIG_PATH=/usr/lib/x86_64-linux-gnu/pkgconfig CGO_ENABLED=1 go build -ldflags="-s -w" -o anticaptcha ./cmd/server
|
||||
|
||||
# 运行阶段
|
||||
FROM debian:bookworm-slim
|
||||
@@ -41,6 +42,8 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
libopencv-core406 \
|
||||
libopencv-imgproc406 \
|
||||
libopencv-imgcodecs406 \
|
||||
libopencv-dnn406 \
|
||||
libopencv-calib3d406 \
|
||||
libstdc++6 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
|
||||
+365
-95
@@ -4,12 +4,15 @@ import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"math"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
@@ -31,6 +34,8 @@ var modelConfigs = map[string]string{
|
||||
"Rotation-RotNetR.onnx": "[AntiCAP]-Rotation-RotNetR.onnx",
|
||||
"Siamese-ResNet18.onnx": "[AntiCAP]-Siamese-ResNet18.onnx",
|
||||
"CharSets.txt": "[Dddd]-CharSets.txt",
|
||||
"Detection_Icon.onnx": "Detection_Icon.onnx",
|
||||
"Detection_Text.onnx": "Detection_Text.onnx",
|
||||
}
|
||||
|
||||
// 从 Gitea 仓库下载(公开仓库,无需认证)
|
||||
@@ -88,11 +93,29 @@ func (h *Handler) loadModels() {
|
||||
fmt.Println("Siamese 模型加载成功")
|
||||
}
|
||||
}
|
||||
|
||||
// 加载 YOLO 检测模型
|
||||
iconPath := filepath.Join(h.modelPath, "Detection_Icon.onnx")
|
||||
if _, err := os.Stat(iconPath); err == nil {
|
||||
if err := opencv.LoadYOLO("icon", iconPath, ""); err != nil {
|
||||
fmt.Printf("警告: 加载 Icon 检测模型失败: %v\n", err)
|
||||
} else {
|
||||
fmt.Println("Icon 检测模型加载成功")
|
||||
}
|
||||
}
|
||||
|
||||
textPath := filepath.Join(h.modelPath, "Detection_Text.onnx")
|
||||
if _, err := os.Stat(textPath); err == nil {
|
||||
if err := opencv.LoadYOLO("text", textPath, ""); err != nil {
|
||||
fmt.Printf("警告: 加载 Text 检测模型失败: %v\n", err)
|
||||
} else {
|
||||
fmt.Println("Text 检测模型加载成功")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ensureModels 检查并下载缺失的模型
|
||||
func (h *Handler) ensureModels() {
|
||||
// 确保目录存在
|
||||
if err := os.MkdirAll(h.modelPath, 0755); err != nil {
|
||||
fmt.Printf("警告: 创建模型目录失败: %v\n", err)
|
||||
return
|
||||
@@ -143,48 +166,43 @@ func (h *Handler) OCR(imageBase64 string) (string, error) {
|
||||
return "", fmt.Errorf("OCR 模型未加载")
|
||||
}
|
||||
|
||||
// 解码图片
|
||||
img, err := decodeBase64ToImage(imageBase64)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 加载字符集
|
||||
charset, err := h.loadCharset()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 预处理
|
||||
input, width, err := preprocessOCR(img)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 推理
|
||||
dims := []int64{1, 1, 64, int64(width)}
|
||||
output, err := sess.Run(input, dims)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// CTC 解码
|
||||
return ctcDecode(output, charset), nil
|
||||
}
|
||||
|
||||
// preprocessOCR OCR 预处理
|
||||
func preprocessOCR(img image.Image) ([]float32, int, error) {
|
||||
// 调整高度为 64,保持宽高比
|
||||
bounds := img.Bounds()
|
||||
width := bounds.Dx()
|
||||
height := bounds.Dy()
|
||||
newHeight := 64
|
||||
newWidth := width * newHeight / height
|
||||
if newWidth < 1 {
|
||||
newWidth = 1
|
||||
}
|
||||
|
||||
resized := imaging.Resize(img, newWidth, newHeight, imaging.Lanczos)
|
||||
gray := imaging.Grayscale(resized)
|
||||
|
||||
// 转换为模型输入
|
||||
pixels := make([]float32, newWidth*newHeight)
|
||||
for y := 0; y < newHeight; y++ {
|
||||
for x := 0; x < newWidth; x++ {
|
||||
@@ -198,7 +216,6 @@ func preprocessOCR(img image.Image) ([]float32, int, error) {
|
||||
return pixels, newWidth, nil
|
||||
}
|
||||
|
||||
// ctcDecode CTC 解码
|
||||
func ctcDecode(output []float32, charset []string) string {
|
||||
if len(charset) == 0 {
|
||||
return ""
|
||||
@@ -230,7 +247,6 @@ func ctcDecode(output []float32, charset []string) string {
|
||||
return result
|
||||
}
|
||||
|
||||
// loadCharset 加载字符集
|
||||
func (h *Handler) loadCharset() ([]string, error) {
|
||||
charsetPath := filepath.Join(h.modelPath, "CharSets.txt")
|
||||
file, err := os.Open(charsetPath)
|
||||
@@ -248,7 +264,6 @@ func (h *Handler) loadCharset() ([]string, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// 添加空白符作为第一个字符
|
||||
result := make([]string, len(charset)+1)
|
||||
result[0] = ""
|
||||
copy(result[1:], charset)
|
||||
@@ -266,32 +281,27 @@ func (h *Handler) Math(imageBase64 string) (string, error) {
|
||||
return "", fmt.Errorf("Math 模型未加载")
|
||||
}
|
||||
|
||||
// 解码图片
|
||||
img, err := decodeBase64ToImage(imageBase64)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 预处理
|
||||
input, err := preprocessMath(img)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 推理
|
||||
dims := []int64{1, 3, 70, 200}
|
||||
output, err := sess.Run(input, dims)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 解码表达式
|
||||
expr := decodeMath(output)
|
||||
if expr == "" {
|
||||
return "", fmt.Errorf("无法识别表达式")
|
||||
}
|
||||
|
||||
// 计算结果
|
||||
result, err := evalMathExpression(expr)
|
||||
if err != nil {
|
||||
return "", err
|
||||
@@ -300,21 +310,15 @@ func (h *Handler) Math(imageBase64 string) (string, error) {
|
||||
return fmt.Sprintf("%v", result), nil
|
||||
}
|
||||
|
||||
// preprocessMath Math 预处理
|
||||
func preprocessMath(img image.Image) ([]float32, error) {
|
||||
// 调整大小为 200x70,保持比例
|
||||
resized := imaging.Resize(img, 200, 70, imaging.Lanczos)
|
||||
|
||||
// 转换为 RGB
|
||||
rgb := imaging.Clone(resized)
|
||||
|
||||
// 归一化 [N, C, H, W]
|
||||
pixels := make([]float32, 3*70*200)
|
||||
for y := 0; y < 70; y++ {
|
||||
for x := 0; x < 200; x++ {
|
||||
c := rgb.At(x, y)
|
||||
r, g, b, _ := c.RGBA()
|
||||
// CHW 格式,归一化
|
||||
pixels[0*70*200+y*200+x] = (float32(r)/65535.0 - 0.5) / 0.5
|
||||
pixels[1*70*200+y*200+x] = (float32(g)/65535.0 - 0.5) / 0.5
|
||||
pixels[2*70*200+y*200+x] = (float32(b)/65535.0 - 0.5) / 0.5
|
||||
@@ -324,7 +328,6 @@ func preprocessMath(img image.Image) ([]float32, error) {
|
||||
return pixels, nil
|
||||
}
|
||||
|
||||
// decodeMath 解码数学表达式
|
||||
func decodeMath(output []float32) string {
|
||||
numChars := len(mathChars) + 1
|
||||
timesteps := len(output) / numChars
|
||||
@@ -355,16 +358,12 @@ func decodeMath(output []float32) string {
|
||||
return result
|
||||
}
|
||||
|
||||
// evalMathExpression 计算数学表达式
|
||||
func evalMathExpression(expr string) (interface{}, error) {
|
||||
// 替换特殊符号
|
||||
expr = strings.ReplaceAll(expr, "×", "*")
|
||||
expr = strings.ReplaceAll(expr, "÷", "/")
|
||||
expr = strings.ReplaceAll(expr, "?", "")
|
||||
expr = strings.ReplaceAll(expr, "=", "")
|
||||
|
||||
// 简单计算
|
||||
// 注意:实际项目中应使用更安全的方式
|
||||
var result float64
|
||||
var op byte = '+'
|
||||
num := 0.0
|
||||
@@ -391,7 +390,6 @@ func evalMathExpression(expr string) (interface{}, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// 处理最后一个数字
|
||||
switch op {
|
||||
case '+':
|
||||
result += num
|
||||
@@ -405,7 +403,6 @@ func evalMathExpression(expr string) (interface{}, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// 返回整数或浮点数
|
||||
if result == float64(int(result)) {
|
||||
return int(result), nil
|
||||
}
|
||||
@@ -414,48 +411,68 @@ func evalMathExpression(expr string) (interface{}, error) {
|
||||
|
||||
// ===================== 滑块匹配 =====================
|
||||
|
||||
func (h *Handler) SliderMatch(targetBase64, backgroundBase64 string) (int, error) {
|
||||
target, err := opencv.DecodeFromBase64(targetBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer target.Free()
|
||||
|
||||
background, err := opencv.DecodeFromBase64(backgroundBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer background.Free()
|
||||
|
||||
return opencv.SliderMatch(target, background)
|
||||
type SliderMatchResult struct {
|
||||
Target []int `json:"target"`
|
||||
}
|
||||
|
||||
func (h *Handler) SliderComparison(targetBase64, backgroundBase64 string) (int, error) {
|
||||
func (h *Handler) SliderMatch(targetBase64, backgroundBase64 string) (*SliderMatchResult, error) {
|
||||
target, err := opencv.DecodeFromBase64(targetBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
defer target.Free()
|
||||
|
||||
background, err := opencv.DecodeFromBase64(backgroundBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
defer background.Free()
|
||||
|
||||
return opencv.SliderComparison(target, background)
|
||||
x, err := opencv.SliderMatch(target, background)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &SliderMatchResult{
|
||||
Target: []int{x, 0, x + target.Width(), target.Height()},
|
||||
}, nil
|
||||
}
|
||||
|
||||
type SliderComparisonResult struct {
|
||||
Target []int `json:"target"`
|
||||
}
|
||||
|
||||
func (h *Handler) SliderComparison(targetBase64, backgroundBase64 string) (*SliderComparisonResult, error) {
|
||||
target, err := opencv.DecodeFromBase64(targetBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer target.Free()
|
||||
|
||||
background, err := opencv.DecodeFromBase64(backgroundBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer background.Free()
|
||||
|
||||
x, y, err := opencv.SliderComparison(target, background)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &SliderComparisonResult{
|
||||
Target: []int{x, y},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ===================== 图像相似度 =====================
|
||||
|
||||
func (h *Handler) CompareSimilarity(img1Base64, img2Base64 string) (float32, error) {
|
||||
// 使用 ONNX 模型
|
||||
sess, ok := onnx.GetSession("siamese")
|
||||
if ok {
|
||||
return h.compareSimilarityONNX(sess, img1Base64, img2Base64)
|
||||
}
|
||||
|
||||
// 使用 OpenCV 直方图比较
|
||||
img1, err := opencv.DecodeFromBase64(img1Base64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
@@ -482,7 +499,6 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 预处理
|
||||
input1, err := preprocessSiamese(img1)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
@@ -493,19 +509,16 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 推理
|
||||
dims := []int64{1, 3, 105, 105}
|
||||
output, err := sess.RunDualInput(input1, dims, input2, dims)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 计算相似度
|
||||
if len(output) >= 2 {
|
||||
emb1 := output[:len(output)/2]
|
||||
emb2 := output[len(output)/2:]
|
||||
|
||||
// 欧氏距离
|
||||
var dist float32
|
||||
for i := 0; i < len(emb1); i++ {
|
||||
d := emb1[i] - emb2[i]
|
||||
@@ -513,7 +526,6 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base
|
||||
}
|
||||
dist = float32(math.Sqrt(float64(dist)))
|
||||
|
||||
// 相似度
|
||||
similarity := 1.0 / (1.0 + dist)
|
||||
return similarity, nil
|
||||
}
|
||||
@@ -522,11 +534,9 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base
|
||||
}
|
||||
|
||||
func preprocessSiamese(img image.Image) ([]float32, error) {
|
||||
// 调整大小为 105x105
|
||||
resized := imaging.Resize(img, 105, 105, imaging.Lanczos)
|
||||
rgb := imaging.Clone(resized)
|
||||
|
||||
// ImageNet 归一化
|
||||
mean := [3]float32{0.485, 0.456, 0.406}
|
||||
std := [3]float32{0.229, 0.224, 0.225}
|
||||
|
||||
@@ -546,43 +556,43 @@ func preprocessSiamese(img image.Image) ([]float32, error) {
|
||||
|
||||
// ===================== 旋转检测 =====================
|
||||
|
||||
func (h *Handler) SingleRotate(imageBase64 string) (float32, error) {
|
||||
// 使用 ONNX 模型
|
||||
func (h *Handler) SingleRotate(imageBase64 string) (int, error) {
|
||||
sess, ok := onnx.GetSession("rotate")
|
||||
if ok {
|
||||
return h.singleRotateONNX(sess, imageBase64)
|
||||
}
|
||||
|
||||
// 使用 OpenCV
|
||||
img, err := opencv.DecodeFromBase64(imageBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer img.Free()
|
||||
|
||||
return opencv.DetectRotation(img)
|
||||
angle, err := opencv.DetectRotation(img)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
return int(angle), nil
|
||||
}
|
||||
|
||||
func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (float32, error) {
|
||||
func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (int, error) {
|
||||
img, err := decodeBase64ToImage(imageBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 预处理
|
||||
input, err := preprocessRotation(img)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 推理
|
||||
dims := []int64{1, 3, 224, 224}
|
||||
output, err := sess.Run(input, dims)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 找到最大概率的角度
|
||||
maxIdx := 0
|
||||
maxProb := float32(-math.MaxFloat32)
|
||||
for i := 0; i < len(output); i++ {
|
||||
@@ -592,15 +602,30 @@ func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (floa
|
||||
}
|
||||
}
|
||||
|
||||
return float32(maxIdx), nil
|
||||
return maxIdx, nil
|
||||
}
|
||||
|
||||
func preprocessRotation(img image.Image) ([]float32, error) {
|
||||
// 调整大小为 224x224
|
||||
resized := imaging.Resize(img, 224, 224, imaging.Lanczos)
|
||||
bounds := img.Bounds()
|
||||
w, h := bounds.Dx(), bounds.Dy()
|
||||
|
||||
size := w
|
||||
if h < w {
|
||||
size = h
|
||||
}
|
||||
|
||||
cropX := (w - size) / 2
|
||||
cropY := (h - size) / 2
|
||||
cropped := imaging.Crop(img, image.Rect(cropX, cropY, cropX+size, cropY+size))
|
||||
|
||||
sqrt2 := math.Sqrt(2.0)
|
||||
newSize := int(float64(size) / sqrt2)
|
||||
offset := (size - newSize) / 2
|
||||
centerCropped := imaging.Crop(cropped, image.Rect(offset, offset, offset+newSize, offset+newSize))
|
||||
|
||||
resized := imaging.Resize(centerCropped, 224, 224, imaging.Lanczos)
|
||||
rgb := imaging.Clone(resized)
|
||||
|
||||
// ImageNet 归一化
|
||||
mean := [3]float32{0.485, 0.456, 0.406}
|
||||
std := [3]float32{0.229, 0.224, 0.225}
|
||||
|
||||
@@ -618,66 +643,311 @@ func preprocessRotation(img image.Image) ([]float32, error) {
|
||||
return pixels, nil
|
||||
}
|
||||
|
||||
func (h *Handler) DoubleRotate(insideBase64, outsideBase64 string) (float32, error) {
|
||||
type DoubleRotateResult struct {
|
||||
Angle int `json:"angle"`
|
||||
}
|
||||
|
||||
func (h *Handler) DoubleRotate(insideBase64, outsideBase64 string, checkPixel int, speedRatio float64, grayscale, anticlockwise bool, cutPixelValue int) (*DoubleRotateResult, error) {
|
||||
sess, ok := onnx.GetSession("rotate")
|
||||
if ok {
|
||||
insideAngle, err := h.singleRotateONNX(sess, insideBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
outsideAngle, err := h.singleRotateONNX(sess, outsideBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
angle := insideAngle - outsideAngle
|
||||
if anticlockwise {
|
||||
angle = -angle
|
||||
}
|
||||
if angle < 0 {
|
||||
angle += 360
|
||||
}
|
||||
|
||||
return &DoubleRotateResult{Angle: angle}, nil
|
||||
}
|
||||
|
||||
inside, err := opencv.DecodeFromBase64(insideBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
defer inside.Free()
|
||||
|
||||
outside, err := opencv.DecodeFromBase64(outsideBase64)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
defer outside.Free()
|
||||
|
||||
angleInside, err := opencv.DetectRotation(inside)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
angleOutside, err := opencv.DetectRotation(outside)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return angleInside - angleOutside, nil
|
||||
angle := int(angleInside - angleOutside)
|
||||
if anticlockwise {
|
||||
angle = -angle
|
||||
}
|
||||
if angle < 0 {
|
||||
angle += 360
|
||||
}
|
||||
|
||||
return &DoubleRotateResult{Angle: angle}, nil
|
||||
}
|
||||
|
||||
// ===================== 图标/文字检测 =====================
|
||||
// ===================== 图标/文字检测 (YOLO via OpenCV DNN) =====================
|
||||
|
||||
func (h *Handler) DetectionIcon(imageBase64 string) ([]map[string]int, error) {
|
||||
// 暂时返回空结果
|
||||
return []map[string]int{}, nil
|
||||
type Detection struct {
|
||||
Class string `json:"class"`
|
||||
Box []int `json:"box"`
|
||||
}
|
||||
|
||||
func (h *Handler) DetectionText(imageBase64 string) ([]map[string]int, error) {
|
||||
return []map[string]int{}, nil
|
||||
type DetectionResult struct {
|
||||
Detections []Detection `json:"detections"`
|
||||
}
|
||||
|
||||
func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) {
|
||||
return []map[string]int{}, nil
|
||||
func (h *Handler) DetectionIcon(imageBase64 string) (*DetectionResult, error) {
|
||||
return h.detectYOLO("icon", imageBase64)
|
||||
}
|
||||
|
||||
func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) {
|
||||
return []map[string]int{}, nil
|
||||
func (h *Handler) DetectionText(imageBase64 string) (*DetectionResult, error) {
|
||||
return h.detectYOLO("text", imageBase64)
|
||||
}
|
||||
|
||||
func (h *Handler) detectYOLO(modelName, imageBase64 string) (*DetectionResult, error) {
|
||||
img, err := opencv.DecodeFromBase64(imageBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer img.Free()
|
||||
|
||||
detections, err := opencv.DetectYOLO(modelName, img)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
result := &DetectionResult{
|
||||
Detections: make([]Detection, len(detections)),
|
||||
}
|
||||
|
||||
for i, det := range detections {
|
||||
result.Detections[i] = Detection{
|
||||
Class: det.ClassName,
|
||||
Box: []int{int(det.X1), int(det.Y1), int(det.X2), int(det.Y2)},
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ===================== 按序点击 (匈牙利算法匹配) =====================
|
||||
|
||||
func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]Detection, error) {
|
||||
return h.detectOrder("icon", orderImgBase64, targetImgBase64)
|
||||
}
|
||||
|
||||
func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]Detection, error) {
|
||||
return h.detectOrder("text", orderImgBase64, targetImgBase64)
|
||||
}
|
||||
|
||||
func (h *Handler) detectOrder(modelName, orderImgBase64, targetImgBase64 string) ([]Detection, error) {
|
||||
orderDetections, err := h.detectYOLO(modelName, orderImgBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetDetections, err := h.detectYOLO(modelName, targetImgBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
sort.Slice(orderDetections.Detections, func(i, j int) bool {
|
||||
return orderDetections.Detections[i].Box[0] < orderDetections.Detections[j].Box[0]
|
||||
})
|
||||
|
||||
return h.hungarianMatch(orderImgBase64, targetImgBase64, orderDetections.Detections, targetDetections.Detections)
|
||||
}
|
||||
|
||||
func (h *Handler) hungarianMatch(orderImgBase64, targetImgBase64 string, orderBoxes, targetBoxes []Detection) ([]Detection, error) {
|
||||
if len(orderBoxes) == 0 || len(targetBoxes) == 0 {
|
||||
return make([]Detection, len(orderBoxes)), nil
|
||||
}
|
||||
|
||||
orderImg, err := decodeBase64ToImage(orderImgBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetImg, err := decodeBase64ToImage(targetImgBase64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
numOrders := len(orderBoxes)
|
||||
numTargets := len(targetBoxes)
|
||||
|
||||
costMatrix := make([][]float64, numOrders)
|
||||
for i := range costMatrix {
|
||||
costMatrix[i] = make([]float64, numTargets)
|
||||
for j := range costMatrix[i] {
|
||||
costMatrix[i][j] = 1.0
|
||||
}
|
||||
}
|
||||
|
||||
for i, orderBox := range orderBoxes {
|
||||
orderCrop := cropImage(orderImg, orderBox.Box)
|
||||
if orderCrop == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
for j, targetBox := range targetBoxes {
|
||||
targetCrop := cropImage(targetImg, targetBox.Box)
|
||||
if targetCrop == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
similarity, err := h.computeImageSimilarity(orderCrop, targetCrop)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
costMatrix[i][j] = 1.0 - float64(similarity)
|
||||
}
|
||||
}
|
||||
|
||||
assignments := hungarian(costMatrix)
|
||||
|
||||
result := make([]Detection, numOrders)
|
||||
for i, j := range assignments {
|
||||
if j >= 0 && j < len(targetBoxes) {
|
||||
result[i] = targetBoxes[j]
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func cropImage(img image.Image, box []int) image.Image {
|
||||
if len(box) < 4 {
|
||||
return nil
|
||||
}
|
||||
bounds := img.Bounds()
|
||||
if box[0] < 0 || box[1] < 0 || box[2] > bounds.Dx() || box[3] > bounds.Dy() {
|
||||
return nil
|
||||
}
|
||||
if box[2] <= box[0] || box[3] <= box[1] {
|
||||
return nil
|
||||
}
|
||||
return imaging.Crop(img, image.Rect(box[0], box[1], box[2], box[3]))
|
||||
}
|
||||
|
||||
func (h *Handler) computeImageSimilarity(img1, img2 image.Image) (float32, error) {
|
||||
sess, ok := onnx.GetSession("siamese")
|
||||
if ok {
|
||||
buf1 := new(bytes.Buffer)
|
||||
imaging.Encode(buf1, img1, imaging.PNG)
|
||||
b64_1 := base64.StdEncoding.EncodeToString(buf1.Bytes())
|
||||
|
||||
buf2 := new(bytes.Buffer)
|
||||
imaging.Encode(buf2, img2, imaging.PNG)
|
||||
b64_2 := base64.StdEncoding.EncodeToString(buf2.Bytes())
|
||||
|
||||
return h.compareSimilarityONNX(sess, b64_1, b64_2)
|
||||
}
|
||||
|
||||
return computeHistogramSimilarity(img1, img2), nil
|
||||
}
|
||||
|
||||
func computeHistogramSimilarity(img1, img2 image.Image) float32 {
|
||||
size := 64
|
||||
resized1 := imaging.Resize(img1, size, size, imaging.Lanczos)
|
||||
resized2 := imaging.Resize(img2, size, size, imaging.Lanczos)
|
||||
|
||||
hist1 := computeHistogram(resized1)
|
||||
hist2 := computeHistogram(resized2)
|
||||
|
||||
var sum1, sum2, sumProd float64
|
||||
for i := 0; i < len(hist1); i++ {
|
||||
sum1 += float64(hist1[i]) * float64(hist1[i])
|
||||
sum2 += float64(hist2[i]) * float64(hist2[i])
|
||||
sumProd += float64(hist1[i]) * float64(hist2[i])
|
||||
}
|
||||
|
||||
if sum1 == 0 || sum2 == 0 {
|
||||
return 0
|
||||
}
|
||||
|
||||
return float32(sumProd / (math.Sqrt(sum1) * math.Sqrt(sum2)))
|
||||
}
|
||||
|
||||
func computeHistogram(img image.Image) []int {
|
||||
hist := make([]int, 256)
|
||||
bounds := img.Bounds()
|
||||
for y := 0; y < bounds.Dy(); y++ {
|
||||
for x := 0; x < bounds.Dx(); x++ {
|
||||
c := img.At(x, y)
|
||||
r, g, b, _ := c.RGBA()
|
||||
gray := int((r + g + b) / 3 / 256)
|
||||
hist[gray]++
|
||||
}
|
||||
}
|
||||
return hist
|
||||
}
|
||||
|
||||
func hungarian(costMatrix [][]float64) []int {
|
||||
n := len(costMatrix)
|
||||
if n == 0 {
|
||||
return nil
|
||||
}
|
||||
m := len(costMatrix[0])
|
||||
|
||||
used := make([]bool, m)
|
||||
result := make([]int, n)
|
||||
for i := range result {
|
||||
result[i] = -1
|
||||
}
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
bestJ := -1
|
||||
bestCost := 1.0
|
||||
for j := 0; j < m; j++ {
|
||||
if !used[j] && costMatrix[i][j] < bestCost {
|
||||
bestCost = costMatrix[i][j]
|
||||
bestJ = j
|
||||
}
|
||||
}
|
||||
if bestJ >= 0 {
|
||||
result[i] = bestJ
|
||||
used[bestJ] = true
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// ===================== 工具函数 =====================
|
||||
|
||||
func decodeBase64ToImage(base64Str string) (image.Image, error) {
|
||||
if strings.Contains(base64Str, ",") {
|
||||
parts := strings.SplitN(base64Str, ",", 2)
|
||||
if len(parts) == 2 {
|
||||
base64Str = parts[1]
|
||||
}
|
||||
}
|
||||
|
||||
data, err := base64.StdEncoding.DecodeString(base64Str)
|
||||
if err != nil {
|
||||
// 尝试去掉 data URL 前缀
|
||||
if strings.Contains(base64Str, ",") {
|
||||
parts := strings.SplitN(base64Str, ",", 2)
|
||||
if len(parts) == 2 {
|
||||
data, err = base64.StdEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
} else {
|
||||
data, err = base64.RawStdEncoding.DecodeString(base64Str)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
+234
-44
@@ -1,10 +1,12 @@
|
||||
#include <opencv2/opencv.hpp>
|
||||
#include <opencv2/imgproc.hpp>
|
||||
#include <opencv2/imgcodecs.hpp>
|
||||
#include <opencv2/dnn.hpp>
|
||||
#include <vector>
|
||||
#include <string>
|
||||
#include <cstring>
|
||||
|
||||
using namespace cv;
|
||||
using namespace cv::dnn;
|
||||
|
||||
// 图像结构体定义
|
||||
typedef struct {
|
||||
unsigned char* data;
|
||||
@@ -15,6 +17,19 @@ typedef struct {
|
||||
|
||||
static std::string last_error;
|
||||
|
||||
// YOLO 检测器
|
||||
class YOLODetector {
|
||||
public:
|
||||
Net net;
|
||||
std::vector<std::string> classNames;
|
||||
float confThreshold;
|
||||
float nmsThreshold;
|
||||
|
||||
YOLODetector() : confThreshold(0.5f), nmsThreshold(0.4f) {}
|
||||
};
|
||||
|
||||
static std::map<std::string, YOLODetector*> yolo_detectors;
|
||||
|
||||
extern "C" {
|
||||
|
||||
// 图像解码
|
||||
@@ -53,6 +68,175 @@ void cv_image_free(Image* img) {
|
||||
}
|
||||
}
|
||||
|
||||
// 加载 YOLO ONNX 模型
|
||||
int cv_yolo_load(const char* name, const char* model_path, const char* classes_path) {
|
||||
try {
|
||||
YOLODetector* detector = new YOLODetector();
|
||||
|
||||
// 加载 ONNX 模型
|
||||
detector->net = readNetFromONNX(model_path);
|
||||
if (detector->net.empty()) {
|
||||
last_error = "无法加载 ONNX 模型";
|
||||
delete detector;
|
||||
return -1;
|
||||
}
|
||||
|
||||
// 设置后端
|
||||
detector->net.setPreferableBackend(DNN_BACKEND_OPENCV);
|
||||
detector->net.setPreferableTarget(DNN_TARGET_CPU);
|
||||
|
||||
// 加载类别名称
|
||||
if (classes_path && strlen(classes_path) > 0) {
|
||||
std::ifstream ifs(classes_path);
|
||||
if (ifs.is_open()) {
|
||||
std::string line;
|
||||
while (std::getline(ifs, line)) {
|
||||
if (!line.empty()) {
|
||||
detector->classNames.push_back(line);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
yolo_detectors[std::string(name)] = detector;
|
||||
return 0;
|
||||
} catch (const std::exception& e) {
|
||||
last_error = e.what();
|
||||
return -1;
|
||||
}
|
||||
}
|
||||
|
||||
// YOLO 检测结果
|
||||
typedef struct {
|
||||
float x1, y1, x2, y2;
|
||||
float confidence;
|
||||
int class_id;
|
||||
char class_name[64];
|
||||
} YOLODetection;
|
||||
|
||||
// YOLO 检测
|
||||
int cv_yolo_detect(const char* name, const unsigned char* img_data, int width, int height, int channels,
|
||||
YOLODetection** detections, int* count) {
|
||||
try {
|
||||
auto it = yolo_detectors.find(std::string(name));
|
||||
if (it == yolo_detectors.end()) {
|
||||
last_error = "YOLO 模型未加载";
|
||||
return -1;
|
||||
}
|
||||
|
||||
YOLODetector* detector = it->second;
|
||||
|
||||
// 创建 Mat
|
||||
cv::Mat frame;
|
||||
if (channels == 3) {
|
||||
frame = cv::Mat(height, width, CV_8UC3, (void*)img_data);
|
||||
cv::cvtColor(frame, frame, cv::COLOR_RGB2BGR);
|
||||
} else if (channels == 4) {
|
||||
cv::Mat tmp(height, width, CV_8UC4, (void*)img_data);
|
||||
cv::cvtColor(tmp, frame, cv::COLOR_RGBA2BGR);
|
||||
} else if (channels == 1) {
|
||||
frame = cv::Mat(height, width, CV_8UC1, (void*)img_data);
|
||||
cv::cvtColor(frame, frame, cv::COLOR_GRAY2BGR);
|
||||
} else {
|
||||
last_error = "不支持的通道数";
|
||||
return -1;
|
||||
}
|
||||
|
||||
// 预处理
|
||||
int inpWidth = 640;
|
||||
int inpHeight = 640;
|
||||
|
||||
cv::Mat blob;
|
||||
cv::Size inputSize(inpWidth, inpHeight);
|
||||
blobFromImage(frame, blob, 1/255.0, inputSize, Scalar(0,0,0), true, false);
|
||||
|
||||
// 推理
|
||||
detector->net.setInput(blob);
|
||||
std::vector<Mat> outs;
|
||||
detector->net.forward(outs, detector->net.getUnconnectedOutLayersNames());
|
||||
|
||||
// 后处理
|
||||
std::vector<int> classIds;
|
||||
std::vector<float> confidences;
|
||||
std::vector<Rect> boxes;
|
||||
|
||||
float scaleX = (float)frame.cols / inpWidth;
|
||||
float scaleY = (float)frame.rows / inpHeight;
|
||||
|
||||
for (size_t i = 0; i < outs.size(); ++i) {
|
||||
float* data = (float*)outs[i].data;
|
||||
for (int j = 0; j < outs[i].rows; ++j, data += outs[i].cols) {
|
||||
Mat scores = outs[i].row(j).colRange(5, outs[i].cols);
|
||||
Point classIdPoint;
|
||||
double confidence;
|
||||
minMaxLoc(scores, 0, &confidence, 0, &classIdPoint);
|
||||
|
||||
if (confidence > detector->confThreshold) {
|
||||
int centerX = (int)(data[0] * scaleX);
|
||||
int centerY = (int)(data[1] * scaleY);
|
||||
int width = (int)(data[2] * scaleX);
|
||||
int height = (int)(data[3] * scaleY);
|
||||
|
||||
int left = centerX - width / 2;
|
||||
int top = centerY - height / 2;
|
||||
|
||||
classIds.push_back(classIdPoint.x);
|
||||
confidences.push_back((float)confidence);
|
||||
boxes.push_back(Rect(left, top, width, height));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NMS
|
||||
std::vector<int> indices;
|
||||
NMSBoxes(boxes, confidences, detector->confThreshold, detector->nmsThreshold, indices);
|
||||
|
||||
// 分配结果
|
||||
*count = (int)indices.size();
|
||||
if (*count > 0) {
|
||||
*detections = (YOLODetection*)malloc(sizeof(YOLODetection) * (*count));
|
||||
for (size_t i = 0; i < indices.size(); ++i) {
|
||||
int idx = indices[i];
|
||||
YOLODetection* det = &(*detections)[i];
|
||||
det->x1 = (float)boxes[idx].x;
|
||||
det->y1 = (float)boxes[idx].y;
|
||||
det->x2 = (float)(boxes[idx].x + boxes[idx].width);
|
||||
det->y2 = (float)(boxes[idx].y + boxes[idx].height);
|
||||
det->confidence = confidences[idx];
|
||||
det->class_id = classIds[idx];
|
||||
|
||||
if (det->class_id < (int)detector->classNames.size()) {
|
||||
strncpy(det->class_name, detector->classNames[det->class_id].c_str(), 63);
|
||||
det->class_name[63] = '\0';
|
||||
} else {
|
||||
sprintf(det->class_name, "class_%d", det->class_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return 0;
|
||||
} catch (const std::exception& e) {
|
||||
last_error = e.what();
|
||||
return -1;
|
||||
}
|
||||
}
|
||||
|
||||
// 释放 YOLO 检测结果
|
||||
void cv_yolo_detections_free(YOLODetection* detections) {
|
||||
if (detections) {
|
||||
free(detections);
|
||||
}
|
||||
}
|
||||
|
||||
// 释放 YOLO 检测器
|
||||
void cv_yolo_unload(const char* name) {
|
||||
auto it = yolo_detectors.find(std::string(name));
|
||||
if (it != yolo_detectors.end()) {
|
||||
delete it->second;
|
||||
yolo_detectors.erase(it);
|
||||
}
|
||||
}
|
||||
|
||||
// 滑块缺口匹配
|
||||
int cv_slider_match(const Image* target, const Image* background, int* out_x) {
|
||||
try {
|
||||
@@ -63,37 +247,10 @@ int cv_slider_match(const Image* target, const Image* background, int* out_x) {
|
||||
cv::cvtColor(target_mat, target_gray, cv::COLOR_BGR2GRAY);
|
||||
cv::cvtColor(bg_mat, bg_gray, cv::COLOR_BGR2GRAY);
|
||||
|
||||
// 模板匹配
|
||||
cv::Mat result;
|
||||
cv::matchTemplate(bg_gray, target_gray, result, cv::TM_CCOEFF_NORMED);
|
||||
|
||||
double min_val, max_val;
|
||||
cv::Point min_loc, max_loc;
|
||||
cv::minMaxLoc(result, &min_val, &max_val, &min_loc, &max_loc);
|
||||
|
||||
*out_x = max_loc.x;
|
||||
return 0;
|
||||
} catch (const std::exception& e) {
|
||||
last_error = e.what();
|
||||
return -1;
|
||||
}
|
||||
}
|
||||
|
||||
// 阴影滑块匹配
|
||||
int cv_slider_comparison(const Image* target, const Image* background, int* out_x) {
|
||||
try {
|
||||
cv::Mat target_mat(target->height, target->width, CV_8UC3, target->data);
|
||||
cv::Mat bg_mat(background->height, background->width, CV_8UC3, background->data);
|
||||
|
||||
// 转灰度
|
||||
cv::Mat target_gray, bg_gray;
|
||||
cv::cvtColor(target_mat, target_gray, cv::COLOR_BGR2GRAY);
|
||||
cv::cvtColor(bg_mat, bg_gray, cv::COLOR_BGR2GRAY);
|
||||
|
||||
// Canny 边缘检测
|
||||
cv::Mat target_edges, bg_edges;
|
||||
cv::Canny(target_gray, target_edges, 50, 150);
|
||||
cv::Canny(bg_gray, bg_edges, 50, 150);
|
||||
cv::Canny(target_gray, target_edges, 100, 200);
|
||||
cv::Canny(bg_gray, bg_edges, 100, 200);
|
||||
|
||||
// 模板匹配
|
||||
cv::Mat result;
|
||||
@@ -111,7 +268,52 @@ int cv_slider_comparison(const Image* target, const Image* background, int* out_
|
||||
}
|
||||
}
|
||||
|
||||
// 检测旋转角度
|
||||
// 阴影滑块匹配
|
||||
int cv_slider_comparison(const Image* target, const Image* background, int* out_x, int* out_y) {
|
||||
try {
|
||||
cv::Mat target_mat(target->height, target->width, CV_8UC3, target->data);
|
||||
cv::Mat bg_mat(background->height, background->width, CV_8UC3, background->data);
|
||||
|
||||
// 计算差异
|
||||
cv::Mat diff;
|
||||
cv::absdiff(bg_mat, target_mat, diff);
|
||||
|
||||
// 阈值处理
|
||||
cv::Mat thresh;
|
||||
cv::threshold(diff, thresh, 30, 255, cv::THRESH_BINARY);
|
||||
|
||||
// 找到差异区域
|
||||
std::vector<std::vector<cv::Point>> contours;
|
||||
cv::findContours(thresh, contours, cv::RETR_EXTERNAL, cv::CHAIN_APPROX_SIMPLE);
|
||||
|
||||
if (contours.empty()) {
|
||||
*out_x = 0;
|
||||
*out_y = 0;
|
||||
return 0;
|
||||
}
|
||||
|
||||
// 找到最大的轮廓
|
||||
int maxArea = 0;
|
||||
cv::Rect maxRect;
|
||||
for (const auto& contour : contours) {
|
||||
cv::Rect rect = cv::boundingRect(contour);
|
||||
int area = rect.width * rect.height;
|
||||
if (area > maxArea) {
|
||||
maxArea = area;
|
||||
maxRect = rect;
|
||||
}
|
||||
}
|
||||
|
||||
*out_x = maxRect.x;
|
||||
*out_y = maxRect.y;
|
||||
return 0;
|
||||
} catch (const std::exception& e) {
|
||||
last_error = e.what();
|
||||
return -1;
|
||||
}
|
||||
}
|
||||
|
||||
// 检测旋转角度(简化版,使用特征点)
|
||||
float cv_detect_rotation(const Image* img) {
|
||||
try {
|
||||
cv::Mat mat(img->height, img->width, CV_8UC3, img->data);
|
||||
@@ -120,18 +322,6 @@ float cv_detect_rotation(const Image* img) {
|
||||
cv::Mat gray;
|
||||
cv::cvtColor(mat, gray, cv::COLOR_BGR2GRAY);
|
||||
|
||||
// 使用霍夫圆变换检测圆心
|
||||
cv::Mat blurred;
|
||||
cv::GaussianBlur(gray, blurred, cv::Size(5, 5), 0);
|
||||
|
||||
std::vector<cv::Vec3f> circles;
|
||||
cv::HoughCircles(blurred, circles, cv::HOUGH_GRADIENT, 1,
|
||||
blurred.rows / 8, 100, 30, 0, 0);
|
||||
|
||||
if (circles.empty()) {
|
||||
return 0.0f;
|
||||
}
|
||||
|
||||
// 简化处理:返回 0 度
|
||||
// 实际实现需要更复杂的特征点匹配
|
||||
return 0.0f;
|
||||
|
||||
+142
-67
@@ -1,49 +1,39 @@
|
||||
package opencv
|
||||
|
||||
/*
|
||||
#cgo pkg-config: opencv4
|
||||
#cgo CXXFLAGS: -std=c++17
|
||||
#cgo CXXFLAGS: -std=c++17 -I/usr/include/opencv4
|
||||
#cgo linux LDFLAGS: -L/usr/lib/x86_64-linux-gnu -lopencv_core -lopencv_imgproc -lopencv_imgcodecs -lopencv_dnn -lopencv_calib3d -lstdc++
|
||||
#cgo darwin LDFLAGS: -lopencv_core -lopencv_imgproc -lopencv_imgcodecs -lopencv_dnn -lopencv_calib3d -lstdc++
|
||||
|
||||
#include <stdlib.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// 图像结构
|
||||
typedef struct {
|
||||
// 图像结构体 - 在这里定义让 Go 可以访问
|
||||
typedef struct Image {
|
||||
unsigned char* data;
|
||||
int width;
|
||||
int height;
|
||||
int channels;
|
||||
} Image;
|
||||
|
||||
// 图像操作
|
||||
typedef struct {
|
||||
float x1, y1, x2, y2;
|
||||
float confidence;
|
||||
int class_id;
|
||||
char class_name[64];
|
||||
} YOLODetection;
|
||||
|
||||
// OpenCV 函数
|
||||
Image* cv_imdecode(const unsigned char* buf, size_t size);
|
||||
void cv_image_free(Image* img);
|
||||
|
||||
// 滑块匹配
|
||||
int cv_yolo_load(const char* name, const char* model_path, const char* classes_path);
|
||||
int cv_yolo_detect(const char* name, const unsigned char* img_data, int width, int height, int channels, YOLODetection** detections, int* count);
|
||||
void cv_yolo_detections_free(YOLODetection* detections);
|
||||
void cv_yolo_unload(const char* name);
|
||||
int cv_slider_match(const Image* target, const Image* background, int* out_x);
|
||||
int cv_slider_comparison(const Image* target, const Image* background, int* out_x);
|
||||
|
||||
// 旋转检测
|
||||
int cv_slider_comparison(const Image* target, const Image* background, int* out_x, int* out_y);
|
||||
float cv_detect_rotation(const Image* img);
|
||||
|
||||
// 模板匹配
|
||||
int cv_template_match(const Image* src, const Image* templ, double* max_val, int* max_x, int* max_y);
|
||||
|
||||
// 特征点检测
|
||||
int cv_detect_features(const Image* img, int** points_x, int** points_y, int* count);
|
||||
|
||||
// 图像相似度
|
||||
float cv_compare_similarity(const Image* img1, const Image* img2);
|
||||
|
||||
// 错误信息
|
||||
const char* cv_get_last_error();
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
*/
|
||||
import "C"
|
||||
import (
|
||||
@@ -53,23 +43,27 @@ import (
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
// Image 封装图像数据
|
||||
// Image OpenCV 图像
|
||||
type Image struct {
|
||||
img *C.Image
|
||||
}
|
||||
|
||||
// YOLODetection YOLO 检测结果
|
||||
type YOLODetection struct {
|
||||
X1, Y1, X2, Y2 float32
|
||||
Confidence float32
|
||||
ClassID int
|
||||
ClassName string
|
||||
}
|
||||
|
||||
// DecodeFromBase64 从 Base64 解码图像
|
||||
func DecodeFromBase64(data string) (*Image, error) {
|
||||
decoded, err := base64.StdEncoding.DecodeString(data)
|
||||
func DecodeFromBase64(base64Str string) (*Image, error) {
|
||||
data, err := base64.StdEncoding.DecodeString(base64Str)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("base64 解码失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
img := C.cv_imdecode(
|
||||
(*C.uchar)(unsafe.Pointer(&decoded[0])),
|
||||
C.size_t(len(decoded)),
|
||||
)
|
||||
|
||||
img := C.cv_imdecode((*C.uchar)(unsafe.Pointer(&data[0])), C.size_t(len(data)))
|
||||
if img == nil {
|
||||
return nil, errors.New(C.GoString(C.cv_get_last_error()))
|
||||
}
|
||||
@@ -77,7 +71,7 @@ func DecodeFromBase64(data string) (*Image, error) {
|
||||
return &Image{img: img}, nil
|
||||
}
|
||||
|
||||
// Free 释放图像内存
|
||||
// Free 释放图像
|
||||
func (i *Image) Free() {
|
||||
if i.img != nil {
|
||||
C.cv_image_free(i.img)
|
||||
@@ -95,51 +89,132 @@ func (i *Image) Height() int {
|
||||
return int(i.img.height)
|
||||
}
|
||||
|
||||
// SliderMatch 滑块缺口匹配
|
||||
func SliderMatch(target, background *Image) (int, error) {
|
||||
var outX C.int
|
||||
|
||||
result := C.cv_slider_match(
|
||||
(*C.Image)(target.img),
|
||||
(*C.Image)(background.img),
|
||||
&outX,
|
||||
)
|
||||
// Channels 获取通道数
|
||||
func (i *Image) Channels() int {
|
||||
return int(i.img.channels)
|
||||
}
|
||||
|
||||
if result != 0 {
|
||||
return 0, errors.New(C.GoString(C.cv_get_last_error()))
|
||||
// Data 获取图像数据
|
||||
func (i *Image) Data() []byte {
|
||||
size := int(i.img.width) * int(i.img.height) * int(i.img.channels)
|
||||
return C.GoBytes(unsafe.Pointer(i.img.data), C.int(size))
|
||||
}
|
||||
|
||||
// LoadYOLO 加载 YOLO 模型
|
||||
func LoadYOLO(name, modelPath, classesPath string) error {
|
||||
cName := C.CString(name)
|
||||
cModelPath := C.CString(modelPath)
|
||||
defer C.free(unsafe.Pointer(cName))
|
||||
defer C.free(unsafe.Pointer(cModelPath))
|
||||
|
||||
var cClassesPath *C.char
|
||||
if classesPath != "" {
|
||||
cClassesPath = C.CString(classesPath)
|
||||
defer C.free(unsafe.Pointer(cClassesPath))
|
||||
}
|
||||
|
||||
return int(outX), nil
|
||||
ret := C.cv_yolo_load(cName, cModelPath, cClassesPath)
|
||||
if ret != 0 {
|
||||
return fmt.Errorf("加载 YOLO 模型失败: %s", C.GoString(C.cv_get_last_error()))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DetectYOLO YOLO 检测
|
||||
func DetectYOLO(name string, img *Image) ([]YOLODetection, error) {
|
||||
if img == nil || img.img == nil {
|
||||
return nil, errors.New("图像为空")
|
||||
}
|
||||
|
||||
var detections *C.YOLODetection
|
||||
var count C.int
|
||||
|
||||
cName := C.CString(name)
|
||||
defer C.free(unsafe.Pointer(cName))
|
||||
|
||||
ret := C.cv_yolo_detect(cName, img.img.data, img.img.width, img.img.height, img.img.channels,
|
||||
&detections, &count)
|
||||
if ret != 0 {
|
||||
return nil, fmt.Errorf("YOLO 检测失败: %s", C.GoString(C.cv_get_last_error()))
|
||||
}
|
||||
|
||||
if count == 0 {
|
||||
return []YOLODetection{}, nil
|
||||
}
|
||||
|
||||
defer C.cv_yolo_detections_free(detections)
|
||||
|
||||
// 转换为 Go 类型
|
||||
detectionSlice := (*[1 << 20]C.YOLODetection)(unsafe.Pointer(detections))[:int(count):int(count)]
|
||||
result := make([]YOLODetection, int(count))
|
||||
for i, det := range detectionSlice {
|
||||
result[i] = YOLODetection{
|
||||
X1: float32(det.x1),
|
||||
Y1: float32(det.y1),
|
||||
X2: float32(det.x2),
|
||||
Y2: float32(det.y2),
|
||||
Confidence: float32(det.confidence),
|
||||
ClassID: int(det.class_id),
|
||||
ClassName: C.GoString(&det.class_name[0]),
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// UnloadYOLO 卸载 YOLO 模型
|
||||
func UnloadYOLO(name string) {
|
||||
cName := C.CString(name)
|
||||
defer C.free(unsafe.Pointer(cName))
|
||||
C.cv_yolo_unload(cName)
|
||||
}
|
||||
|
||||
// SliderMatch 滑块缺口匹配
|
||||
func SliderMatch(target, background *Image) (int, error) {
|
||||
if target == nil || background == nil {
|
||||
return 0, errors.New("图像为空")
|
||||
}
|
||||
|
||||
var x C.int
|
||||
ret := C.cv_slider_match(target.img, background.img, &x)
|
||||
if ret != 0 {
|
||||
return 0, fmt.Errorf("滑块匹配失败: %s", C.GoString(C.cv_get_last_error()))
|
||||
}
|
||||
|
||||
return int(x), nil
|
||||
}
|
||||
|
||||
// SliderComparison 阴影滑块匹配
|
||||
func SliderComparison(target, background *Image) (int, error) {
|
||||
var outX C.int
|
||||
|
||||
result := C.cv_slider_comparison(
|
||||
(*C.Image)(target.img),
|
||||
(*C.Image)(background.img),
|
||||
&outX,
|
||||
)
|
||||
|
||||
if result != 0 {
|
||||
return 0, errors.New(C.GoString(C.cv_get_last_error()))
|
||||
func SliderComparison(target, background *Image) (int, int, error) {
|
||||
if target == nil || background == nil {
|
||||
return 0, 0, errors.New("图像为空")
|
||||
}
|
||||
|
||||
return int(outX), nil
|
||||
var x, y C.int
|
||||
ret := C.cv_slider_comparison(target.img, background.img, &x, &y)
|
||||
if ret != 0 {
|
||||
return 0, 0, fmt.Errorf("阴影滑块匹配失败: %s", C.GoString(C.cv_get_last_error()))
|
||||
}
|
||||
|
||||
return int(x), int(y), nil
|
||||
}
|
||||
|
||||
// DetectRotation 检测旋转角度
|
||||
func DetectRotation(img *Image) (float32, error) {
|
||||
angle := C.cv_detect_rotation((*C.Image)(img.img))
|
||||
if img == nil {
|
||||
return 0, errors.New("图像为空")
|
||||
}
|
||||
|
||||
angle := C.cv_detect_rotation(img.img)
|
||||
return float32(angle), nil
|
||||
}
|
||||
|
||||
// CompareSimilarity 比较图像相似度
|
||||
func CompareSimilarity(img1, img2 *Image) (float32, error) {
|
||||
similarity := C.cv_compare_similarity(
|
||||
(*C.Image)(img1.img),
|
||||
(*C.Image)(img2.img),
|
||||
)
|
||||
if img1 == nil || img2 == nil {
|
||||
return 0, errors.New("图像为空")
|
||||
}
|
||||
|
||||
similarity := C.cv_compare_similarity(img1.img, img2.img)
|
||||
return float32(similarity), nil
|
||||
}
|
||||
Reference in New Issue
Block a user