pytorch框架实现图像分类任务pipeline总结(以猫狗二分类为例)
pipeline:
1.数据准备以及处理
2.模型准备(resnet50)
3.训练、推理脚本
(attention:本文实现的代码可能会和其他文章的代码有不太一样的地方,但是只要把整个逻辑梳理清楚,用什么样的数据结构使用代码就看个人习惯了)
1.数据准备以及处理
首先了解一下Dataset和Dataloader类:
类所在库:from torch.utils.data import Dataset,DataLoader
一个Dataset类的对象包含了数据集中的数据和该数据相对应的标签,一个Dataloader类的对象是一个迭代器,将Dataset类对象中的数据和对象的标签按照批次组织起来,便于数据读取。
情况一:(需要自己手写Dataset类)
训练数据的目录被整理为以下格式:
train
----cat0.jpg
----cat1.jpg
----cat2.jpg
----cat3.jpg
.....
----dog0.jpg
----dog1.jpg
----dog2.jpg
----dog3.jpg
可以看到训练数据在一个train目录下,每张图片的文件名表明了这张图片所属于的类别。这种情况下就需要我们重写Dataset类(必须重写,因为Dataset类是抽象类,不能实例化对象),需要重新实现三个方法:
__init__
__getitem__
__len__
init方法是初始化方法,可以把对数据集的预处理放在这个方法里,一般预处理的流程代码如下:
def __init__(self, path):
print("building dataset.......")
#图像预处理,也可以作为对象传入init方法,主要就是写死或者不写死的问题
self.transform = transforms.Compose([
transforms.CenterCrop(224),
transforms.Resize([224, 224]),
transforms.ToTensor(),
transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)) # 归一化
])
#创建数据列表和标签列表
self.data = []
self.label = []
#从数据集目录下一一读取数据并为数据打标签,数据和标签的索引要一一对应
for filename in os.listdir(path):
if os.path.isfile(os.path.join(path, filename)):
#处理标签
if filename.startswith('c'):
self.label.append(torch.tensor(1))
elif filename.startswith('d'):
self.label.append(torch.tensor(0))
#处理数据
img_file = os.path.join(path, filename)
img = Image.open(img_file) #从磁盘加载图像文件
img = self.transform(img) #将数据变为张量
self.data.append(img)
print("building dataset success!!!")
getitem方法主要是对了便于数据读取的,比如创建了一个数据集对象dataset,根据上面的代码,你可以通过dataset.data[0],dataset.label[0]来取第一个数据和它的标签,这样有点麻烦,getitem方法可以直接通过返回元组的形式获取数据和对应的标签,即dataset[0]就可以直接返回。
len方法主要是为了返回数据集的长度
之后有了dataset对象,我们就可以使用Dataloader类进行对dataset对象中的数据进行分批处理,Dataloader迭代器对象中的每个元素也是一个元组,不过元组中的数据和标签都是按批次加了一个维度,代码如下:
train_iter = DataLoader(dataset, batch_size=32, shuffle=True)
情况二:(不需要自己手写Dataset类)
训练数据的目录被整理为以下格式:
train
---cat
---dog
如果数据集以这种方式处理的话,我们直接调用 torchvision.datasets模块中的ImageFolder类,让该类实例化数据集,代码如下:
train_data = datasets.ImageFolder("C:\\Users\\25190\\Desktop\\NLP作业\\作业一\\flower_data",transform)
train_data对象可以理解为一个Dataset类对象,省去了我们重新书写Dataset类的繁琐,可以直接传入Dataloader类,这个形式很方便。
2.模型准备(resnet50)
本次二分类使用torchvision库中的resnet50模型,由于resnet50做的事1000分类的任务,所以我们应该把模型的最后一个的out_features改为我们任务的分类数(即2),方法和流程如下代码所示:
resnet = torchvision.models.resnet50(pretrained=True)
input_f = resnet.fc.in_features
resnet.fc = nn.Linear(input_f, 5)
3.训练、推理脚本
接下来就是训练模型了,代码如下:
resnet.train()
for epoch in range(epochs):
n = 0
acc_num = 0.0
loss_sum = 0.0
for X,y in train_iter:
batch_count = 0
X = X.to(device)
y = y.to(device)
y_hat = resnet(X)
acc_num += (y_hat.argmax(dim=1) == y).sum().cpu().item()
loss = loss_func(y_hat, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
n += y.shape[0]
loss_sum += loss.cpu().item()
batch_count += 1
print("epoch: %d, loss: %.5f, train_acc: %.5f" % (epoch, loss_sum / batch_count, acc_num / n))
推理:
resnet.eval()
acc_num = 0.0
n = 0
for X,y in test_iter:
X = X.to(device)
y = y.to(device)
acc_num += (resnet(X).argmax(dim=1) == y).sum().cpu().item()
n += y.shape[0]
print("acc: %.5f" % (acc_num / n))
更多推荐
所有评论(0)