栏目分类:
子分类:
返回
名师互学网用户登录
快速导航关闭
当前搜索
当前分类
子分类
实用工具
热门搜索
名师互学网 > IT > 软件开发 > 后端开发 > Python

【无标题】

Python 更新时间: 发布时间: IT归档 最新发布 模块sitemap 名妆网 法律咨询 聚返吧 英语巴士网 伯小乐 网商动力

【无标题】

import os#操作系统接口模块
import torch
import torch.nn as nn#涵盖深度学习网络模型搭建和参数优化过程中的常用内容
import torch.nn.functional as F#functional 是以函数的方式实现的,nn 是以类的方式实现的
import torch.optim as optim#参数自动优化
import torchvision.transforms as transforms#用于对载入数据的进行变换
from torch.utils.data import DataLoader, Dataset#用于数据集的装载
from PIL import Image#PIL.Image读取图像
import argparse

parse = argparse.ArgumentParser(description='Params for training. ')

数据集根目录

parse.add_argument(’–root’, type=str, default=’/home/wn/data’, help=‘path to data set’)

模式,3选1

parse.add_argument(’–mode’, type=str, default=‘train’, choices=[‘train’, ‘validation’, ‘inference’])

checkpoint 路径

parse.add_argument(’–log_path’, type=str, default=os.path.abspath(’.’) + ‘/log.pth’, help=‘dir of checkpoints’)

parse.add_argument(’–restore’, type=bool, default=True, help=‘whether to restore checkpoints’)

parse.add_argument(’–batch_size’, type=int, default=16, help=‘size of mini-batch’)
parse.add_argument(’–image_size’, type=int, default=64, help=‘resize image’)
parse.add_argument(’–epoch’, type=int, default=100)
#数据集类别数是3755,所以给定了一个选择范围
parse.add_argument(’–num_class’, type=int, default=100, choices=range(10, 50))
args = parse.parse_args()

#自定义数据集
“”"
重写 Dataset 里的 init, getitem, len
__getitem__在训练的时候返回输入网络的数据,图片和标签等
len 返回数据集长度
“”"
class MyDataset(Dataset):
def init(self, txt_path, num_class, transforms=None):
super(MyDataset, self).init()
images = []# 存储图片路径
labels = []# 存储类别名,在本例中是数字
# 打开生成的txt文件
with open(txt_path, ‘r’) as f:
for line in f:
if int(line.split(’/’)[-2]) >= num_class: # 只读取前 num_class 个类
break
line = line.strip(’n’)
images.append(line)
labels.append(int(line.split(’/’)[-2]))
self.images = images
self.labels = labels
self.transforms = transforms# 图片需要进行的变换,ToTensor()等等

def __getitem__(self, index):
    image = Image.open(self.images[index]).convert('RGB') # 用PIL.Image读取图像
    label = self.labels[index]
    if self.transforms is not None:
        image = self.transforms(image)
    return image, label

def __len__(self):
    return len(self.labels)

#改进后的简单神经网络:四层卷积,三层全连接
class NetBig(nn.Module):
def init(self):
super(NetBig, self).init()
self.conv1 = nn.Conv2d(in_channels=1, out_channels=64, kernel_size=3, padding=1)
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1)
self.conv3 = nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=1)
self.conv4 = nn.Conv2d(in_channels=256, out_channels=512, kernel_size=3, padding=1)
self.fc1 = nn.Linear(8192, 4096)
self.fc2 = nn.Linear(4096, 1024)
self.fc3 = nn.Linear(1024, args.num_class)
# self.dropout = nn.Dropout(p=0.8)

def forward(self, x):
    x = self.pool(F.relu(self.conv1(x)))
    x = self.pool(F.relu(self.conv2(x)))
    x = self.pool(F.relu(self.conv3(x)))
    x = self.pool(F.relu(self.conv4(x)))
    x = x.view(-1, self.num_flat_features(x))
    x = F.relu(self.fc1(x))
    x = F.relu(self.fc2(x))
    x = self.fc3(x)
    return x

@staticmethod
def num_flat_features(x):
    size = x.size()[1:]
    num_features = 1
    for s in size:
        num_features *= s
    return num_features

#简单的网络:两层卷积,三层全连接,20个类别的情况下可以训练至95%以上的准确率
class NetSmall(nn.Module):
def init(self):
super(NetSmall, self).init()
self.conv1 = nn.Conv2d(1, 6, 3)# 3个参数分别是in_channels,out_channels,kernel_size,还可以加padding
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(6, 16, 5)
self.fc1 = nn.Linear(2704, 512)
self.fc2 = nn.Linear(512, 84)
self.fc3 = nn.Linear(84, args.num_class)

def forward(self, x):
    x = self.pool(F.relu(self.conv1(x)))
    x = self.pool(F.relu(self.conv2(x)))
    x = x.view(-1, self.num_flat_features(x))
    x = F.relu(self.fc1(x))
    x = F.relu(self.fc2(x))
    x = self.fc3(x)
    return x

