1、项目简介

1.1 项目名称

基于DesNet169的鸟分类研究

1.2 项目简介

该项目主要为基于DesNet对6种鸟类进行分类研究。本研究数据来源为kaggle数据集(Bird Speciees Dataset | Kaggle)。本研究基于DesNet网络基础上进行调整,与Resne网络相比DesNet的参数量更少,因此运行速度也更快。此外,项目还将探索数据增强、迁移学习等技术手段,以提高模型的泛化能力和鲁棒性,最终实现高精度的鸟分类。

2、数据

2.1 数据来源

数据来源为kaggle数据集 bird-speciees-dataset (kaggle.com)

该数据集中主要为针对鸟的6种小颗粒特征分类,其中数据集分为American goldfinch、barn owl、carmine bee-eater、downy woodpecker、emperor penguin、flamingo以及相对应的image图像。

2.2 数据切分

该数据集中未区分训练集与验证集,因此需要对数据进行切分。按照数据的label标签进行训练集、验证集的切分,保证数据的均匀性,同时将切分好的data数据分别生成train_data、test_data两个txt文本,便于后续使用。

import os
import torch
from sklearn.model_selection import train_test_split
from torch.utils.data import Dataset, DataLoader
from PIL import Image
​
def delete_datasets(base_path, test_size=0.2,random_state=42):
    # 存储所有类别的训练集文件路径
    train_files = []
    test_files = []
    # 遍历每个类别文件夹
    for class_name in os.listdir(base_path):   #os.listdir(base_path):获取根目录下的所有文件和文件夹
        # print(class_name)  # 6种鸟类名称
        class_path = os.path.join(base_path, class_name)
        # 检查class_path是否是一个目录(文件夹)
        if not os.path.isdir(class_path):
            continue
        # 获取当前类别的所有文件
        files = [os.path.join(class_path, file) for file in os.listdir(class_path)]
        # 划分训练集和测试集
        train_data, test_data = train_test_split(files, test_size=test_size, random_state=random_state)
        # 将结果添加到总列表
        train_files.extend(train_data)  
        test_files.extend(test_data)
    return train_files, test_files
​
def save_file_list(file_list, output_path):
    with open(output_path, 'w', encoding='utf-8') as f:
        for file in file_list:
            f.write(file + '\n')
            
def load_file_list(file_path):
    with open(file_path, 'r') as f:
        file_list = [line.strip() for line in f.readlines()]
    return file_list
​
if __name__ == '__main__':
    train_files, test_files = delete_datasets(base_path='D:\HQYJ\py\code\pytorch/data/Bird')
    # 保存训练集和测试集文件列表
    save_file_list(train_files, 'train_files.txt')
    save_file_list(test_files, 'test_files.txt')

2.3 数据处理

2.3.1 数据加载器

1.通过自定义 Dataset,直接从 train_files.txt 和 test_files.txt 中读取文件路径,并加载对应的图片数据。

2.采用整数索引将类别名称转换为对应的整数索引。

3.采用数据增强器对图像数据进行处理,增加图像的翻转、旋转、标准化等相关操作,增强模 型的鲁棒性。

4.在dataset数据回传过程中采用Image进行图像读取,对应位置返回,完成数据集加载。

import os
import torch
from sklearn.model_selection import train_test_split
from torch.utils.data import Dataset, DataLoader
from PIL import Image
​
​
​
class BirdDataset(Dataset):
    def __init__(self,file_list,transform=None,class_to_idx=None):
        super(BirdDataset, self).__init__()
        self.file_list = file_list
        self.transform = transform
        self.class_to_idx = {cls: idx for idx, cls in enumerate(sorted(set([os.path.basename(os.path.dirname(file)) for file in file_list])))}
​
    def __len__(self):
        return len(self.file_list)
​
    def __getitem__(self, index):
        file_path = self.file_list[index]
        # 加载图片
        image = Image.open(file_path).convert('RGB')
        # 获取标签
        label_name = os.path.basename(os.path.dirname(file_path))
        label = self.class_to_idx[label_name]
        if self.transform is not None:
            image = self.transform(image)
        return image,torch.tensor(label)

2.3.2 数据增强器

为区分训练集与验证集数据增强器,在dataset构建过程中添加控制变量,对数据加强器进行变更。同时对数据增强处理时,数据增强的顺序具有一定的重要性。在对数据集进行加载后,首先将数据转换为Tensor,然后根据转换后的Tensor数据进行相关变换操作,最后进行数据的Normalize,将有助于 数据特征提取与数据分类,提高模型训练与预测的精度,降低训练时间与时长。

