一、引入包

import random
import torch
import torch.nn as nn
import numpy as np
import os
from PIL import Image #读取图片数据
from torch.utils.data import Dataset, DataLoader
from tqdm import tqdm
from torchvision import transforms
import time
import matplotlib.pyplot as plt
import matplotlib
matplotlib.use('TkAgg')
from model_utils.model import initialize_model

二、设置随机种子,确保结果可复现

def seed_everything(seed):
    torch.manual_seed(seed)
    torch.cuda.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.benchmark = False
    torch.backends.cudnn.deterministic = True
    random.seed(seed)
    np.random.seed(seed)
    os.environ['PYTHONHASHSEED'] = str(seed)
#################################################################
seed_everything(0)#随机种子?
###############################################

三、数据模块

1、对输入图像进行“数据增强”

HW = 224
#用于训练深度学习模型时对输入图像进行“数据增强”(不断喂各种图片,让模型认识)的处理流程
train_transform = transforms.Compose(#把多个图像变换操作按顺寻“串起来”,作为一个整体的处理流程
    [
        transforms.ToPILImage(),#把输入的图片转为PIL图像格式,224, 224, 3模型  :3, 224, 224
        transforms.RandomResizedCrop(224),#在图片中随即裁剪出一个区域,然后放大到224*224
        transforms.RandomRotation(50),#将图片随机旋转一个角度
        transforms.ToTensor()#将图片转成Tensor,像素归一化,维度重排(听到维度移到最前面,符合卷积)
    ]
)
val_transform = transforms.Compose(
    [
        transforms.ToPILImage(),   #224, 224, 3模型  :3, 224, 224
        transforms.ToTensor()
    ]
)

2、主数据集

class food_Dataset(Dataset):
    def __init__(self, path, mode="train"):
        self.mode = mode#设置实例化属性mode
        if mode == "semi":
            self.X = self.read_file(path)#从指定路径 path 读取数据,并将结果赋值给当前实例的属性 X,以便后续在类的其他方法中使用。
        else:
            self.X, self.Y = self.read_file(path)#从指定路径path读取样本及其对应标签,并分别存储为当前实例的两个属性X和Y,供后续训练使用。
            self.Y = torch.LongTensor(self.Y)  #标签转为长整形

        if mode == "train":#根据模式选择增强策略
            self.transform = train_transform
        else:
            self.transform = val_transform

    def read_file(self, path):
        if self.mode == "semi":
            file_list = os.listdir(path)#列出目录下所有文件
            xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)#预分配内存(len*HW*Hw*3)uint8无符号整数类型
            # 列出文件夹下所有文件名字
            for j, img_name in enumerate(file_list):#j为序号,img_name为图片名字
                img_path = os.path.join(path, img_name)#组合成每个图片的完整路径
                img = Image.open(img_path)#Image.open()函数打开指定路径的图像
                img = img.resize((HW, HW))#将图片设为224*224
                xi[j, ...] = img#存入数组xi[j][...]
            print("读到了%d个数据" % len(xi))#打印行数(既图片数)
            return xi
        else:
            for i in tqdm(range(11)):
                file_dir = path + "/%02d" % i
                file_list = os.listdir(file_dir)

                xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)
                yi = np.zeros(len(file_list), dtype=np.uint8)

                # 列出文件夹下所有文件名字
                for j, img_name in enumerate(file_list):
                    img_path = os.path.join(file_dir, img_name)
                    img = Image.open(img_path)
                    img = img.resize((HW, HW))
                    xi[j, ...] = img
                    yi[j] = i

                if i == 0:
                    X = xi
                    Y = yi
                else:
                    X = np.concatenate((X, xi), axis=0)#样本数量增多0个224*224*3->18个224*224*3
                    Y = np.concatenate((Y, yi), axis=0)#将当前类别的标签数组 yi 追加到已有标签数组 Y 的末尾
            print("读到了%d个数据" % len(Y))
            return X, Y

    def __getitem__(self, item):
        if self.mode == "semi":
            return self.transform(self.X[item]), self.X[item]#如果是无标签模式返回(增强图像(这里的不裁剪什么的),原始图像),增强图像用于模型预测,原始图像用于后续构建伪标签数据集
        else:
            return self.transform(self.X[item]), self.Y[item]#有标签模式返回(增强图像(要裁剪什么的),标签)y[item]为第item张图像的编号

    def __len__(self):#返回数据集大小
        return len(self.X)

