项目背景

利用ResNet50实现在beans数据集上面的微调,其主要有有如下亮点:

  1. 根据数据集特点,自写datasets来加载数据集
  2. 修改ResNet结构,使其适合数据集三分类输出
  3. 加载模型参数
  4. 冻结网络层
  5. 微调和验证,保存模型数据

数据集:beans

https://huggingface.co/datasets/AI-Lab-Makerere/beans
目录结构:
解压缩train.zip和validation.zip之后,每个文件夹下面有三个子文件夹,文件名表示三个类别,每个类别下面保存若干jpg图像数据。
在这里插入图片描述
注意:了解数据集文件特征,是自写数据集加载的必要条件,也就是重新__getitem__() 和__len__() 这两个方法,详细写法在下文给出。

预训练模型:

下载地址:wget https://download.pytorch.org/models/resnet50-0676ba61.pth
注意:也可以不下载预训练模型,直接从零开始训练模型

源码

import torch
from torchvision import models
from torch import nn

from torch.utils.data import Dataset
from torch.utils.data import DataLoader

import glob
import os

import torchvision.transforms as transforms
from PIL import Image
from tqdm import tqdm
torch.backends.cudnn.enabled = False

resnet50 = models.resnet50(pretrained=False) #pretrained=True
model_path = "./resnet50-0676ba61.pth"
pth = torch.load(model_path)
# print(pth.keys())

def load_state_dict_by_name(model_struct, model_state_dict, strict = True):
    """
    根据名称加载模型权重。

    :param model_struct: PyTorch 模型(如 ResNet-50)
    :param model_state_dict: 权重字典,通常是通过 torch.load() 加载的模型权重
    :param strict: 是否严格匹配所有的层,默认 True
    """
    model_sd = model_struct.state_dict()
    
    match_state_dict = {}
    for name,param in zip(model_state_dict.keys(),model_state_dict.values() ):
        if name in model_sd:# 如果模型中存在相同名称的层,则加载权重
            if param.shape == model_sd[name].shape:# 比较形状是否匹配
                match_state_dict[name] = param
                # print(f"orinig weight = {model_sd[name].mean()}, load weight ={param.mean()},shape = {param.shape}")
            else:
                print(f"Warning: Skipping layer {name},due to shape mismatch.Expected {param.shape}, but got {model_sd[name].shape}")
        else:
            print(f"Warning: Skipping layer {name} (not found in model)")
    
    model_sd.update(match_state_dict)  # 加载匹配的权重到模型
    if strict:
        missing_keys = set(model_sd.keys()) - set(match_state_dict.keys())
        if missing_keys:
            print(f"Warning: The following {len(missing_keys)} layers were not loaded! ")
    
    model_struct.load_state_dict(model_sd)# 更新模型的状态字典
    print(f"Successfully loaded {len(match_state_dict)} layers. missed load {len(missing_keys)}")
    return model_struct

resnet50 = load_state_dict_by_name(resnet50, pth, strict = True) #可用 resnet50.load_state_dict(pth) 替换


class_name_list = ["healthy","angular_leaf_spot","bean_rust"]
class_name_dict = {"healthy":0,"angular_leaf_spot":1,"bean_rust":2}

out_feature = len(class_name_list)

# 冻结模型参数,也就是采用模型已经训练好的原始参数
for name,param in resnet50.named_parameters():
#     # print(name, param.shape)
     param.requires_grad = False

resnet50.fc = nn.Sequential(nn.Linear(2048,len(class_name_list)),
                            nn.LogSoftmax(dim=1))
# print("model = ", resnet50)