transform = transforms.Compose([
        transforms.RandomHorizontalFlip(), # 随机水平翻转
        transforms.RandomVerticalFlip(p=0.5),  # 以50%的概率随机垂直翻转图像
        transforms.Resize((224, 224)),  # 调整图片大小
        transforms.ToTensor() , # 转换为张量
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])

3、神经网络

本项目采用的基础网络模型为DesNet169,在此基础上分别进行模型改进,对小颗粒图像特征的分类 进行效果进行对比。

1、对图像输入7x7卷积核进行修改,采用3*3并行通道处理的方式进行图像特征的提取与输入;

2、在DesNet169基础训练模型基础上修改classifier使其输出从1000变为6,满足我自己的训练集。

4、模型训练

4.1 损失函数

本项目损失函数采用交叉熵损失函数

criterion = nn.CrossEntropyLoss()

4.2 优化器

使用Adam优化器,同时采用学习率调度器对在训练过程中对优化器学习率进行调整。

optimizer = torch.optim.Adam(model.parameters(), lr=lr) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)
​

4.3 训练过程可视化

在训练过程中采用tensorboarder进行数据可视化操作,对模型训练过程中训练精度、损失效果进行展示。

5、模型验证

5.1 验证过程数据化

5.2 指标报表

5.2 混淆矩阵

6、模型移植

6.1 导出onnx

import torch
import torch.nn as nn
from torchvision.models import densenet169
​
​
def onnx_export():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    # 加载模型
    model = densenet169()
​
    # 根据数据调整模型
    # model.fc.in_features返回的是最后一层全连接层的输入维度
    in_features = model.classifier.in_features
    model.classifier = nn.Linear(in_features, 6)
    model.conv0 = nn.Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
    model.maxpool = nn.MaxPool2d(kernel_size=1, stride=1)
​
    model.to(device)
    # 加载权重文件
    model.load_state_dict(torch.load("./desnet_bird.pth"))
​
    # 创建一个实例输入
    x = torch.randn(1, 3, 224, 224, device=device)
    # 导出onnx
    onnxpath = "desnet_bird.onnx"
    torch.onnx.export(
        model,
        x,
        onnxpath,
        #
        verbose=True,  # 输出转换过程
        input_names=["input"],
        output_names=["output"],
    )
    print("onnx导出成功")
​
​
if __name__ == '__main__':
    onnx_export()

6.2 onnx推理

import numpy as np
import torch
from PIL import Image
from torchvision.transforms import transforms
import onnxruntime as ort
​
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def inference():
    # 加载数据
    transform = transforms.Compose([
        transforms.Resize((226, 226)),
        transforms.CenterCrop(224),
        transforms.ToTensor(),
    ])
​
    img_path = "test_img/啄木鸟1.webp"
    img = Image.open(img_path).convert("RGB")
    img = transform(img)
    img = img.unsqueeze(0)
    # print(img.shape)
    # 将图片转化为ONNx运行时所需要的格式
    img = img.numpy()
​
    # 加载模型
    onnx_path = "desnet_bird.onnx"
    # 设置ONNx使用GPU
    provides = ["CUDAExecutionProvider"]
    # 加载ONNX模型
    sess = ort.InferenceSession(onnx_path,provider=provides)
​
    # 运行onnx模型
    outputs = sess.run(None, {"input": img})
    output = outputs[0]
    class_name = ["美洲金翅雀", "谷仓猫头鹰", "洋红蜂虎", "绒啄木鸟", "皇帝企鹅", "火烈鸟"]
    print(class_name[np.argmax(output)])
​
​
if __name__ == "__main__":
    inference()

7、项目总结

7.1 遇到的问题和解决办法

在项目数据集的划分时,由于我在kaggle上下载的数据集并没有划分为测试集和训练集,和直接使用ImageFolder对划分好的数据集使用有所不同,我想对数据集划分且不使用ImageFolder,我询问了AI,有了AI的帮助我的问题很快就迎刃而解了。当我的数据集、模型构建好后,我开始了第一轮训练,慢慢地我发现,训练集的精度非常的高,但是测试集的精度却很低,模型过拟合了,我以为是训练轮次的问题,于是我开始了第二轮训练,结果还是不行。于是我增加了数据增强的多样性以及对学习率进行了调整,然后模型过拟合问题就完美的解决了。

7.2 收获

通过这个项目我收获颇多,我深入的了解了如何使用PyTorch构建、训练和评估深度学习模型,还掌握了创建自定义数据集类(Dataset)、使用DataLoader进行批量处理等技巧。同时也学会了计算和解读各种评估指标,如准确率、精确度、召回率和F1得分。尽管在项目过程中遇到了一些问题,例如权重加载错误、数据格式不匹配等,但是经过这个项目,不仅提升了我的调试和解决问题的能力,还使我收益颇多。

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