Compare commits
2 Commits
70b846b41c
...
44b0d6b756
| Author | SHA1 | Date | |
|---|---|---|---|
| 44b0d6b756 | |||
| a217077c13 |
@@ -27,6 +27,7 @@ core/dataset/toy/*
|
|||||||
core/dataset/benchmark/*
|
core/dataset/benchmark/*
|
||||||
core/models/*.pth
|
core/models/*.pth
|
||||||
core/models/*.pt
|
core/models/*.pt
|
||||||
|
core/models/*.onnx
|
||||||
|
|
||||||
build/bin
|
build/bin
|
||||||
frontend/node_modules
|
frontend/node_modules
|
||||||
|
|||||||
+1
-44
@@ -1,44 +1 @@
|
|||||||
mode = "toy"
|
from .const import *
|
||||||
|
|
||||||
"""
|
|
||||||
epoch 训练多少轮
|
|
||||||
lr 学习率
|
|
||||||
batch_size 一轮跑多少张图片
|
|
||||||
input_size 训练图片输入尺寸
|
|
||||||
label_name 分类
|
|
||||||
"""
|
|
||||||
|
|
||||||
if mode == "toy":
|
|
||||||
epoch = 15
|
|
||||||
lr = 2e-4
|
|
||||||
batch_size = 2
|
|
||||||
input_size = 224
|
|
||||||
|
|
||||||
# 分类
|
|
||||||
label_name = [
|
|
||||||
"american_shorthair",
|
|
||||||
"bengal",
|
|
||||||
"british_shorthair",
|
|
||||||
"exotic_shorthair",
|
|
||||||
"maine_coon",
|
|
||||||
"ragdoll",
|
|
||||||
"sphynx",
|
|
||||||
]
|
|
||||||
num_classes = len(label_name)
|
|
||||||
elif mode == "benchmark":
|
|
||||||
epoch = 15
|
|
||||||
lr = 2e-4
|
|
||||||
batch_size = 2
|
|
||||||
input_size = 224
|
|
||||||
|
|
||||||
# 分类
|
|
||||||
label_name = [
|
|
||||||
"american_shorthair",
|
|
||||||
"bengal",
|
|
||||||
"british_shorthair",
|
|
||||||
"exotic_shorthair",
|
|
||||||
"maine_coon",
|
|
||||||
"ragdoll",
|
|
||||||
"sphynx",
|
|
||||||
]
|
|
||||||
num_classes = len(label_name)
|
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
mode = "toy"
|
||||||
|
|
||||||
|
"""
|
||||||
|
epoch 训练多少轮
|
||||||
|
lr 学习率
|
||||||
|
batch_size 一轮跑多少张图片
|
||||||
|
input_size 训练图片输入尺寸
|
||||||
|
label_name 分类
|
||||||
|
"""
|
||||||
|
|
||||||
|
if mode == "toy":
|
||||||
|
epoch = 100
|
||||||
|
lr = 5e-4
|
||||||
|
batch_size = 8
|
||||||
|
input_size = 224
|
||||||
|
|
||||||
|
# 分类
|
||||||
|
label_name = [
|
||||||
|
"american_shorthair",
|
||||||
|
"bengal",
|
||||||
|
"british_shorthair",
|
||||||
|
"exotic_shorthair",
|
||||||
|
"maine_coon",
|
||||||
|
"ragdoll",
|
||||||
|
"sphynx",
|
||||||
|
]
|
||||||
|
num_classes = len(label_name)
|
||||||
|
elif mode == "benchmark":
|
||||||
|
epoch = 50
|
||||||
|
lr = 2e-4
|
||||||
|
batch_size = 2
|
||||||
|
input_size = 224
|
||||||
|
|
||||||
|
# 分类
|
||||||
|
label_name = [
|
||||||
|
"american_shorthair",
|
||||||
|
"bengal",
|
||||||
|
"british_shorthair",
|
||||||
|
"exotic_shorthair",
|
||||||
|
"maine_coon",
|
||||||
|
"ragdoll",
|
||||||
|
"sphynx",
|
||||||
|
]
|
||||||
|
num_classes = len(label_name)
|
||||||
@@ -3,7 +3,7 @@ import glob
|
|||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from torch.utils.data import DataLoader, Dataset
|
from torch.utils.data import DataLoader, Dataset
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from core.const import label_name, input_size, batch_size
|
from core.const import mode, label_name, input_size, batch_size
|
||||||
|
|
||||||
|
|
||||||
label_dict = {}
|
label_dict = {}
|
||||||
@@ -17,7 +17,7 @@ def default_loader(path):
|
|||||||
|
|
||||||
|
|
||||||
train_transform = transforms.Compose([
|
train_transform = transforms.Compose([
|
||||||
transforms.Resize((input_size)),
|
transforms.Resize(input_size),
|
||||||
transforms.CenterCrop(input_size),
|
transforms.CenterCrop(input_size),
|
||||||
transforms.RandomHorizontalFlip(p=0.5), # 50%的概率(p=0.5)水平翻转图片
|
transforms.RandomHorizontalFlip(p=0.5), # 50%的概率(p=0.5)水平翻转图片
|
||||||
transforms.RandomRotation(10), # 轻微旋转
|
transforms.RandomRotation(10), # 轻微旋转
|
||||||
@@ -62,8 +62,8 @@ class MyDataset(Dataset):
|
|||||||
|
|
||||||
|
|
||||||
dataset_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
dataset_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
im_train_list = glob.glob(os.path.join(dataset_root, "dataset", "toy", "train", "*", "*.jpg"))
|
im_train_list = glob.glob(os.path.join(dataset_root, "dataset", mode, "train", "*", "*.jpg"))
|
||||||
im_test_list = glob.glob(os.path.join(dataset_root, "dataset", "toy", "test", "*", "*.jpg"))
|
im_test_list = glob.glob(os.path.join(dataset_root, "dataset", mode, "test", "*", "*.jpg"))
|
||||||
|
|
||||||
|
|
||||||
train_dataset = MyDataset(im_train_list, transform=train_transform)
|
train_dataset = MyDataset(im_train_list, transform=train_transform)
|
||||||
|
|||||||
Binary file not shown.
@@ -1,3 +1,4 @@
|
|||||||
|
import os
|
||||||
import cv2
|
import cv2
|
||||||
import glob
|
import glob
|
||||||
import torch
|
import torch
|
||||||
@@ -5,21 +6,32 @@ from torchvision import transforms
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from core.nets.resnet import resnet
|
from core.nets.resnet import resnet
|
||||||
from core.const import label_name, input_size
|
from core.const import mode, label_name, input_size
|
||||||
|
|
||||||
|
|
||||||
def test():
|
def test():
|
||||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
print(device)
|
print(device)
|
||||||
|
|
||||||
net = resnet()
|
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
net.load_state_dict(torch.load("./model/resnet_epoch_15.pth", weights_only=True))
|
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 = resnet()
|
||||||
|
net.load_state_dict(torch.load(os.path.join(model_dir, "resnet_epoch_50.pth"), weights_only=True))
|
||||||
|
|
||||||
|
print("111")
|
||||||
|
|
||||||
|
im_list = glob.glob(os.path.join(dataset_dir, "*", "*.jpg"))
|
||||||
np.random.shuffle(im_list)
|
np.random.shuffle(im_list)
|
||||||
|
|
||||||
net.to(device)
|
net.to(device)
|
||||||
|
|
||||||
|
print("222")
|
||||||
|
|
||||||
test_transform = transforms.Compose([
|
test_transform = transforms.Compose([
|
||||||
transforms.Resize(input_size),
|
transforms.Resize(input_size),
|
||||||
transforms.CenterCrop(input_size),
|
transforms.CenterCrop(input_size),
|
||||||
@@ -27,7 +39,11 @@ def test():
|
|||||||
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])
|
||||||
])
|
])
|
||||||
|
|
||||||
|
print("333")
|
||||||
|
|
||||||
|
ary = []
|
||||||
for im_path in im_list:
|
for im_path in im_list:
|
||||||
|
# print("im_path", im_path)
|
||||||
net.eval()
|
net.eval()
|
||||||
im_data = Image.open(im_path)
|
im_data = Image.open(im_path)
|
||||||
|
|
||||||
@@ -40,11 +56,15 @@ def test():
|
|||||||
_, pred = torch.max(outputs.data, dim=1)
|
_, pred = torch.max(outputs.data, dim=1)
|
||||||
print(label_name[pred.cpu().numpy()[0]], " ", im_path)
|
print(label_name[pred.cpu().numpy()[0]], " ", im_path)
|
||||||
|
|
||||||
img = np.asarray(im_data)
|
result = label_name[pred.cpu().numpy()[0]]
|
||||||
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
|
|
||||||
cv2.imshow("img", img)
|
|
||||||
cv2.waitKey(0)
|
|
||||||
|
|
||||||
|
if result in im_path:
|
||||||
|
ary.append(True)
|
||||||
|
else:
|
||||||
|
ary.append(False)
|
||||||
|
|
||||||
|
print("ary", ary)
|
||||||
|
print(sum(ary) / len(ary))
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
test()
|
test()
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
import cv2
|
||||||
|
import glob
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def test():
|
||||||
|
# 【修改2】选择ONNX Runtime的执行提供程序,自动选择CPU或CUDA
|
||||||
|
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if 'CUDAExecutionProvider' in ort.get_available_providers() else ['CPUExecutionProvider']
|
||||||
|
print(f"使用设备: {providers[0]}")
|
||||||
|
|
||||||
|
# 【修改3】加载ONNX模型,替代原来的PyTorch模型加载
|
||||||
|
session = ort.InferenceSession("./model/resnet_final.onnx", providers=providers)
|
||||||
|
|
||||||
|
# 获取输入名称(用于后续推理时指定输入)
|
||||||
|
input_name = session.get_inputs()[0].name
|
||||||
|
|
||||||
|
im_list = glob.glob("./dataset/test/*/*.jpg")
|
||||||
|
np.random.shuffle(im_list)
|
||||||
|
|
||||||
|
# 预处理完全不变
|
||||||
|
test_transform = transforms.Compose([
|
||||||
|
transforms.Resize(input_size),
|
||||||
|
transforms.CenterCrop(input_size),
|
||||||
|
transforms.ToTensor(),
|
||||||
|
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
|
||||||
|
])
|
||||||
|
|
||||||
|
for im_path in im_list:
|
||||||
|
im_data = Image.open(im_path)
|
||||||
|
|
||||||
|
inputs = test_transform(im_data)
|
||||||
|
inputs = torch.unsqueeze(inputs, dim=0) # 这里还在用torch,下面会改
|
||||||
|
|
||||||
|
# 【修改4】将输入转为numpy,ONNX Runtime需要numpy输入
|
||||||
|
inputs = inputs.numpy()
|
||||||
|
if providers[0] == 'CPUExecutionProvider':
|
||||||
|
inputs = inputs.astype(np.float32)
|
||||||
|
|
||||||
|
# 【修改5】ONNX Runtime推理,输出直接是numpy数组
|
||||||
|
outputs = session.run(None, {input_name: inputs})[0]
|
||||||
|
|
||||||
|
# 【修改6】解析结果,直接用numpy操作
|
||||||
|
pred = np.argmax(outputs, axis=1)
|
||||||
|
print(label_name[pred[0]], " ", im_path)
|
||||||
|
|
||||||
|
img = np.asarray(im_data)
|
||||||
|
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
|
||||||
|
cv2.imshow("img", img)
|
||||||
|
cv2.waitKey(0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
test()
|
||||||
@@ -41,10 +41,12 @@ def train():
|
|||||||
scheduler.step()
|
scheduler.step()
|
||||||
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
|
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
|
||||||
|
|
||||||
if not os.path.exists("./model"):
|
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
os.makedirs("./model")
|
model_dir = os.path.join(script_dir, "..", "models")
|
||||||
|
if not os.path.exists(model_dir):
|
||||||
|
os.makedirs(model_dir)
|
||||||
|
|
||||||
torch.save(net.state_dict(), "./model/resnet_epoch_{}.pth".format(e + 1))
|
torch.save(net.state_dict(), os.path.join(model_dir, "resnet_epoch_{}.pth".format(e + 1)))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -41,10 +41,12 @@ def train():
|
|||||||
scheduler.step()
|
scheduler.step()
|
||||||
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
|
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
|
||||||
|
|
||||||
if not os.path.exists("./models"):
|
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
os.makedirs("./models")
|
model_dir = os.path.join(script_dir, "..", "models")
|
||||||
|
if not os.path.exists(model_dir):
|
||||||
|
os.makedirs(model_dir)
|
||||||
|
|
||||||
torch.save(net.state_dict(), "./models/resnet18_epoch_{}.pth".format(e + 1))
|
torch.save(net.state_dict(), os.path.join(model_dir, "resnet18_epoch_{}.pth".format(e + 1)))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import sys
|
||||||
|
from core.nets.resnet import resnet
|
||||||
|
|
||||||
|
# 加载 pth
|
||||||
|
net = resnet() # 实例化你的模型
|
||||||
|
|
||||||
|
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.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"),
|
||||||
|
export_params=True,
|
||||||
|
opset_version=11,
|
||||||
|
input_names=["input"],
|
||||||
|
output_names=["output"],
|
||||||
|
dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}}
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user