@staticmethod
def num_flat_features(x):
    size = x.size()[1:]
    num_features = 1
    for s in size:
        num_features *= s
    return num_features

def train():
#由于数据集图片尺寸不一,因此要进行resize,这里还可以加入数据增强,灰度变换,随机剪切等等
transform = transforms.Compose([transforms.Resize((args.image_size, args.image_size)),
transforms.Grayscale(),
transforms.ToTensor()])

train_set = MyDataset(args.root + '/train.txt', num_class=args.num_class, transforms=transform)
train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=True)
# 选择使用的设备
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
print(device)

model = NetSmall()
model.to(device)
# 训练模式
model.train()

criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

if args.restore:
    checkpoint = torch.load(args.log_path)
    model.load_state_dict(checkpoint['model_state_dict'])
    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
    loss = checkpoint['loss']
    epoch = checkpoint['epoch']
else:
    loss = 0.0
    epoch = 0

while epoch < args.epoch:
    running_loss = 0.0

    for i, data in enumerate(train_loader):
        inputs, labels = data[0].to(device), data[1].to(device)

        optimizer.zero_grad()
        outs = model(inputs)
        loss = criterion(outs, labels)
        loss.backward()
        optimizer.step()

        running_loss += loss.item()

        if i % 200 == 199:  # every 200 steps
            print('epoch %5d: batch: %5d, loss: %f' % (epoch + 1, i + 1, running_loss / 200))
            running_loss = 0.0
    # 保存 checkpoint
    if epoch % 10 == 9:
        print('Save checkpoint...')
        torch.save({'epoch': epoch,
                    'model_state_dict': model.state_dict(),
                    'optimizer_state_dict': optimizer.state_dict(),
                    'loss': loss},
                   args.log_path)
    epoch += 1

print('Finish training')

def validation():
transform = transforms.Compose([transforms.Resize((args.image_size, args.image_size)),
transforms.Grayscale(),
transforms.ToTensor()])

test_set = MyDataset(args.root + '/test.txt', num_class=args.num_class, transforms=transform)
test_loader = DataLoader(test_set, batch_size=args.batch_size)

device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
model = NetSmall()
model.to(device)

checkpoint = torch.load(args.log_path)
model.load_state_dict(checkpoint['model_state_dict'])

model.eval()

total = 0.0
correct = 0.0
with torch.no_grad():
    for i, data in enumerate(test_loader):
        inputs, labels = data[0].cuda(), data[1].cuda()
        outputs = model(inputs)
        _, predict = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += sum(int(predict == labels)).item()

        if i % 100 == 99:
            print('batch: %5d,t acc: %f' % (i + 1, correct / total))
print('Accuracy: %.2f%%' % (correct / total * 100))

def inference():
print(‘Start inference…’)
transform = transforms.Compose([transforms.Resize((args.image_size, args.image_size)),
transforms.Grayscale(),
transforms.ToTensor()])

f = open(args.root + '/test.txt')
num_line = sum(line.count('n') for line in f)
f.seek(0, 0)
line = int(torch.rand(1).data * num_line - 10) # -10 for 'n's are more than lines
while line > 0:
    f.readline()
    line -= 1
img_path = f.readline().rstrip('n')
f.close()
label = int(img_path.split('/')[-2])
print('label:t%4d' % label)
input = Image.open(img_path).convert('RGB')
input = transform(input)
input = input.unsqueeze(0)
model = NetSmall()
model.eval()
checkpoint = torch.load(args.log_path)
model.load_state_dict(checkpoint['model_state_dict'])
output = model(input)
_, pred = torch.max(output.data, 1)

print('predict:t%4d' % pred)

#提取图片路径
def classes_txt(root, out_path, num_class=None):
‘’’
write image paths (containing class name) into a txt file.
:param root: data set path
:param out_path: txt file path
:param num_class: how many classes needed
:return: None
‘’’
dirs = os.listdir(root)
if not num_class:
num_class = len(dirs)

if not os.path.exists(out_path):
    f = open(out_path, 'w')
    f.close()

with open(out_path, 'r+') as f:
    try:
        end = int(f.readlines()[-1].split('/')[-2]) + 1
    except:
        end = 0
    if end < num_class - 1:
        dirs.sort()
        dirs = dirs[end:num_class]
        for dir in dirs:
            files = os.listdir(os.path.join(root, dir))
            for file in files:
                f.write(os.path.join(root, dir, file) + 'n')

if name == ‘main’:

classes_txt(args.root + '/train', args.root + '/train.txt', num_class=args.num_class)
classes_txt(args.root + '/test', args.root + '/test.txt', num_class=args.num_class)

if args.mode == 'train':
    train()
elif args.mode == 'validation':
    validation()
elif args.mode == 'inference':
    inference()
转载请注明:文章转载自 www.mshxw.com
本文地址:https://www.mshxw.com/it/656505.html
我们一直用心在做
关于我们 文章归档 网站地图 联系我们

版权所有 (c)2021-2022 MSHXW.COM

ICP备案号:晋ICP备2021003244-6号