3、半监督伪标签数据集

class semiDataset(Dataset):#半监督数据集
    def __init__(self, no_label_loder, model, device, thres=0.99):
        x, y = self.get_label(no_label_loder, model, device, thres)#得到满足高置信度的原始数据和伪标签
        if x == []:#没有超过阈值的样本,标记数据集无效,判断是否可以跳过数据集的训练
            self.flag = False

        else:#有超过阈值的样本,标记有效
            self.flag = True
            self.X = np.array(x)#将x转为数组
            self.Y = torch.LongTensor(y)#将y转为长整型张量
            self.transform = train_transform#应用数据增强变换,为了之后将伪标签数据当样本对模型训练
    def get_label(self, no_label_loder, model, device, thres):
        model = model.to(device)
        pred_prob = []
        labels = []
        x = []
        y = []
        soft = nn.Softmax()#将模型输出转为概率分部
        with torch.no_grad():#半监督不更新参数(无符号数没有真正参与反向传播)
            for bat_x, _ in no_label_loder:#得到无标签数据
                bat_x = bat_x.to(device)
                pred = model(bat_x)#经过模型前向传播
                pred_soft = soft(pred)#转为概率
                pred_max, pred_value = pred_soft.max(1)#得到最大的概率值和对应的索引
                pred_prob.extend(pred_max.cpu().numpy().tolist())#将最大概率值转为数组再转为列表
                labels.extend(pred_value.cpu().numpy().tolist())#将最大概率值索引转为数组再转为列表

        for index, prob in enumerate(pred_prob):#将置信度符合条件的加到有标签样本中
            if prob > thres:
                x.append(no_label_loder.dataset[index][1])   #调用到原始的getitem,第一个是增强图像
                y.append(labels[index])
        return x, y

    def __getitem__(self, item):
        return self.transform(self.X[item]), self.Y[item]
    def __len__(self):
        return len(self.X)

4、伪标签数据加载器

def get_semi_loader(no_label_loder, model, device, thres):
    semiset = semiDataset(no_label_loder, model, device, thres)#创建semiDataset类的实例
    if semiset.flag == False:#没有样本置信度超过阈值 返回None
        return None
    else:#有样本置信度超过阈值,
        semi_loader = DataLoader(semiset, batch_size=16, shuffle=False)#封装数据集为加载器,每次取出16个样本,不打乱数据样本
        return semi_loader

四、模型

