diff --git a/app.db b/app.db index 1a90b8c..d07816c 100644 Binary files a/app.db and b/app.db differ diff --git a/backend/app.go b/backend/app.go index af77d9c..3e3dc6e 100644 --- a/backend/app.go +++ b/backend/app.go @@ -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()} diff --git a/core/const/const.py b/core/const/const.py index 38a9810..3399357 100644 --- a/core/const/const.py +++ b/core/const/const.py @@ -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) \ No newline at end of file diff --git a/core/dataloader/dataloader.py b/core/dataloader/dataloader.py index a68558e..166e00b 100644 --- a/core/dataloader/dataloader.py +++ b/core/dataloader/dataloader.py @@ -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)), ]) diff --git a/core/models/resnet_epoch_100.onnx b/core/models/resnet_epoch_100.onnx index 7b9c81d..83bf702 100644 Binary files a/core/models/resnet_epoch_100.onnx and b/core/models/resnet_epoch_100.onnx differ diff --git a/core/test/test_resnet_onnx.py b/core/test/test_onnx.py similarity index 77% rename from core/test/test_resnet_onnx.py rename to core/test/test_onnx.py index e4cec1f..a54ed83 100644 --- a/core/test/test_resnet_onnx.py +++ b/core/test/test_onnx.py @@ -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) diff --git a/core/test/test_resnet18.py b/core/test/test_resnet18.py index 7e8a4e3..457e16e 100644 --- a/core/test/test_resnet18.py +++ b/core/test/test_resnet18.py @@ -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) diff --git a/core/train/train_resnet_onnx.py b/core/train/to_onnx.py similarity index 61% rename from core/train/train_resnet_onnx.py rename to core/train/to_onnx.py index 3c67f24..3b42e6a 100644 --- a/core/train/train_resnet_onnx.py +++ b/core/train/to_onnx.py @@ -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"], diff --git a/frontend/src/components/History.tsx b/frontend/src/components/History.tsx index a16d767..0700395 100644 --- a/frontend/src/components/History.tsx +++ b/frontend/src/components/History.tsx @@ -27,7 +27,7 @@ const History = () => { const [showClear, setShowClear] = useState(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) diff --git a/frontend/src/components/ImagePreview.tsx b/frontend/src/components/ImagePreview.tsx index 2409b4a..27ca69d 100644 --- a/frontend/src/components/ImagePreview.tsx +++ b/frontend/src/components/ImagePreview.tsx @@ -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) } diff --git a/frontend/src/components/Main.tsx b/frontend/src/components/Main.tsx index deff108..93aa970 100644 --- a/frontend/src/components/Main.tsx +++ b/frontend/src/components/Main.tsx @@ -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) diff --git a/go.mod b/go.mod index b1fc58f..15d9410 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index 0450608..84c7329 100644 --- a/go.sum +++ b/go.sum @@ -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=