new train

This commit is contained in:
2026-07-28 18:06:10 +08:00
parent 44b0d6b756
commit a55d449028
13 changed files with 224 additions and 27 deletions
+6 -3
View File
@@ -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)
+2 -1
View File
@@ -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)),
])
Binary file not shown.
@@ -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)
+16 -5
View File
@@ -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)
@@ -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"],