This commit is contained in:
2026-07-24 09:54:42 +08:00
parent 4d1ca8ffaf
commit 70b846b41c
17 changed files with 17 additions and 17 deletions
+1
View File
@@ -0,0 +1 @@
# Neural Network Package
+44
View File
@@ -0,0 +1,44 @@
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)
View File
+78
View File
@@ -0,0 +1,78 @@
import os
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
label_dict = {}
for idx, name in enumerate(label_name):
label_dict[name] = idx
def default_loader(path):
return Image.open(path).convert('RGB')
train_transform = transforms.Compose([
transforms.Resize((input_size)),
transforms.CenterCrop(input_size),
transforms.RandomHorizontalFlip(p=0.5), # 50%的概率(p=0.5)水平翻转图片
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])
])
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])
])
class MyDataset(Dataset):
def __init__(self, im_list, transform=None, loader=default_loader):
super(MyDataset, self).__init__()
imgs = []
for im_item in im_list:
im_label_name = os.path.basename(os.path.dirname(im_item))
imgs.append([im_item, label_dict[im_label_name]])
self.imgs = imgs
self.transform = transform
self.loader = loader
def __getitem__(self, index):
im_path, im_label = self.imgs[index]
im_data = self.loader(im_path)
if self.transform is not None:
im_data = self.transform(im_data)
return im_data, im_label
def __len__(self):
return len(self.imgs)
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"))
train_dataset = MyDataset(im_train_list, transform=train_transform)
test_dataset = MyDataset(im_test_list, transform=test_transform)
print("train_dataset", len(train_dataset))
print("test_dataset", len(test_dataset))
train_dataloader = DataLoader(dataset=train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)
test_dataloader = DataLoader(dataset=test_dataset, batch_size=batch_size, shuffle=False, num_workers=4)
+83
View File
@@ -0,0 +1,83 @@
import torch.nn as nn
import torch.nn.functional as F
from core.const import num_classes
class ResBlock(nn.Module):
def __init__(self, in_channel, out_channel, stride=1):
super(ResBlock, self).__init__()
self.layer = nn.Sequential(
nn.Conv2d(in_channel, out_channel, kernel_size=3, stride=stride, padding=1),
nn.BatchNorm2d(out_channel),
nn.ReLU(),
nn.Conv2d(out_channel, out_channel, kernel_size=3, stride=1, padding=1),
nn.BatchNorm2d(out_channel)
)
if in_channel != out_channel or stride > 1:
self.shortcut = nn.Sequential(
nn.Conv2d(in_channel, out_channel, kernel_size=1, stride=stride),
nn.BatchNorm2d(out_channel),
)
else:
self.shortcut = nn.Sequential()
def forward(self, x):
out = self.layer(x)
shortcut = self.shortcut(x)
out = out + shortcut
out = F.relu(out)
return out
class ResNet(nn.Module):
def make_layer(self, block, out_channel, stride, num_block):
layer_list = []
for i in range(num_block):
if i == 0:
in_stride = stride
else:
in_stride = 1
layer_list.append(block(self.in_channel, out_channel, in_stride))
self.in_channel = out_channel
return nn.Sequential(*layer_list)
def __init__(self):
super(ResNet, self).__init__()
self.in_channel = 32
self.conv1 = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, stride=1, padding=1),
nn.BatchNorm2d(32),
nn.ReLU()
)
self.layer1 = self.make_layer(ResBlock, 64, 2, 2) # 32 -> 64
self.layer2 = self.make_layer(ResBlock, 128, 2, 2) # 64 -> 128
self.layer3 = self.make_layer(ResBlock, 256, 2, 2) # 128 -> 256
self.layer4 = self.make_layer(ResBlock, 512, 2, 2) # 256 -> 512
# 添加全局平均池化
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
self.fc = nn.Linear(512, num_classes)
def forward(self, x):
out = self.conv1(x)
out = self.layer1(out)
out = self.layer2(out)
out = self.layer3(out)
out = self.layer4(out)
# 使用全局平均池化
out = self.avgpool(out)
out = out.view(out.size(0), -1)
out = self.fc(out)
return out
def resnet():
return ResNet()
+19
View File
@@ -0,0 +1,19 @@
import torch.nn as nn
from torchvision import models
from core.const import num_classes
class resnet18(nn.Module):
def __init__(self):
super(resnet18, self).__init__()
self.model = models.resnet18(weights='IMAGENET1K_V1')
self.num_features = self.model.fc.in_features
self.model.fc = nn.Linear(self.num_features, num_classes)
def forward(self, x):
out = self.model(x)
return out
def pytorch_resnet18():
return resnet18()
+67
View File
@@ -0,0 +1,67 @@
class,image_count,avg_width,avg_height,min_width,min_height,max_width,max_height,formats,corrupt_files
abyssinian,200,148,117,88,46,162,140,"jpeg, png",0
cyprus,200,159,100,93,40,162,140,"jpeg, png",0
lykoi,200,142,122,68,55,162,140,"jpeg, png",0
donskoy,200,141,123,67,54,162,140,jpeg,0
chausie,200,143,119,78,46,162,140,"jpeg, png",0
european_shorthair,200,149,118,70,49,162,140,"jpeg, png",0
turkish_van,200,149,120,87,56,162,140,"jpeg, png",0
pixie_bob,200,147,119,78,49,162,140,"jpeg, png",0
ragdoll,199,149,116,78,50,162,140,jpeg,0
german_rex,199,138,125,72,50,162,140,"jpeg, png",0
american_shorthair,199,147,114,65,50,162,140,"jpeg, png",0
sokoke,199,145,122,70,49,162,140,"jpeg, png",0
khao_manee,198,144,120,78,49,162,140,"jpeg, png",0
thai,198,143,118,68,60,162,140,"jpeg, png",0
cymric,198,148,120,63,43,162,140,"jpeg, png",0
oriental_shorthair,197,141,123,78,54,162,140,jpeg,0
cornish_rex,197,147,117,75,56,162,140,"jpeg, png",0
burmese,197,147,117,45,58,162,140,"jpeg, png",0
savannah,197,153,106,78,40,162,140,"jpeg, png",0
american_wirehair,196,149,117,76,50,162,140,"jpeg, png",0
peterbald,196,144,123,77,50,162,140,"jpeg, png",0
karelian_bobtail,196,145,125,67,65,162,140,"jpeg, png",0
tonkinese,195,148,116,70,53,162,140,"jpeg, png",0
balinese,195,154,112,79,54,162,140,jpeg,0
japanese_bobtail,194,149,118,87,49,162,140,"jpeg, png",0
nebelung,194,143,119,64,46,162,140,jpeg,0
selkirk_rex,192,144,122,78,76,162,140,jpeg,0
persian,192,150,114,78,46,162,140,jpeg,0
manx,192,155,113,86,50,162,140,"jpeg, png",0
himalayan,192,158,109,93,40,162,140,jpeg,0
munchkin,191,147,117,61,48,162,140,jpeg,0
bengal,189,151,115,78,71,162,140,jpeg,0
turkish_angora,188,145,119,66,52,162,140,jpeg,0
vankedisi,187,147,116,63,47,162,140,"jpeg, png",0
scottish_fold,184,143,120,78,50,162,140,jpeg,0
egyptian_mau,184,144,121,72,74,162,140,"jpeg, png",0
ocicat,182,150,115,62,44,162,140,"jpeg, png",0
ragamuffin,182,149,116,78,44,162,140,"jpeg, png",0
serengeti,175,159,100,78,53,300,140,"jpeg, png",0
british_shorthair,174,148,117,78,50,162,140,jpeg,0
toyger,160,145,120,78,50,162,140,"jpeg, png",0
siberian,159,153,116,78,55,162,140,jpeg,0
havana_brown,159,133,127,75,34,300,140,"jpeg, png",0
exotic_shorthair,157,148,118,93,55,300,140,"jpeg, png",0
bombay,154,139,123,46,50,162,140,"jpeg, png",0
korat,152,147,117,93,50,162,140,"jpeg, png",0
safari,150,158,105,79,38,300,140,"jpeg, png",0
american_bobtail,140,155,114,64,48,162,140,jpeg,0
mekong_bobtail,140,149,118,44,50,162,140,"jpeg, png",0
korean_bobtail,139,140,129,63,81,162,140,jpeg,0
siamese,139,152,115,78,53,300,140,"jpeg, png",0
somali,139,150,112,61,53,162,140,jpeg,0
devon_rex,138,150,116,78,55,162,140,jpeg,0
american_curl,138,155,112,91,51,300,140,"jpeg, png",0
ural_rex,137,144,121,48,65,162,140,"jpeg, png",0
singapura,136,158,109,93,54,300,140,"jpeg, png",0
ukrainian_levkoy,134,140,123,78,59,162,140,jpeg,0
maine_coon,133,141,122,66,79,162,140,jpeg,0
birman,131,155,113,93,56,162,140,jpeg,0
oregon_rex,121,147,123,85,66,300,140,"jpeg, png",0
kurilian_bobtail,120,149,121,78,81,162,140,jpeg,0
laperm,120,149,116,78,55,162,140,"jpeg, png",0
sphynx,120,149,117,47,63,162,140,jpeg,0
chartreux,114,146,117,75,61,162,140,jpeg,0
russian_blue,109,152,111,88,36,162,140,jpeg,0
norwegian_forest_cat,97,150,114,78,56,162,140,jpeg,0
1 class image_count avg_width avg_height min_width min_height max_width max_height formats corrupt_files
2 abyssinian 200 148 117 88 46 162 140 jpeg, png 0
3 cyprus 200 159 100 93 40 162 140 jpeg, png 0
4 lykoi 200 142 122 68 55 162 140 jpeg, png 0
5 donskoy 200 141 123 67 54 162 140 jpeg 0
6 chausie 200 143 119 78 46 162 140 jpeg, png 0
7 european_shorthair 200 149 118 70 49 162 140 jpeg, png 0
8 turkish_van 200 149 120 87 56 162 140 jpeg, png 0
9 pixie_bob 200 147 119 78 49 162 140 jpeg, png 0
10 ragdoll 199 149 116 78 50 162 140 jpeg 0
11 german_rex 199 138 125 72 50 162 140 jpeg, png 0
12 american_shorthair 199 147 114 65 50 162 140 jpeg, png 0
13 sokoke 199 145 122 70 49 162 140 jpeg, png 0
14 khao_manee 198 144 120 78 49 162 140 jpeg, png 0
15 thai 198 143 118 68 60 162 140 jpeg, png 0
16 cymric 198 148 120 63 43 162 140 jpeg, png 0
17 oriental_shorthair 197 141 123 78 54 162 140 jpeg 0
18 cornish_rex 197 147 117 75 56 162 140 jpeg, png 0
19 burmese 197 147 117 45 58 162 140 jpeg, png 0
20 savannah 197 153 106 78 40 162 140 jpeg, png 0
21 american_wirehair 196 149 117 76 50 162 140 jpeg, png 0
22 peterbald 196 144 123 77 50 162 140 jpeg, png 0
23 karelian_bobtail 196 145 125 67 65 162 140 jpeg, png 0
24 tonkinese 195 148 116 70 53 162 140 jpeg, png 0
25 balinese 195 154 112 79 54 162 140 jpeg 0
26 japanese_bobtail 194 149 118 87 49 162 140 jpeg, png 0
27 nebelung 194 143 119 64 46 162 140 jpeg 0
28 selkirk_rex 192 144 122 78 76 162 140 jpeg 0
29 persian 192 150 114 78 46 162 140 jpeg 0
30 manx 192 155 113 86 50 162 140 jpeg, png 0
31 himalayan 192 158 109 93 40 162 140 jpeg 0
32 munchkin 191 147 117 61 48 162 140 jpeg 0
33 bengal 189 151 115 78 71 162 140 jpeg 0
34 turkish_angora 188 145 119 66 52 162 140 jpeg 0
35 vankedisi 187 147 116 63 47 162 140 jpeg, png 0
36 scottish_fold 184 143 120 78 50 162 140 jpeg 0
37 egyptian_mau 184 144 121 72 74 162 140 jpeg, png 0
38 ocicat 182 150 115 62 44 162 140 jpeg, png 0
39 ragamuffin 182 149 116 78 44 162 140 jpeg, png 0
40 serengeti 175 159 100 78 53 300 140 jpeg, png 0
41 british_shorthair 174 148 117 78 50 162 140 jpeg 0
42 toyger 160 145 120 78 50 162 140 jpeg, png 0
43 siberian 159 153 116 78 55 162 140 jpeg 0
44 havana_brown 159 133 127 75 34 300 140 jpeg, png 0
45 exotic_shorthair 157 148 118 93 55 300 140 jpeg, png 0
46 bombay 154 139 123 46 50 162 140 jpeg, png 0
47 korat 152 147 117 93 50 162 140 jpeg, png 0
48 safari 150 158 105 79 38 300 140 jpeg, png 0
49 american_bobtail 140 155 114 64 48 162 140 jpeg 0
50 mekong_bobtail 140 149 118 44 50 162 140 jpeg, png 0
51 korean_bobtail 139 140 129 63 81 162 140 jpeg 0
52 siamese 139 152 115 78 53 300 140 jpeg, png 0
53 somali 139 150 112 61 53 162 140 jpeg 0
54 devon_rex 138 150 116 78 55 162 140 jpeg 0
55 american_curl 138 155 112 91 51 300 140 jpeg, png 0
56 ural_rex 137 144 121 48 65 162 140 jpeg, png 0
57 singapura 136 158 109 93 54 300 140 jpeg, png 0
58 ukrainian_levkoy 134 140 123 78 59 162 140 jpeg 0
59 maine_coon 133 141 122 66 79 162 140 jpeg 0
60 birman 131 155 113 93 56 162 140 jpeg 0
61 oregon_rex 121 147 123 85 66 300 140 jpeg, png 0
62 kurilian_bobtail 120 149 121 78 81 162 140 jpeg 0
63 laperm 120 149 116 78 55 162 140 jpeg, png 0
64 sphynx 120 149 117 47 63 162 140 jpeg 0
65 chartreux 114 146 117 75 61 162 140 jpeg 0
66 russian_blue 109 152 111 88 36 162 140 jpeg 0
67 norwegian_forest_cat 97 150 114 78 56 162 140 jpeg 0
+50
View File
@@ -0,0 +1,50 @@
import cv2
import glob
import torch
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
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))
im_list = glob.glob("./dataset/test/*/*.jpg")
np.random.shuffle(im_list)
net.to(device)
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:
net.eval()
im_data = Image.open(im_path)
inputs = test_transform(im_data)
inputs = torch.unsqueeze(inputs, dim=0)
inputs = inputs.to(device)
outputs = net.forward(inputs)
_, 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)
if __name__ == "__main__":
test()
+55
View File
@@ -0,0 +1,55 @@
import cv2
import glob
import torch
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
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))
im_list = glob.glob("./dataset/test/*/*.jpg")
np.random.shuffle(im_list)
net.to(device)
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:
net.eval()
im_data = Image.open(im_path)
inputs = test_transform(im_data)
inputs = torch.unsqueeze(inputs, dim=0)
inputs = inputs.to(device)
outputs = net.forward(inputs)
# print("outputs", outputs)
_, 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)
# 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)
cv2.imshow("img", img)
cv2.waitKey(0)
if __name__ == "__main__":
test()
+51
View File
@@ -0,0 +1,51 @@
import os
import torch
from core.nets.resnet import resnet
from core.dataloader.dataloader import train_dataloader
from core.const import epoch, lr, batch_size
def train():
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device: ", device)
net = resnet().to(device)
loss_func = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)
for e in range(epoch):
print("epoch: ", e)
net.train()
for i, data in enumerate(train_dataloader):
inputs, labels = data
inputs, labels = inputs.to(device), labels.to(device)
outputs = net(inputs)
loss = loss_func(outputs, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
_, pred = torch.max(outputs, dim=1)
correct = pred.eq(labels.data).cpu().sum()
print("step: ", i, "loss: ", loss.item(), "correct: ", 1.0 * correct / batch_size)
scheduler.step()
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
if not os.path.exists("./model"):
os.makedirs("./model")
torch.save(net.state_dict(), "./model/resnet_epoch_{}.pth".format(e + 1))
if __name__ == "__main__":
train()
+51
View File
@@ -0,0 +1,51 @@
import os
import torch
from core.nets.resnet18 import resnet18
from core.dataloader.dataloader import train_dataloader
from core.const import epoch, lr, batch_size
def train():
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device: ", device)
net = resnet18().to(device)
loss_func = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6) # 余弦退火
for e in range(epoch):
print("epoch: ", e)
net.train()
for i, data in enumerate(train_dataloader):
inputs, labels = data
inputs, labels = inputs.to(device), labels.to(device)
outputs = net(inputs)
loss = loss_func(outputs, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
_, pred = torch.max(outputs, dim=1)
correct = pred.eq(labels.data).cpu().sum()
print("step: ", i, "loss: ", loss.item(), "correct: ", 1.0 * correct / batch_size)
scheduler.step()
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
if not os.path.exists("./models"):
os.makedirs("./models")
torch.save(net.state_dict(), "./models/resnet18_epoch_{}.pth".format(e + 1))
if __name__ == "__main__":
train()