class myModel(nn.Module):
    def __init__(self, num_class):#接受参数num_class表示分类的类别数量
        super(myModel, self).__init__()#调用父类初始化方法的方式
        #3 *224 *224  -> 512*7*7 -> 拉直 -》全连接分类
        self.conv1 = nn.Conv2d(3, 64, 3, 1, 1)    # 卷积层64*224*224
        self.bn1 = nn.BatchNorm2d(64)#批量归一化层,输出64个
        self.relu = nn.ReLU()#激活层
        self.pool1 = nn.MaxPool2d(2)#最大池   #64*112*112


        self.layer1 = nn.Sequential(#第一组卷积块
            nn.Conv2d(64, 128, 3, 1, 1),    # 128*112*112
            nn.BatchNorm2d(128),
            nn.ReLU(),
            nn.MaxPool2d(2)   #128*56*56
        )
        self.layer2 = nn.Sequential(
            nn.Conv2d(128, 256, 3, 1, 1),
            nn.BatchNorm2d(256),
            nn.ReLU(),
            nn.MaxPool2d(2)   #256*28*28
        )
        self.layer3 = nn.Sequential(
            nn.Conv2d(256, 512, 3, 1, 1),
            nn.BatchNorm2d(512),
            nn.ReLU(),
            nn.MaxPool2d(2)   #512*14*14
        )

        self.pool2 = nn.MaxPool2d(2)    #512*7*7,最大池化
        self.fc1 = nn.Linear(25088, 1000)   #25088->1000 全连接层
        self.relu2 = nn.ReLU()
        self.fc2 = nn.Linear(1000, num_class)  #1000-11

    def forward(self, x):
        x = self.conv1(x)
        x = self.bn1(x)
        x = self.relu(x)
        x = self.pool1(x)
        x = self.layer1(x)
        x = self.layer2(x)
        x = self.layer3(x)
        x = self.pool2(x)
        x = x.view(x.size()[0], -1)#展平
        x = self.fc1(x)
        x = self.relu2(x)
        x = self.fc2(x)
        return x

五、训练

def train_val(model, train_loader, val_loader, no_label_loader, device, epochs, optimizer, loss, thres, save_path):
    model = model.to(device)
    semi_loader = None #无符号数据加载器
    plt_train_loss = []#训练损失
    plt_val_loss = []#测试损失

    plt_train_acc = []#训练准确度
    plt_val_acc = []#测试准确度

    max_acc = 0.0#最大准确度

    for epoch in range(epochs):
        train_loss = 0.0
        val_loss = 0.0
        train_acc = 0.0
        val_acc = 0.0
        semi_loss = 0.0
        semi_acc = 0.0

        start_time = time.time()

        model.train()
        for batch_x, batch_y in train_loader:
            x, target = batch_x.to(device), batch_y.to(device)#取出原始值和真实值
            pred = model(x)#得到预测值
            train_bat_loss = loss(pred, target)#计算误差
            train_bat_loss.backward()#反向回传
            optimizer.step()  # 更新参数 之后要梯度清零否则会累积梯度
            optimizer.zero_grad()
            train_loss += train_bat_loss.cpu().item()#统计误差
            train_acc += np.sum(np.argmax(pred.detach().cpu().numpy(), axis=1) == target.cpu().numpy())
            #每次一组数[0.3,0.5,0.2]为猪,狗,猫的可能性猜测概率,detach()去除梯度,避免影响,取argmax(,axis=1)取最大值第二列的数既1,对应就是狗,如果真实是狗,双等号返回True(1),用train_acc统计
        plt_train_loss.append(train_loss / train_loader.__len__())#计算平均损失并加到plt_train_loss后
        plt_train_acc.append(train_acc/train_loader.dataset.__len__()) #记录准确率,总的1处以数据个数

        if semi_loader!= None:#如果有伪标签数据,加入训练
            for batch_x, batch_y in semi_loader:
                x, target = batch_x.to(device), batch_y.to(device)
                pred = model(x)
                semi_bat_loss = loss(pred, target)
                semi_bat_loss.backward()
                optimizer.step()  # 更新参数 之后要梯度清零否则会累积梯度
                optimizer.zero_grad()
                semi_loss += train_bat_loss.cpu().item()
                semi_acc += np.sum(np.argmax(pred.detach().cpu().numpy(), axis=1) == target.cpu().numpy())
            print("半监督数据集的训练准确率为", semi_acc/train_loader.dataset.__len__())


        model.eval()
        with torch.no_grad():
            for batch_x, batch_y in val_loader:
                x, target = batch_x.to(device), batch_y.to(device)
                pred = model(x)
                val_bat_loss = loss(pred, target)
                val_loss += val_bat_loss.cpu().item()
                val_acc += np.sum(np.argmax(pred.detach().cpu().numpy(), axis=1) == target.cpu().numpy())
        plt_val_loss.append(val_loss / val_loader.dataset.__len__())
        plt_val_acc.append(val_acc / val_loader.dataset.__len__())

        if epoch%3 == 0 and plt_val_acc[-1] > 0.6:#每训练三轮模型,进行一次判断无标签数据的准确率是否大于0.6,若大于则加入到标签数据
            semi_loader = get_semi_loader(no_label_loader, model, device, thres)

        if val_acc > max_acc:#记录最好模型
            torch.save(model, save_path)
            max_acc = val_loss

        print('[%03d/%03d] %2.2f sec(s) TrainLoss : %.6f | valLoss: %.6f Trainacc : %.6f | valacc: %.6f' % \
              (epoch, epochs, time.time() - start_time, plt_train_loss[-1], plt_val_loss[-1], plt_train_acc[-1], plt_val_acc[-1])
              )  # 打印训练结果。 注意python语法, %2.2f 表示小数位为2的浮点数, 后面可以对应。

    plt.plot(plt_train_loss)
    plt.plot(plt_val_loss)
    plt.title("loss")
    plt.legend(["train", "val"])
    plt.show()


    plt.plot(plt_train_acc)
    plt.plot(plt_val_acc)
    plt.title("acc")
    plt.legend(["train", "val"])
    plt.show()

