resnet parm

This commit is contained in:
2026-07-24 17:44:47 +08:00
parent 70b846b41c
commit a217077c13
9 changed files with 171 additions and 62 deletions
+1 -44
View File
@@ -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)
+44
View File
@@ -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)
+4 -4
View File
@@ -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.
+28 -8
View File
@@ -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()
+57
View File
@@ -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】将输入转为numpyONNX 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()
+5 -3
View File
@@ -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__":
+5 -3
View File
@@ -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__":
+27
View File
@@ -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"}}
)