new train

This commit is contained in:
2026-07-28 18:06:10 +08:00
parent 44b0d6b756
commit a55d449028
13 changed files with 224 additions and 27 deletions
BIN
View File
Binary file not shown.
+160 -3
View File
@@ -1,19 +1,39 @@
package backend
import (
"bytes"
"context"
"encoding/base64"
"fmt"
"math"
"os"
"path/filepath"
"time"
"github.com/disintegration/imaging"
ort "github.com/yalue/onnxruntime_go"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
var (
publicImagePath = "./frontend/public/images"
modelPath = "resnet_epoch_100.onnx"
// ImageNet 标准化参数
mean = []float32{0.485, 0.456, 0.406}
std = []float32{0.229, 0.224, 0.225}
// 猫品种标签
labelName = []string{
"american_shorthair",
"bengal",
"british_shorthair",
"exotic_shorthair",
"maine_coon",
"ragdoll",
"sphynx",
}
)
func NewApp() *App {
@@ -22,6 +42,113 @@ func NewApp() *App {
func (a *App) Startup(ctx context.Context) {
a.ctx = ctx
// 设置 ONNX Runtime DLL 路径
// ort.SetSharedLibraryPath("onnxruntime.dll")
exeDir, _ := os.Executable()
println("exeDir", exeDir)
ort.SetSharedLibraryPath(filepath.Join(filepath.Dir(exeDir), "onnxruntime.dll"))
// 初始化 ONNX Runtime 环境
err := ort.InitializeEnvironment()
if err != nil {
panic("Failed to initialize ONNX runtime: " + err.Error())
}
}
func preprocessImage(imgData []byte) ([]float32, error) {
// 解码图片
reader := bytes.NewReader(imgData)
img, err := imaging.Decode(reader, imaging.AutoOrientation(true))
if err != nil {
return nil, fmt.Errorf("decode image: %w", err)
}
// 缩放到 224x224
img = imaging.Resize(img, 224, 224, imaging.Lanczos)
// 转换为 float32 数组 (NCHW 格式: 1, 3, 224, 224)
input := make([]float32, 1*3*224*224)
bounds := img.Bounds()
idx := 0
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()
// RGBA 返回 0-65535,需要转换到 0-255
rf := float32(r>>8) / 255.0
gf := float32(g>>8) / 255.0
bf := float32(b>>8) / 255.0
// ImageNet 标准化
input[idx] = (rf - mean[0]) / std[0] // R channel
input[idx+224*224] = (gf - mean[1]) / std[1] // G channel
input[idx+224*224*2] = (bf - mean[2]) / std[2] // B channel
idx++
}
}
return input, nil
}
func runInference(input []float32) (string, float64, error) {
// 创建输入张量 [1, 3, 224, 224]
inputTensor, err := ort.NewTensor(ort.Shape{1, 3, 224, 224}, input)
if err != nil {
return "", 0, fmt.Errorf("create input tensor: %w", err)
}
defer inputTensor.Destroy()
// 创建输出张量 [1, 7]
outputTensor, err := ort.NewTensor(ort.Shape{1, 7}, make([]float32, 7))
if err != nil {
return "", 0, fmt.Errorf("create output tensor: %w", err)
}
defer outputTensor.Destroy()
exeDir, _ := os.Executable()
println("exeDir222", exeDir)
// 创建 session 并运行推理
session, err := ort.NewAdvancedSession(
filepath.Join(filepath.Dir(exeDir), modelPath),
[]string{"input"},
[]string{"output"},
[]ort.Value{inputTensor},
[]ort.Value{outputTensor},
nil,
)
if err != nil {
return "", 0, fmt.Errorf("create session: %w", err)
}
defer session.Destroy()
err = session.Run()
if err != nil {
return "", 0, fmt.Errorf("run inference: %w", err)
}
// 获取输出
outputData := outputTensor.GetData()
// 找最大值的索引
maxIdx := 0
maxVal := outputData[0]
for i := 1; i < len(outputData); i++ {
if outputData[i] > maxVal {
maxVal = outputData[i]
maxIdx = i
}
}
// Softmax 计算置信度
var sum float64
for _, v := range outputData {
sum += math.Exp(float64(v))
}
confidence := math.Exp(float64(maxVal)) / sum
return labelName[maxIdx], confidence, nil
}
func (a *App) GormDB() (*gorm.DB, error) {
@@ -33,6 +160,7 @@ func (a *App) GormDB() (*gorm.DB, error) {
}
func (a *App) UploadImage(data []byte, filename string) Response {
println("UploadImage")
uploadsDir := publicImagePath
err := os.MkdirAll(uploadsDir, 0755)
if err != nil {
@@ -48,6 +176,8 @@ func (a *App) UploadImage(data []byte, filename string) Response {
return Response{Code: 1, Message: "failed", Data: err.Error()}
}
println("333")
return Response{Code: 0, Message: "success", Data: newFilename}
}
@@ -94,14 +224,41 @@ func (a *App) GetHistory(page int, pageSize int) Response {
return Response{Code: 0, Message: "success", Data: historyData}
}
func (a *App) Detect(img string) Response {
func (a *App) Detect(filename string) Response {
db, err := a.GormDB()
if err != nil {
return Response{Code: 1, Message: err.Error()}
}
/**
这里调用模型
**/
filePath := filepath.Join(publicImagePath, filename)
imgData, err := os.ReadFile(filePath)
if err != nil {
return Response{Code: 1, Message: "failed to read image: " + err.Error()}
}
input, err := preprocessImage(imgData)
if err != nil {
return Response{Code: 1, Message: "failed to preprocess: " + err.Error()}
}
// 推理
detectRet, confidence, err := runInference(input)
if err != nil {
return Response{Code: 1, Message: "model inference failed: " + err.Error()}
}
println("confidence", confidence)
println("detectRet", detectRet)
/**
结束
**/
var breed Breed
detectRet := "british_shorthair"
err = db.Table("breeds_test").Where("code = ?", detectRet).First(&breed).Error
if err != nil {
return Response{Code: 1, Message: err.Error()}
@@ -109,7 +266,7 @@ func (a *App) Detect(img string) Response {
now := int(time.Now().Unix())
one := HistoryItem{Img: img, Breed: int(breed.Id), Date: now}
one := HistoryItem{Img: filename, Breed: int(breed.Id), Date: now}
result := db.Table("history_test").Create(&one)
if result.Error != nil {
return Response{Code: 1, Message: result.Error.Error()}
+6 -3
View File
@@ -1,4 +1,4 @@
mode = "toy"
mode = "benchmark"
"""
epoch 训练多少轮
@@ -27,8 +27,8 @@ if mode == "toy":
num_classes = len(label_name)
elif mode == "benchmark":
epoch = 50
lr = 2e-4
batch_size = 2
lr = 1e-4
batch_size = 8
input_size = 224
# 分类
@@ -39,6 +39,9 @@ elif mode == "benchmark":
"exotic_shorthair",
"maine_coon",
"ragdoll",
"scottish_fold",
"siamese",
"sphynx",
"turkish_van",
]
num_classes = len(label_name)
+2 -1
View File
@@ -23,7 +23,8 @@ train_transform = transforms.Compose([
transforms.RandomRotation(10), # 轻微旋转
transforms.ColorJitter(brightness=0.1, contrast=0.1),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
transforms.RandomErasing(p=0.5, scale=(0.02, 0.2), ratio=(0.3, 3.3)),
])
Binary file not shown.
@@ -1,10 +1,12 @@
import os
import cv2
import glob
import torch
import numpy as np
from PIL import Image
import onnxruntime as ort # 【修改1】替换torch导入
from torchvision import transforms # 保留transforms,仍用于预处理
from core.const import label_name, input_size
from core.const import mode, label_name, input_size
def test():
@@ -12,13 +14,18 @@ def test():
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if 'CUDAExecutionProvider' in ort.get_available_providers() else ['CPUExecutionProvider']
print(f"使用设备: {providers[0]}")
script_dir = os.path.dirname(os.path.abspath(__file__))
core_dir = os.path.dirname(script_dir)
model_dir = os.path.join(core_dir, "models")
dataset_dir = os.path.join(core_dir, "dataset", mode, "test")
# 【修改3】加载ONNX模型,替代原来的PyTorch模型加载
session = ort.InferenceSession("./model/resnet_final.onnx", providers=providers)
session = ort.InferenceSession(os.path.join(model_dir, "resnet18_epoch_50_bak2.onnx"), providers=providers)
# 获取输入名称(用于后续推理时指定输入)
input_name = session.get_inputs()[0].name
im_list = glob.glob("./dataset/test/*/*.jpg")
im_list = glob.glob(os.path.join(dataset_dir, "*", "*.jpg"))
np.random.shuffle(im_list)
# 预处理完全不变
@@ -49,6 +56,7 @@ def test():
img = np.asarray(im_data)
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
img = cv2.resize(img, (200, int(img.shape[0] * 200 / img.shape[1])))
cv2.imshow("img", img)
cv2.waitKey(0)
+16 -5
View File
@@ -1,3 +1,4 @@
import os
import cv2
import glob
import torch
@@ -5,17 +6,26 @@ from torchvision import transforms
from PIL import Image
import numpy as np
from core.nets.resnet18 import resnet18
from core.const import label_name, input_size
from core.const import mode, label_name, input_size
def test():
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(device)
net = resnet18()
net.load_state_dict(torch.load("./model/resnet_epoch_14.pth", weights_only=True))
script_dir = os.path.dirname(os.path.abspath(__file__))
core_dir = os.path.dirname(script_dir)
model_dir = os.path.join(core_dir, "models")
dataset_dir = os.path.join(core_dir, "dataset", mode, "test")
im_list = glob.glob("./dataset/test/*/*.jpg")
print("model_dir", model_dir)
net = resnet18()
# net.load_state_dict(torch.load("./models/resnet18_epoch_100.pth", weights_only=True))
net.load_state_dict(torch.load(os.path.join(model_dir, "resnet18_epoch_50_bak2.pth"), weights_only=True))
# im_list = glob.glob("./dataset/test/*/*.jpg")
im_list = glob.glob(os.path.join(dataset_dir, "*", "*.jpg"))
np.random.shuffle(im_list)
net.to(device)
@@ -41,12 +51,13 @@ def test():
_, pred = torch.max(outputs.data, dim=1)
print(label_name[pred.cpu().numpy()[0]], " ", im_path)
# prob, pred = torch.topk(outputs.data, k=3, dim=1)
prob, pred = torch.topk(outputs.data, k=3, dim=1)
# for i in range(3):
# print(label_name[pred[0, i].item()], " ", prob[0, i].item(), " ", im_path)
img = np.asarray(im_data)
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
img = cv2.resize(img, (200, 200))
cv2.imshow("img", img)
cv2.waitKey(0)
@@ -1,24 +1,29 @@
import os
import torch
import sys
from core.nets.resnet import resnet
# from core.nets.resnet import resnet
from core.nets.resnet18 import resnet18
# 加载 pth
net = resnet() # 实例化你的模型
# net = resnet() # 实例化你的模型
net = resnet18()
script_dir = os.path.dirname(os.path.abspath(__file__))
core_dir = os.path.dirname(script_dir)
model_dir = os.path.join(core_dir, "models")
net.load_state_dict(torch.load(os.path.join(model_dir, "resnet_epoch_100.pth"), map_location="cpu"))
net.load_state_dict(torch.load(os.path.join(model_dir, "resnet18_epoch_50_bak2.pth"), map_location="cpu"))
net.eval()
# 导出 ONNX
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
net,
dummy_input,
os.path.join(model_dir, "resnet_epoch_100.onnx"),
os.path.join(model_dir, "resnet18_epoch_50_bak2.onnx"),
export_params=True,
opset_version=11,
input_names=["input"],
+4 -4
View File
@@ -27,7 +27,7 @@ const History = () => {
const [showClear, setShowClear] = useState<boolean>(false)
useEffect(() => {
if (!(window as any).go?.main?.App?.GetHistory) {
if (!(window as any).go?.backend?.App?.GetHistory) {
message.error('Wails runtime not ready')
return
}
@@ -35,7 +35,7 @@ const History = () => {
}, [])
const fetchData = async(page: number) => {
const result = await (window as any).go.main.App.GetHistory(page, pageSize)
const result = await (window as any).go.backend.App.GetHistory(page, pageSize)
if (result.code === 0) {
setHistoryList(result.data.list)
setTotal(result.data.total)
@@ -60,7 +60,7 @@ const History = () => {
message.error('currentId为空')
return
}
const result = await (window as any).go.main.App.DeleteOneHistory(currentId)
const result = await (window as any).go.backend.App.DeleteOneHistory(currentId)
if (result.code === 0) {
message.success('删除成功')
setCurrentId(null)
@@ -79,7 +79,7 @@ const History = () => {
const handleClearOk = async() => {
setShowClear(false)
const result = await (window as any).go.main.App.ClearHistory()
const result = await (window as any).go.backend.App.ClearHistory()
if (result.code === 0) {
message.success('已清空历史')
setCurrentId(null)
+1 -1
View File
@@ -11,7 +11,7 @@ const ImagePreview = ({ filename, ...props }: Props) => {
useEffect(() => {
const loadImage = async () => {
const result = await (window as any).go.main.App.GetImage(filename)
const result = await (window as any).go.backend.App.GetImage(filename)
if (result.code === 0) {
setSrc(result.data)
}
+3 -2
View File
@@ -60,7 +60,7 @@ const Main = () => {
const arrayBuffer = await processedFile.arrayBuffer()
const uint8Array = new Uint8Array(arrayBuffer)
const result = await (window as any).go.main.App.UploadImage(
const result = await (window as any).go.backend.App.UploadImage(
Array.from(uint8Array),
file.name
)
@@ -92,11 +92,12 @@ const Main = () => {
setStep(2)
setTimeout(async() => {
const result = await (window as any).go.main.App.Detect(fileSrc)
const result = await (window as any).go.backend.App.Detect(fileSrc)
if (result.code === 0) {
setDetectResult(result.data)
setStep(3)
} else if (result.code === 1) {
console.log("result.message", result.message)
message.error(result.message)
}
}, 1000)
+3
View File
@@ -3,7 +3,9 @@ module sortmeow
go 1.25.0
require (
github.com/disintegration/imaging v1.6.2
github.com/wailsapp/wails/v2 v2.13.0
github.com/yalue/onnxruntime_go v1.31.0
gorm.io/driver/sqlite v1.6.0
gorm.io/gorm v1.31.2
)
@@ -37,6 +39,7 @@ require (
github.com/wailsapp/go-webview2 v1.0.22 // indirect
github.com/wailsapp/mimetype v1.4.1 // indirect
golang.org/x/crypto v0.51.0 // indirect
golang.org/x/image v0.40.0 // indirect
golang.org/x/net v0.54.0 // indirect
golang.org/x/sys v0.46.0 // indirect
golang.org/x/text v0.40.0 // indirect
+8
View File
@@ -4,6 +4,8 @@ github.com/bep/debounce v1.2.1 h1:v67fRdBA9UQu2NhLFXrSg0Brw7CexQekrBwDMM8bzeY=
github.com/bep/debounce v1.2.1/go.mod h1:H8yggRPQKLUhUoqrJC1bO2xNya7vanpDl7xR3ISbCJ0=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/disintegration/imaging v1.6.2 h1:w1LecBlG2Lnp8B3jk5zSuNqd7b4DXhcjwek1ei82L+c=
github.com/disintegration/imaging v1.6.2/go.mod h1:44/5580QXChDfwIclfc/PCwrr44amcmDAg8hxG0Ewe4=
github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE=
github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78=
github.com/godbus/dbus/v5 v5.1.0 h1:4KLkAxT3aOY8Li4FRJe/KvhoNFFxo0m6fNuFUO8QJUk=
@@ -67,8 +69,13 @@ github.com/wailsapp/mimetype v1.4.1 h1:pQN9ycO7uo4vsUUuPeHEYoUkLVkaRntMnHJxVwYhw
github.com/wailsapp/mimetype v1.4.1/go.mod h1:9aV5k31bBOv5z6u+QP8TltzvNGJPmNJD4XlAL3U+j3o=
github.com/wailsapp/wails/v2 v2.13.0 h1:S7OgXWpj72V91unF8iDWJKbcS9ZpwCT3R0QVru4v2Mg=
github.com/wailsapp/wails/v2 v2.13.0/go.mod h1:nVr/wSIEZ7xxKPkzK65mjpKpaOPQI2k4pvLwGR/i4kc=
github.com/yalue/onnxruntime_go v1.31.0 h1:1ln4YW1SFOFfGJZXe3jNOb2JUSt+l2pEneZfV8HdtFA=
github.com/yalue/onnxruntime_go v1.31.0/go.mod h1:b4X26A8pekNb1ACJ58wAXgNKeUCGEAQ9dmACut9Sm/4=
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
golang.org/x/image v0.0.0-20191009234506-e7c1f5e7dbb8/go.mod h1:FeLwcggjj3mMvU+oOTbSwawSJRM1uh48EjtB4UJZlP0=
golang.org/x/image v0.40.0 h1:Tw4GyDXMo+daZN1znreBRC3VayR1aLFUyUEOLUdW1a8=
golang.org/x/image v0.40.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA=
golang.org/x/net v0.0.0-20210505024714-0287a6fb4125/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
@@ -81,6 +88,7 @@ golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=