# 加载数据集,自己重写DataSet类
class MyDataset(Dataset):
    # image_dir为数据目录,label_file,为标签文件
    def __init__(self, image_dir,transform=None):
        super(MyDataset, self).__init__()    # 添加对父类的初始化
        self.image_dir = image_dir         # 图像文件所在路径
        self.transform = transform
        
        self.img_file_list = []
        img_extensions = ['*.jpg','*.jpeg','*.png']
        for ext in img_extensions:
            self.img_file_list.extend(glob.glob(os.path.join(self.image_dir,"**", ext)))
    
    # 加载每一项数据
    def __getitem__(self, idx):
        image_index = self.img_file_list[idx]    # 图片相对路径
       
        label_name = image_index.split('/')[-2]     # 图片对应的父文件名称,也是图片类别
        label = class_name_dict[label_name]        # 把标签转换为类别0,1,2
        
        image = Image.open(image_index)
        transformed_image = self.transform(image)
        return transformed_image, label,label_name
    
    # 数据集大小
    def __len__(self):
        return (len(self.img_file_list))

transform_pre_preocess = transforms.Compose([
    transforms.Resize((224, 224)),                # 调整图像大小为 224x224
    transforms.ToTensor(),                        # 转换为 Tensor
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # 归一化
])

batch_size = 128
img_dir_path = "./train"
datasets_train = MyDataset(img_dir_path, transform=transform_pre_preocess)
dataloader_train = DataLoader(datasets_train,batch_size=batch_size,shuffle=True)
print("num of datasets_train = " , len(datasets_train))

val_img_dir_path = "./validation"
datasets_val = MyDataset(val_img_dir_path, transform=transform_pre_preocess)
dataloader_val = DataLoader(datasets_val,batch_size=batch_size,shuffle=True)
print("num of datasets_val = " , len(datasets_val))

if torch.cuda.is_available():
    resnet50 = resnet50.cuda()

loss_fn = nn.CrossEntropyLoss()
#optimizer
learning_rate = 1e-4
optimizer = torch.optim.SGD(resnet50.parameters(),lr=learning_rate,)

def metric(pred, lable):
    equal_elements = torch.eq(pred, lable)
    equal_count = torch.sum(equal_elements)
    return equal_count


epoch_num = 100
step_iter = 0
loss_train_max = 1e10
for e in tqdm(range(epoch_num)):
    # print(f"-------第{e} 个epoch训练开始----------")
    resnet50.train()
    
    for data in dataloader_train:
        imgs, lable, lable_name = data
        
        if torch.cuda.is_available():
            imgs = imgs.cuda()
            lable = lable.cuda()
            loss_fn = loss_fn.cuda()
        
        pred = resnet50(imgs)
        loss = loss_fn(pred, lable)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()
        step_iter += 1
        if loss_train_max > loss:
            loss_train_max = loss
            # 训练过程中保存模型的 checkpoint
            checkpoint = {
                'epoch': e,  # 当前训练的 epoch
                'model_state_dict': resnet50.state_dict(),  # 模型的权重
                'optimizer_state_dict': optimizer.state_dict(),  # 优化器的状态
                'loss': loss  # 当前的损失值(可选)
            }
            # 保存模型和优化器的状态
            torch.save(checkpoint, f'./out_resnet50/model_checkpoint_{e}.pth')
            print(f"epoch = {e}, save model_weight to the path: ./out_resnet50/model_checkpoint_{e}.pth, loss = {loss}")
            
        if (step_iter % (len(datasets_train)/batch_size)) ==0:
            print(f"\n 训练epoch={e},迭代次数={step_iter}, loss = {loss},acc = {metric(torch.argmax(pred, dim=1), lable) / batch_size}")
    
    
    if (e+1) % 20 == 0:
        print(f"-------第{e} 个epoch评估开始----------")
        resnet50.eval()
        total_true_count = 0
        for val_data in tqdm(dataloader_val):
            img_val, lable, lable_name = val_data
            if torch.cuda.is_available():
                img_val = img_val.cuda()
                lable = lable.cuda()
                loss_fn = loss_fn.cuda()
            pred_val = resnet50(img_val)
            max_indices_val = torch.argmax(pred_val, dim=1)
            # print(max_indices_val, lable)
            loss_val = loss_fn(pred_val,lable )
            true_count = metric(max_indices_val, lable)
            total_true_count += true_count
        print(f"第{e}个epoch 上面验证集:loss ={loss_val},Acc = {total_true_count/len(datasets_val)} ")   #Acc = {total_true_count/len(datasets_val)}  

Logo

更多推荐