六、其他

train_path = r"D:\pytest\clarify\food_classification\food-11_sample\training\labeled"
val_path = r"D:\pytest\clarify\food_classification\food-11_sample\validation"
no_label_path = r"D:\pytest\clarify\food_classification\food-11_sample\training\unlabeled\00"
#数据集
train_set = food_Dataset(train_path, "train")
val_set = food_Dataset(val_path, "val")
no_label_set = food_Dataset(no_label_path, "semi")
#定义加载器
train_loader = DataLoader(train_set, batch_size=16, shuffle=True)
val_loader = DataLoader(val_set, batch_size=16, shuffle=True)
no_label_loader = DataLoader(no_label_set, batch_size=16, shuffle=False)

# model = myModel(11)
model, _ = initialize_model("vgg", 11, use_pretrained=True)#加载一个训练好的模型


lr = 0.001
loss = nn.CrossEntropyLoss()#交叉熵损失函数
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
#自适应优化算法AdamW优化器
#Adam是自适应学习率优化算法,各参数有独立的学习率(梯度大的参数步长小,梯度小的步长大)且自动更新,引入动量(利用历史梯度方向,加速收敛),当使用L2正则化,Adam使之效果不稳定
  #这个正则项和梯度,然后被 Adam 的自适应机制(除以 √v̂ₜ)缩放。导致梯度大的参数,其权重衰减也被削弱了 → 正则化效果不一致!
  #weight_decay=λ(权重衰减系数)  L′=L+0.5λθ² 梯度▽L′=▽L+λθ 动量m₂=β₁m₁+(1-β₁)▽L′ 自适应项v₂=β₁v₁+(1-β₁)(▽L′)² θ₃=θ₂-ŋm₂/(√v₂+ε)当θ大->v大->历史梯度大的,二阶矩(自适应项)大导致有效学习率小,
#AdamW是将梯度只来源于损失▽L,无正则项,独立权重衰减θ₃=θ₂-ŋm₂/(√v₂+ε)-ŋλθ 衰减不受历史梯度影响
#在 Adam 中,由于权重衰减项λθ被混入梯度并参与二阶矩v的计算,导致大权重参数的有效更新步长被自适应机制过度压缩,使得其实际受到的正则化强度反而减弱,违背了 L2 正则化的本意。
#而 AdamW 通过解耦权重衰减,确保正则化效果不被自适应学习率干扰,从而获得更稳定、更强的泛化能力。
device = "cuda" if torch.cuda.is_available() else "cpu"
save_path = "model_save/best_model.pth"
epochs = 15
thres = 0.99

train_val(model, train_loader, val_loader, no_label_loader, device, epochs, optimizer, loss, thres, save_path)

更多推荐