diff --git a/core/const/__init__.py b/core/const/__init__.py index 3d09d5f..983f876 100644 --- a/core/const/__init__.py +++ b/core/const/__init__.py @@ -1,44 +1 @@ -mode = "toy" - -""" -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) \ No newline at end of file +from .const import * \ No newline at end of file diff --git a/core/const/const.py b/core/const/const.py new file mode 100644 index 0000000..38a9810 --- /dev/null +++ b/core/const/const.py @@ -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) \ No newline at end of file diff --git a/core/dataloader/dataloader.py b/core/dataloader/dataloader.py index a730ff4..a68558e 100644 --- a/core/dataloader/dataloader.py +++ b/core/dataloader/dataloader.py @@ -3,7 +3,7 @@ import glob from torchvision import transforms from torch.utils.data import DataLoader, Dataset 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 = {} @@ -17,7 +17,7 @@ def default_loader(path): train_transform = transforms.Compose([ - transforms.Resize((input_size)), + transforms.Resize(input_size), transforms.CenterCrop(input_size), transforms.RandomHorizontalFlip(p=0.5), # 50%的概率(p=0.5)水平翻转图片 transforms.RandomRotation(10), # 轻微旋转 @@ -62,8 +62,8 @@ class MyDataset(Dataset): 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_test_list = glob.glob(os.path.join(dataset_root, "dataset", "toy", "test", "*", "*.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", mode, "test", "*", "*.jpg")) train_dataset = MyDataset(im_train_list, transform=train_transform) diff --git a/core/models/resnet_epoch_100.onnx b/core/models/resnet_epoch_100.onnx new file mode 100644 index 0000000..7b9c81d Binary files /dev/null and b/core/models/resnet_epoch_100.onnx differ diff --git a/core/test/test_resnet.py b/core/test/test_resnet.py index 098e2d5..2a2b698 100644 --- a/core/test/test_resnet.py +++ b/core/test/test_resnet.py @@ -1,3 +1,4 @@ +import os import cv2 import glob import torch @@ -5,21 +6,32 @@ from torchvision import transforms from PIL import Image import numpy as np 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(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(device) - net = resnet() - net.load_state_dict(torch.load("./model/resnet_epoch_15.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 = 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) net.to(device) + print("222") + test_transform = transforms.Compose([ transforms.Resize(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]) ]) + print("333") + + ary = [] for im_path in im_list: + # print("im_path", im_path) net.eval() im_data = Image.open(im_path) @@ -40,11 +56,15 @@ def test(): _, pred = torch.max(outputs.data, dim=1) print(label_name[pred.cpu().numpy()[0]], " ", im_path) - img = np.asarray(im_data) - img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR) - cv2.imshow("img", img) - cv2.waitKey(0) + result = label_name[pred.cpu().numpy()[0]] + if result in im_path: + ary.append(True) + else: + ary.append(False) + + print("ary", ary) + print(sum(ary) / len(ary)) if __name__ == "__main__": test() diff --git a/core/test/test_resnet_onnx.py b/core/test/test_resnet_onnx.py new file mode 100644 index 0000000..e4cec1f --- /dev/null +++ b/core/test/test_resnet_onnx.py @@ -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() \ No newline at end of file diff --git a/core/train/train_resnet.py b/core/train/train_resnet.py index 053b183..89ee1cc 100644 --- a/core/train/train_resnet.py +++ b/core/train/train_resnet.py @@ -41,10 +41,12 @@ def train(): scheduler.step() print("lr: ", optimizer.state_dict()['param_groups'][0]['lr']) - if not os.path.exists("./model"): - os.makedirs("./model") + script_dir = os.path.dirname(os.path.abspath(__file__)) + 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__": diff --git a/core/train/train_resnet18.py b/core/train/train_resnet18.py index 3cd7e73..b42fe69 100644 --- a/core/train/train_resnet18.py +++ b/core/train/train_resnet18.py @@ -41,10 +41,12 @@ def train(): scheduler.step() print("lr: ", optimizer.state_dict()['param_groups'][0]['lr']) - if not os.path.exists("./models"): - os.makedirs("./models") + script_dir = os.path.dirname(os.path.abspath(__file__)) + 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__": diff --git a/core/train/train_resnet_onnx.py b/core/train/train_resnet_onnx.py new file mode 100644 index 0000000..3c67f24 --- /dev/null +++ b/core/train/train_resnet_onnx.py @@ -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"}} +)