工程化
This commit is contained in:
+16
-6
@@ -1,3 +1,8 @@
|
|||||||
|
__pycache__/
|
||||||
|
*.py[cod]
|
||||||
|
*.pyo
|
||||||
|
*.pyd
|
||||||
|
|
||||||
*.exe
|
*.exe
|
||||||
*.exe~
|
*.exe~
|
||||||
*.dll
|
*.dll
|
||||||
@@ -8,14 +13,19 @@
|
|||||||
*.log
|
*.log
|
||||||
*.txt
|
*.txt
|
||||||
|
|
||||||
tmp/
|
venv/
|
||||||
.claude/
|
env/
|
||||||
|
.venv/
|
||||||
|
|
||||||
.vscode/
|
.vscode/
|
||||||
|
.idea/
|
||||||
|
.claude/
|
||||||
|
|
||||||
|
nn/dataset/toy/*
|
||||||
|
nn/dataset/benchmark/*
|
||||||
|
nn/models/*.pth
|
||||||
|
nn/models/*.pt
|
||||||
|
|
||||||
build/bin
|
build/bin
|
||||||
frontend/node_modules
|
frontend/node_modules
|
||||||
frontend/dist
|
frontend/dist
|
||||||
|
|
||||||
nn/dataset/*
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
# Neural Network Package
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
-81
@@ -1,81 +0,0 @@
|
|||||||
#训练多少轮
|
|
||||||
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,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)
|
||||||
@@ -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 const import label_name, input_size, batch_size
|
from nn.const import label_name, input_size, batch_size
|
||||||
|
|
||||||
|
|
||||||
label_dict = {}
|
label_dict = {}
|
||||||
@@ -13,21 +13,13 @@ for idx, name in enumerate(label_name):
|
|||||||
|
|
||||||
|
|
||||||
def default_loader(path):
|
def default_loader(path):
|
||||||
img = Image.open(path).convert('RGB')
|
return 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([
|
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), # 轻微旋转
|
||||||
transforms.ColorJitter(brightness=0.1, contrast=0.1),
|
transforms.ColorJitter(brightness=0.1, contrast=0.1),
|
||||||
transforms.ToTensor(),
|
transforms.ToTensor(),
|
||||||
@@ -69,8 +61,9 @@ class MyDataset(Dataset):
|
|||||||
return len(self.imgs)
|
return len(self.imgs)
|
||||||
|
|
||||||
|
|
||||||
im_train_list = glob.glob("dataset/train/*/*.jpg")
|
dataset_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||||
im_test_list = glob.glob("dataset/test/*/*.jpg")
|
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)
|
train_dataset = MyDataset(im_train_list, transform=train_transform)
|
||||||
Binary file not shown.
@@ -1,5 +1,6 @@
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
from nn.const import num_classes
|
||||||
|
|
||||||
|
|
||||||
class ResBlock(nn.Module):
|
class ResBlock(nn.Module):
|
||||||
@@ -44,7 +45,7 @@ class ResNet(nn.Module):
|
|||||||
|
|
||||||
return nn.Sequential(*layer_list)
|
return nn.Sequential(*layer_list)
|
||||||
|
|
||||||
def __init__(self, num_classes=67):
|
def __init__(self):
|
||||||
super(ResNet, self).__init__()
|
super(ResNet, self).__init__()
|
||||||
self.in_channel = 32
|
self.in_channel = 32
|
||||||
|
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
import torch.nn as nn
|
||||||
|
from torchvision import models
|
||||||
|
from nn.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()
|
||||||
@@ -4,17 +4,16 @@ import torch
|
|||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
import numpy as np
|
import numpy as np
|
||||||
# from resnet import resnet
|
from nn.nets.resnet import resnet
|
||||||
from torchvision.models import resnet18
|
from nn.const import label_name, input_size
|
||||||
from const import 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 = resnet18(weights=None)
|
net = resnet()
|
||||||
net.load_state_dict(torch.load("./model/resnet_epoch_31.pth", weights_only=True))
|
net.load_state_dict(torch.load("./model/resnet_epoch_15.pth", weights_only=True))
|
||||||
|
|
||||||
im_list = glob.glob("./dataset/test/*/*.jpg")
|
im_list = glob.glob("./dataset/test/*/*.jpg")
|
||||||
np.random.shuffle(im_list)
|
np.random.shuffle(im_list)
|
||||||
@@ -22,8 +21,10 @@ def test():
|
|||||||
net.to(device)
|
net.to(device)
|
||||||
|
|
||||||
test_transform = transforms.Compose([
|
test_transform = transforms.Compose([
|
||||||
transforms.Resize((input_size, input_size)),
|
transforms.Resize(input_size),
|
||||||
transforms.ToTensor()
|
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:
|
for im_path in im_list:
|
||||||
@@ -35,7 +36,6 @@ def test():
|
|||||||
|
|
||||||
inputs = inputs.to(device)
|
inputs = inputs.to(device)
|
||||||
outputs = net.forward(inputs)
|
outputs = net.forward(inputs)
|
||||||
print("outputs", outputs)
|
|
||||||
|
|
||||||
_, 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)
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
import cv2
|
||||||
|
import glob
|
||||||
|
import torch
|
||||||
|
from torchvision import transforms
|
||||||
|
from PIL import Image
|
||||||
|
import numpy as np
|
||||||
|
from nn.nets.resnet18 import resnet18
|
||||||
|
from nn.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()
|
||||||
@@ -1,23 +1,21 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
# from resnet import resnet
|
from nn.nets.resnet import resnet
|
||||||
from torchvision.models import resnet18
|
from nn.dataloader.dataloader import train_dataloader
|
||||||
from dataloader import train_dataloader
|
from nn.const import epoch, lr, batch_size
|
||||||
from const import epoch, lr, batch_size
|
|
||||||
|
|
||||||
|
|
||||||
def train():
|
def train():
|
||||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
print("device: ", device)
|
print("device: ", device)
|
||||||
|
|
||||||
# net = resnet().to(device)
|
net = resnet().to(device)
|
||||||
net = resnet18(weights=None).to(device)
|
|
||||||
|
|
||||||
loss_func = torch.nn.CrossEntropyLoss()
|
loss_func = torch.nn.CrossEntropyLoss()
|
||||||
|
|
||||||
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
|
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
|
||||||
|
|
||||||
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5)
|
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)
|
||||||
|
|
||||||
for e in range(epoch):
|
for e in range(epoch):
|
||||||
print("epoch: ", e)
|
print("epoch: ", e)
|
||||||
@@ -0,0 +1,51 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
from nn.nets.resnet18 import resnet18
|
||||||
|
from nn.dataloader.dataloader import train_dataloader
|
||||||
|
from nn.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()
|
||||||
Reference in New Issue
Block a user