From bf4d764e58dda8e72b894970b508a511c5d08e0c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=8F=E4=B8=9C=E4=BA=AE?= Date: Thu, 3 Sep 2026 14:48:13 +0800 Subject: [PATCH] lots change --- .gitignore | 3 + app.db | Bin 24576 -> 24576 bytes backend/app.go | 111 +++++++++++++++------------- backend/types.go | 22 +++--- core/train/to_onnx.py | 7 +- frontend/src/components/Footer.tsx | 2 +- frontend/src/components/History.tsx | 4 +- frontend/src/components/Main.tsx | 31 ++++---- frontend/src/components/Modal.tsx | 9 +-- frontend/src/preact.d.ts | 1 + frontend/src/utils/toast.ts | 8 +- 11 files changed, 101 insertions(+), 97 deletions(-) diff --git a/.gitignore b/.gitignore index b058b8f..754d768 100644 --- a/.gitignore +++ b/.gitignore @@ -23,6 +23,7 @@ env/ uploads/* +core/dataset/*.zip core/dataset/toy/* core/dataset/benchmark/* core/models/*.pth @@ -33,3 +34,5 @@ build/bin frontend/node_modules frontend/dist frontend/public + +static/images/* diff --git a/app.db b/app.db index f017e0e5248996825cba643a7c35c96a55ff4ebf..2c7ca50940eb0a7aa454f5ffd327d5d2224b0dde 100644 GIT binary patch delta 3312 zcmZ9P%WoUk6~?JZmTM^q(_M@dG1|CE0i%gkIIz1)f&LL$w16O6swLNvQb{eG)$lD+ zd`e7Rbh#d+*Gm(LSqwXD&UG4QMf(xsP*x=Q|HM zz8P|SGZcGoXm;%vf4nkl`}pgBA5RzTW5f1{&Hi6|!yc*s{2hl~xj8;+`()_q$i&@y z_ij(z`SjkM+Y|Rbe)Q$T;m^+h_nVLyu0C2n|LNw~E32^XcXub`su!Q{`TqH#^U?Z` z;hWc9S#=vH-o5kMSwFt{md$1xv43Z?pWDB;5C6yh9SmvMuaA!m!`Qi{<%Pws#>a+- zuMgc{oSB(k8r%5(laJmV8+qUM_W4`?db2&Qj7?wr)zxoDx2`-HdH>k?#)#tj9OivA z;rZLrgN5bgxux03zbrhS)`L-A=rC`UucmZIU4#yp_f!O0tmovZ3d^f3QxN-37EQ5A zM_*}(v-@5$TUSfD|!)Z3OWYI1t*{%*q(9ODtr;<;4 zTpuE+*Kc_~xIZ)haOyET%Cpr14+iw(7Kn=Yv8F0%-l)mnFo%^-@qy=Dnp$M}7A(*m zA?}Zn;U+_4B?2}q?9@-|xX)Nw(*rf`PZ%4g^pIE+u9m)iz$*k*>u#{kpvWJ52D z%!438A%mbbf6fUhj-hgqh=`tXb=3zC6k zb2(=sSSeruB1JMC4_)NK}odeyzQuSGQUJ z0Pl^RJOror+;EGQQJ%-8US%!N8x6g`1|f;#Hj5^3ng!L;H`RsR7258O!eXtt1_!uO zI76uPtrG9ViSTEmqqg^McqW(TroNb2n&Ih|=q4@g@p_eIx>P5kMX;7%Q66wA$UtfO zS`wiq8u=qhU$WI&LEL5$mdUL3%Z7r&m-fg^MdU4|8Y^K+WfG(x^%7DfU6>^n4=j_> zUE7v8r82>XpB=**O2|}1Q_DD6-^k$*o{MU2C(5mbWk9c}Bv7$0OOEYSjeMPLgt4Ds znY9+B32-5IR%z*#9ORBqvc8&8N#iEeN7Pzk`0kZa+b?f;#=n@Fo1d9{u&^*MI$mlR z&}b-$3J*{KsKec~;<0j7lET4~20JUsa7$SHWP??FOUNpVCNWecig zEXuFsDTJ7hlbL5F(-ranX{S(|iw0T3V?lV$YfW88k;6+7PwOjn-QA^&gUP|Ik>@ zbc0%asm>s+^=J(QC3{j^eOu-|f&-?P_g-!wY(*unU(N6VQk9tk5CU}NC;(ML&SVC?|CI|87Hz%)_1dBk`>dI_+IvTv~-1f zx<#D33P!Sqk}jJ&nyT#Yut4v%NL$y2DL>syi&GD$7al)mO*eB=&qP^~5MZXSkuMZL!-8d zTkb!hsxQtxn3|s)pi<@*Ur!!$%8)M!45Z$U0zXcq0(iEob+c9;qS@2Tw^L>zj2>02 z0V*mf7;eo3v4F|Yqd#ET1*E1j7;a^o0wp*HEHrNxJ-8l7M5MyD-8h>ZU1aqjKRr&+#V?OSt+{SRp)FYaidpp=j|4Ll3*GK>G(~bWF-tu|L delta 79 zcmZoTz}Rqrae_3X`$QRMR(A%y-d`J27VvX02rz&E%Vt4^&HNK5IP&lUc})Cw8TjAx W--QawO}?wI&cq-9<+?9&PyhghNEZzN diff --git a/backend/app.go b/backend/app.go index 925a989..ae3d349 100644 --- a/backend/app.go +++ b/backend/app.go @@ -3,6 +3,7 @@ package backend import ( "bytes" "context" + "encoding/base64" "fmt" "math" "os" @@ -17,23 +18,12 @@ import ( ) var ( - publicImagePath = "./frontend/public/images" - modelPath = "resnet_epoch_100.onnx" + staticImagesPath = "./static/images" + modelPath = "resnet18_epoch_50.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 { @@ -90,16 +80,14 @@ func preprocessImage(imgData []byte) ([]float32, error) { return input, nil } -func runInference(input []float32) (string, float64, error) { - // 创建输入张量 [1, 3, 224, 224] +func runInference(input []float32, labels []string) (string, float64, error) { 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)) + outputTensor, err := ort.NewTensor(ort.Shape{1, 12}, make([]float32, 12)) if err != nil { return "", 0, fmt.Errorf("create output tensor: %w", err) } @@ -147,7 +135,7 @@ func runInference(input []float32) (string, float64, error) { } confidence := math.Exp(float64(maxVal)) / sum - return labelName[maxIdx], confidence, nil + return labels[maxIdx], confidence, nil } func (a *App) GormDB() (*gorm.DB, error) { @@ -159,7 +147,7 @@ func (a *App) GormDB() (*gorm.DB, error) { } func (a *App) UploadImage(data []byte, filename string) Response { - uploadsDir := publicImagePath + uploadsDir := staticImagesPath err := os.MkdirAll(uploadsDir, 0755) if err != nil { return Response{Code: 1, Message: "failed", Data: err.Error()} @@ -174,32 +162,12 @@ 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} -} - -func (a *App) GetHistory(page int, pageSize int) Response { - db, err := a.GormDB() - if err != nil { - return Response{Code: 1, Message: err.Error()} + imageResult := map[string]any{ + "filename": newFilename, + "data": base64.StdEncoding.EncodeToString(data), } - var total int64 - db.Table("history_test").Count(&total) - - var historyList []HistoryWithBreed - err = db.Table("history_test").Select("history_test.*, breeds_test.brief, breeds_test.name"). - Joins("LEFT JOIN breeds_test ON history_test.breed = breeds_test.id"). - Order("history_test.id DESC").Limit(pageSize).Offset((page - 1) * pageSize).Find(&historyList).Error - - if err != nil { - return Response{Code: 1, Message: err.Error()} - } - - historyData := HistoryData{Page: page, PageSize: pageSize, Total: total, List: historyList} - - return Response{Code: 0, Message: "success", Data: historyData} + return Response{Code: 0, Message: "success", Data: imageResult} } func (a *App) Detect(filename string) Response { @@ -208,7 +176,7 @@ func (a *App) Detect(filename string) Response { return Response{Code: 1, Message: err.Error()} } - filePath := filepath.Join(publicImagePath, filename) + filePath := filepath.Join(staticImagesPath, filename) imgData, err := os.ReadFile(filePath) if err != nil { return Response{Code: 1, Message: "failed to read image: " + err.Error()} @@ -219,22 +187,31 @@ func (a *App) Detect(filename string) Response { return Response{Code: 1, Message: "failed to preprocess: " + err.Error()} } + // 查询breeds表的Code列,组成切片 + var labels []string + err = db.Table("breeds").Pluck("code", &labels).Error + if err != nil { + return Response{Code: 1, Message: "failed to query breed codes: " + err.Error()} + } + // 推理 - detectRet, confidence, err := runInference(input) + detectRet, confidence, err := runInference(input, labels) if err != nil { return Response{Code: 1, Message: "model inference failed: " + err.Error()} } + confidence = math.Round(confidence*10000) / 10000 var breed Breed - err = db.Table("breeds_test").Where("code = ?", detectRet).First(&breed).Error + err = db.Table("breeds").Where("code = ?", detectRet).First(&breed).Error if err != nil { return Response{Code: 1, Message: err.Error()} } now := int(time.Now().Unix()) + print("confidence", confidence) - one := HistoryItem{Img: filename, Breed: int(breed.Id), Date: now} - result := db.Table("history_test").Create(&one) + one := HistoryItem{Img: filename, Breed: int(breed.Id), Confidence: confidence, Date: now} + result := db.Table("history").Create(&one) if result.Error != nil { return Response{Code: 1, Message: result.Error.Error()} } @@ -250,13 +227,43 @@ func (a *App) Detect(filename string) Response { return Response{Code: 0, Message: "success", Data: detectData} } +func (a *App) GetHistory(page int, pageSize int) Response { + db, err := a.GormDB() + if err != nil { + return Response{Code: 1, Message: err.Error()} + } + + var total int64 + db.Table("history").Count(&total) + + var historyList []HistoryWithBreed + err = db.Table("history").Select("history.*, breeds.brief, breeds.name"). + Joins("LEFT JOIN breeds ON history.breed = breeds.id"). + Order("history.id DESC").Limit(pageSize).Offset((page - 1) * pageSize).Find(&historyList).Error + + if err != nil { + return Response{Code: 1, Message: err.Error()} + } + + for i := range historyList { + imgPath := filepath.Join(staticImagesPath, historyList[i].Img) + if data, err := os.ReadFile(imgPath); err == nil { + historyList[i].ImgData = base64.StdEncoding.EncodeToString(data) + } + } + + historyData := HistoryData{Page: page, PageSize: pageSize, Total: total, List: historyList} + + return Response{Code: 0, Message: "success", Data: historyData} +} + func (a *App) DeleteOneHistory(id uint) Response { db, err := a.GormDB() if err != nil { return Response{Code: 1, Message: err.Error()} } - result := db.Table("history_test").Delete(&HistoryItem{}, id) + result := db.Table("history").Delete(&HistoryItem{}, id) if result.Error != nil { return Response{Code: 1, Message: result.Error.Error()} } @@ -270,15 +277,15 @@ func (a *App) ClearHistory() Response { return Response{Code: 1, Message: err.Error()} } - result := db.Exec("DELETE FROM history_test") + result := db.Exec("DELETE FROM history") if result.Error != nil { return Response{Code: 1, Message: result.Error.Error()} } - if err := os.RemoveAll(publicImagePath); err != nil { + if err := os.RemoveAll(staticImagesPath); err != nil { return Response{Code: 1, Message: fmt.Sprintf("failed to remove images: %v", err)} } - if err := os.MkdirAll(publicImagePath, 0755); err != nil { + if err := os.MkdirAll(staticImagesPath, 0755); err != nil { return Response{Code: 1, Message: fmt.Sprintf("failed to recreate images dir: %v", err)} } diff --git a/backend/types.go b/backend/types.go index f551309..adad6fe 100644 --- a/backend/types.go +++ b/backend/types.go @@ -16,19 +16,21 @@ type ( } HistoryItem struct { - Id uint `gorm:"primaryKey"` - Img string `gorm:"column:img"` - Breed int `gorm:"column:breed"` - Date int `gorm:"column:date"` + Id uint `gorm:"primaryKey"` + Img string `gorm:"column:img"` + Breed int `gorm:"column:breed"` + Confidence float64 `gorm:"column:confidence"` + Date int `gorm:"column:date"` } HistoryWithBreed struct { - Id uint `gorm:"column:id" json:"id"` - Img string `gorm:"column:img" json:"img"` - Breed int `gorm:"column:breed" json:"breed"` - Date int `gorm:"column:date" json:"date"` - Name string `gorm:"column:name" json:"name"` - Brief string `gorm:"column:brief" json:"brief"` + Id uint `gorm:"column:id" json:"id"` + Img string `gorm:"column:img" json:"img"` + ImgData string `gorm:"-" json:"img_data"` + Breed int `gorm:"column:breed" json:"breed"` + Date int `gorm:"column:date" json:"date"` + Name string `gorm:"column:name" json:"name"` + Brief string `gorm:"column:brief" json:"brief"` } HistoryData struct { diff --git a/core/train/to_onnx.py b/core/train/to_onnx.py index 3b42e6a..74b093e 100644 --- a/core/train/to_onnx.py +++ b/core/train/to_onnx.py @@ -1,11 +1,8 @@ import os import torch -import sys -# from core.nets.resnet import resnet from core.nets.resnet18 import resnet18 # 加载 pth -# net = resnet() # 实例化你的模型 net = resnet18() @@ -14,7 +11,7 @@ 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, "resnet18_epoch_50_bak2.pth"), map_location="cpu")) +net.load_state_dict(torch.load(os.path.join(model_dir, "resnet18_epoch_50.pth"), map_location="cpu")) net.eval() @@ -23,7 +20,7 @@ dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( net, dummy_input, - os.path.join(model_dir, "resnet18_epoch_50_bak2.onnx"), + os.path.join(model_dir, "resnet18_epoch_50.onnx"), export_params=True, opset_version=11, input_names=["input"], diff --git a/frontend/src/components/Footer.tsx b/frontend/src/components/Footer.tsx index 7147586..52a6dbf 100644 --- a/frontend/src/components/Footer.tsx +++ b/frontend/src/components/Footer.tsx @@ -2,7 +2,7 @@ export const Footer = () => (
📩 xiadongliang88@163.com | - 📦 https://git.leonstack.com/owner + 📦 https://git.leonstack.com/xiadongliang | 🌏 https://www.leonstack.com/
diff --git a/frontend/src/components/History.tsx b/frontend/src/components/History.tsx index da5fced..55fbcf1 100644 --- a/frontend/src/components/History.tsx +++ b/frontend/src/components/History.tsx @@ -4,6 +4,7 @@ import Modal from './Modal' import { message } from '../utils/toast' import type { HistoryItem } from '../preact' + const formatDate = (timestamp: number) => { const date = new Date(timestamp * 1000) return date.toLocaleString('zh-CN', { @@ -64,7 +65,6 @@ const History = () => { message.error('currentId为空') return } - // dfdfdf const result = await (window as any).go.backend.App.DeleteOneHistory(currentId) if (result.code === 0) { message.success('删除成功') @@ -109,7 +109,7 @@ const History = () => { {historyList.map((item: HistoryItem) =>
- handleShowItem(item)} /> + handleShowItem(item)} />
diff --git a/frontend/src/components/Main.tsx b/frontend/src/components/Main.tsx index dda72bf..357c187 100644 --- a/frontend/src/components/Main.tsx +++ b/frontend/src/components/Main.tsx @@ -2,9 +2,11 @@ import { useState, useRef } from 'preact/hooks' import type { DetectResult } from '../preact' import { message } from '../utils/toast' + const Main = () => { const fileInputRef = useRef(null) const [fileSrc, setFileSrc] = useState('') + const [filename, setFilename] = useState('') const [step, setStep] = useState(0) const [detectResult, setDetectResult] = useState(null) @@ -46,7 +48,6 @@ const Main = () => { } const handleFileChange = (e: Event) => { - console.log('handleFileChange') const target = e.target as HTMLInputElement const file = target.files?.[0] if (file) { @@ -64,8 +65,8 @@ const Main = () => { file.name ) if (result.code === 0) { - console.log("rrr", result) - setFileSrc(result.data) + setFileSrc(result.data.data) + setFilename(result.data.filename) setStep(1) } else if (result.code === 1) { message.error(result.message) @@ -92,14 +93,16 @@ const Main = () => { setStep(2) setTimeout(async() => { - 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) { - message.error(result.message) + if (filename) { + const result = await (window as any).go.backend.App.Detect(filename) + if (result.code === 0) { + setDetectResult(result.data) + setStep(3) + } else if (result.code === 1) { + message.error(result.message) + } } - }, 1000) + }, 500) } const handleTryOther = () => { @@ -108,8 +111,6 @@ const Main = () => { resetUpload() } - console.log("fff", fileSrc) - return (
@@ -135,7 +136,7 @@ const Main = () => { > {fileSrc.length > 0 && step == 1 ?
- + @@ -194,7 +195,7 @@ const Main = () => {

完成!

- +

{detectResult?.name}

@@ -202,7 +203,7 @@ const Main = () => {
- 置信度 {detectResult ? detectResult.confidence_level * 100 + '%' : ''} + 置信度 {detectResult ? (detectResult.confidence_level * 100).toFixed(2) + '%' : ''}

{detectResult?.brief}

diff --git a/frontend/src/components/Modal.tsx b/frontend/src/components/Modal.tsx index c0e6f2c..201c6d1 100644 --- a/frontend/src/components/Modal.tsx +++ b/frontend/src/components/Modal.tsx @@ -1,9 +1,8 @@ import type { ModalProps } from '../preact' + const Modal = ({ open, title, onClick, onClose, children }: ModalProps) => { - const handleClose = () => { - onClose?.() - } + const handleClose = () => onClose?.() const handleConfirm = () => { onClick?.() @@ -11,9 +10,7 @@ const Modal = ({ open, title, onClick, onClose, children }: ModalProps) => { } const handleMaskClick = (e: MouseEvent) => { - if (e.target === e.currentTarget) { - handleClose() - } + if (e.target === e.currentTarget) handleClose() } return ( diff --git a/frontend/src/preact.d.ts b/frontend/src/preact.d.ts index d1434a7..004a965 100644 --- a/frontend/src/preact.d.ts +++ b/frontend/src/preact.d.ts @@ -1,6 +1,7 @@ export interface HistoryItem { id: number img: string + img_data: string breed: number date: number name: string diff --git a/frontend/src/utils/toast.ts b/frontend/src/utils/toast.ts index 2e07d99..0830d04 100644 --- a/frontend/src/utils/toast.ts +++ b/frontend/src/utils/toast.ts @@ -23,9 +23,7 @@ export const message: MessageAPI = { div.appendChild(subDiv) document.body.appendChild(div) - setTimeout(() => { - div.remove() - }, 3000) + setTimeout(() => div.remove(), 3000) }, error: (text: string) => { const div = document.createElement('div') @@ -46,8 +44,6 @@ export const message: MessageAPI = { div.appendChild(subDiv) document.body.appendChild(div) - setTimeout(() => { - div.remove() - }, 3000) + setTimeout(() => div.remove(), 3000) } }