first commit
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+81
@@ -0,0 +1,81 @@
|
||||
#训练多少轮
|
||||
epoch = 50
|
||||
|
||||
#学习率
|
||||
lr = 0.0005
|
||||
|
||||
#一轮多少张图片
|
||||
batch_size = 64
|
||||
|
||||
#训练图片输入尺寸
|
||||
input_size = 128
|
||||
|
||||
#分类
|
||||
label_name = [
|
||||
"abyssinian",
|
||||
"cyprus",
|
||||
"lykoi",
|
||||
"donskoy",
|
||||
"chausie",
|
||||
"european_shorthair",
|
||||
"turkish_van",
|
||||
"pixie_bob",
|
||||
"ragdoll",
|
||||
"german_rex",
|
||||
"american_shorthair",
|
||||
"sokoke",
|
||||
"khao_manee",
|
||||
"thai",
|
||||
"cymric",
|
||||
"oriental_shorthair",
|
||||
"cornish_rex",
|
||||
"burmese",
|
||||
"savannah",
|
||||
"american_wirehair",
|
||||
"peterbald",
|
||||
"karelian_bobtail",
|
||||
"tonkinese",
|
||||
"balinese",
|
||||
"japanese_bobtail",
|
||||
"nebelung",
|
||||
"selkirk_rex",
|
||||
"persian",
|
||||
"manx",
|
||||
"himalayan",
|
||||
"munchkin",
|
||||
"bengal",
|
||||
"turkish_angora",
|
||||
"vankedisi",
|
||||
"scottish_fold",
|
||||
"egyptian_mau",
|
||||
"ocicat",
|
||||
"ragamuffin",
|
||||
"serengeti",
|
||||
"british_shorthair",
|
||||
"toyger",
|
||||
"siberian",
|
||||
"havana_brown",
|
||||
"exotic_shorthair",
|
||||
"bombay",
|
||||
"korat",
|
||||
"safari",
|
||||
"american_bobtail",
|
||||
"mekong_bobtail",
|
||||
"korean_bobtail",
|
||||
"siamese",
|
||||
"somali",
|
||||
"devon_rex",
|
||||
"american_curl",
|
||||
"ural_rex",
|
||||
"singapura",
|
||||
"ukrainian_levkoy",
|
||||
"maine_coon",
|
||||
"birman",
|
||||
"oregon_rex",
|
||||
"kurilian_bobtail",
|
||||
"laperm",
|
||||
"sphynx",
|
||||
"chartreux",
|
||||
"russian_blue",
|
||||
"norwegian_forest_cat",
|
||||
]
|
||||
@@ -0,0 +1,85 @@
|
||||
import os
|
||||
import glob
|
||||
from torchvision import transforms
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from PIL import Image
|
||||
from 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):
|
||||
img = Image.open(path).convert('RGB')
|
||||
w, h = img.size
|
||||
|
||||
if w > 200: # 如果宽度超过200,可能是错误数据,等比缩放到162
|
||||
ratio = 162 / w
|
||||
new_h = int(h * ratio)
|
||||
img = img.resize((162, new_h), Image.BILINEAR)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
im_train_list = glob.glob("dataset/train/*/*.jpg")
|
||||
im_test_list = glob.glob("dataset/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)
|
||||
Binary file not shown.
@@ -0,0 +1,82 @@
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
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, num_classes=67):
|
||||
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()
|
||||
@@ -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
|
||||
|
+50
@@ -0,0 +1,50 @@
|
||||
import cv2
|
||||
import glob
|
||||
import torch
|
||||
from torchvision import transforms
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
# from resnet import resnet
|
||||
from torchvision.models import resnet18
|
||||
from const import label_name, input_size
|
||||
|
||||
|
||||
def test():
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(device)
|
||||
|
||||
net = resnet18(weights=None)
|
||||
net.load_state_dict(torch.load("./model/resnet_epoch_31.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, input_size)),
|
||||
transforms.ToTensor()
|
||||
])
|
||||
|
||||
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)
|
||||
|
||||
img = np.asarray(im_data)
|
||||
img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)
|
||||
cv2.imshow("img", img)
|
||||
cv2.waitKey(0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test()
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
import os
|
||||
import torch
|
||||
# from resnet import resnet
|
||||
from torchvision.models import resnet18
|
||||
from dataloader import train_dataloader
|
||||
from 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)
|
||||
net = resnet18(weights=None).to(device)
|
||||
|
||||
loss_func = torch.nn.CrossEntropyLoss()
|
||||
|
||||
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
|
||||
|
||||
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, 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()
|
||||
Reference in New Issue
Block a user