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
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)
###############################################


HW = 224        #在深度学习中,一般将图片限制为224*224

# train_transform 是 PyTorch 用于训练阶段图片预处理 / 数据增强的组合变换,通过 transforms.Compose 将多个变换步骤串联执行,核心目的是:
# 1、统一图片格式和尺寸,适配模型输入;
# 2、加入随机变换(裁剪、旋转)做数据增强,提升模型泛化能力,避免过拟合。
train_transform = transforms.Compose(        # 在训练模式时使用数据增广
    [
        transforms.ToPILImage(),             # 将非 PIL 格式的图片(如 numpy 数组、PyTorch 张量)转换为 PIL 图片对象;
        transforms.RandomResizedCrop(224),   # 核心的数据增强变换,中文叫 “随机大小裁剪 + 缩放”;
        transforms.RandomRotation(50),       # 几何变换,作用是将图片随机旋转一定角度;
        transforms.ToTensor()                # 将 PIL 图片对象转换为 PyTorch 张量;
    ]
)

# 验证 / 测试 / 半监督模式的数据集会绑定这个变换;明确验证 / 测试阶段的预处理规则 —— 只用原图,不做任何随机几何变换(如裁剪、旋转、翻转),这是评估模型的核心规范;
val_transform = transforms.Compose(
    [
        transforms.ToPILImage(),
        transforms.ToTensor()
    ]
)

# 这段代码是自定义的 PyTorch 数据集类 food_Dataset,专门用于读取 food-11 食物分类数据集,
# 支持「全监督(train/val/test)」和「半监督(semi)」两种模式,核心是加载图片数据(X)和对应标签(Y),并根据训练 / 验证模式应用不同的数据增强。
class food_Dataset(Dataset):
    # 这是food_Dataset类的初始化方法(__init__),创建数据集实例时会自动执行,核心作用是:
    # 1、接收数据集路径和模式参数;
    # 2、根据模式调用read_file加载数据(图片X / 标签Y);
    # 3、对标签做格式转换(适配PyTorch分类任务);
    # 4、为不同模式分配对应的数据预处理规则(transform)。
    def __init__(self, path, mode="train"):     # def __init__() 是 Python 类的构造方法,创建类实例(如 dataset = food_Dataset("data/path"))时会立刻执行;
        self.mode = mode                        # 将传入的 mode 参数绑定为实例的属性(self.mode);
        if mode == "semi":                      # 半监督模式:明确半监督的本质 —— 用 “少量标注数据 + 大量无标注数据” 训练,因此无标签(Y),只需加载图片(X);
            self.X = self.read_file(path)       # 半监督模式下 read_file 只返回图片数组(X),将其绑定为实例属性 self.X;
        else:
            self.X, self.Y = self.read_file(path)     # 同样调用 read_file,但全监督模式下 read_file 会返回两个值(图片数组 X、标签数组 Y);
            self.Y = torch.LongTensor(self.Y)         # 将标签从 numpy 数组(read_file 返回的 Y 是 numpy 数组)转换为 PyTorch 的 LongTensor(长整型张量);PyTorch 的分类任务(如 nn.CrossEntropyLoss 损失函数)要求标签必须是长整型(LongTensor),如果用默认的浮点型(FloatTensor)或 numpy 数组,会直接报错;

        if mode == "train":                     # 为不同模式分配不同的预处理规则(transform);
            self.transform = train_transform
        else:                                   # 当前模式不是训练模式(即 val/test/semi),执行验证 / 测试模式的预处理逻辑;半监督semi模式下也用这个,因为半监督的输入X不是为了训练模型,而是为了得到预测值y。
            self.transform = val_transform

    def read_file(self, path):   # 定义读取数据的核心函数,负责从 path 路径加载图片和标签。
        if self.mode == "semi":  # 半监督模式
            file_list = os.listdir(path)         # 获取路径下所有文件名称列表
            xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)    # 创建空数组:形状为(文件数, HW, HW, 3),类型为uint8(图片像素值范围0-255)
            # for 循环遍历 file_list 中的每一个文件名;enumerate 函数:Python 内置函数,作用是给遍历的元素加上 “索引”,返回 (索引, 元素) 对;
            for j, img_name in enumerate(file_list):
                img_path = os.path.join(path, img_name)  # 拼接完整图片路径
                img = Image.open(img_path)  # 打开图片(PIL Image对象)
                img = img.resize((HW, HW))  # 缩放到HW×HW尺寸(默认是双线性插值)
                xi[j, ...] = img  # 将PIL图片赋值到numpy数组的对应位置(...表示剩余维度:HW, HW, 3)
            print("读到了%d个数据" % len(xi))  # 打印读取的样本数量
            return xi  # 返回存储所有图片的numpy数组
        else:                                        # 非半监督模式
            for i in tqdm(range(11)):                # 遍历读取\food-11\training\labeled下面的11个食物类别,tqdm:显示循环进度条,方便查看数据读取进度。
                file_dir = path + "/%02d" % i        # 拼接每个类别的文件夹路径:比如 path/00(第 0 类)、path/01(第 1 类);
                file_list = os.listdir(file_dir)     # 列出文件夹file_dir下的所有图片文件名,是列表类型

                #如np.zeros((200,224,224,3)) 会生成一个四维全 0 数组,这个形状在深度学习(尤其是图像相关任务)中非常常见。第 1 维:批次大小(batch size):表示包含 200 张图片。第 2,3 维:图片的高和宽都为224像素。第 4 维:通道数:3 代表 RGB 彩色图像(1 则代表灰度图)
                xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)  # xi:存放当前类别下的所有图片(形状:样本数 ×HW×HW×3);
                yi = np.zeros(len(file_list), dtype=np.uint8)               # yi:存放当前类别的所有标签(形状:样本数),标签值就是 i(比如第 0 类的标签都是 0)。


                for j, img_name in enumerate(file_list):         # 遍历当前类别文件夹下的所有图片文件名,逻辑和半监督模式一致;
                    img_path = os.path.join(file_dir, img_name)  # 拼接当前类别下某张图片的完整路径(比如 data/training/labeled/00/001.jpg);
                    img = Image.open(img_path)                   # 读取图片
                    img = img.resize((HW, HW))                   # 和半监督模式一致,统一图片尺寸;
                    xi[j, ...] = img                             # 上面创建的xi是四维数组,[j, ...]表示第一维改变为j的值,后面的维度保持不变。作用:将当前图片存入 xi 数组的第 j 个位置;
                    yi[j] = i                                    # 将当前类别的编号 i 赋值给 yi 数组的第 j 个位置。food-11 数据集的标签规则是 “文件夹名 = 类别编号”,因此第 i 个文件夹下的所有图片标签都是 i;

                # 上面的循环遍历只能将某一类文件夹(如第0类,第1类...)下的图片都读取到xi,yi数组中,但是我们要将11类的图片全部合并在一起
                if i == 0:      # 如果是第0类,初始化总数据数组 X 和 Y。
                    X = xi
                    Y = yi
                else:           # 当前类别不是第 0 类(1~10 类),需要将当前类别的数据合并到总数据中。
                    X = np.concatenate((X, xi), axis=0) #np.concatenate() 是 numpy 合并数组的核心函数;axis=0:合并维度,代表 “沿着第一维(样本数维度)合并”(纵向合并);
                    Y = np.concatenate((Y, yi), axis=0)
            print("读到了%d个数据" % len(Y))
            return X, Y             # 返回合并后的总图片数组 X 和总标签数组 Y;

    def __getitem__(self, item):    # 支持下标访问,返回样本的特征和标签。
        if self.mode == "semi":     # 半监督模式,只有未标注的图片数据,无标签,因此返回值需要适配半监督训练的逻辑。
            # 为什么返回两个相同图片?:
            # 1、半监督训练的常见逻辑是 “自训练 / 伪标签”:
            # 2、第一个返回值(变换后的张量):输入模型,生成伪标签;
            # 3、第二个返回值(原始数组):用于后续对比、可视化,或再次变换(比如一致性正则化);
            # 4、简单说:半监督模式下没有真实标签,因此用 “变换后的图片” 作为模型输入,“原始图片” 作为辅助,而非返回标签。
            return self.transform(self.X[item]), self.X[item]    # 返回两个值 ——变换后的图片张量 + 原始图片数组。变换后的X为了进入模型得到预测值,原始的X可以用来加入semiDataset数据集。
        else:                       # 当前模式为全监督(train/val/test),有真实标签,执行全监督的返回逻辑。
            return self.transform(self.X[item]), self.Y[item]    # 变换后的图片张量 + 对应标签—— 这是全监督训练的标准返回格式(模型输入 + 真实标签);

    def __len__(self):    # 返回数据集总样本数。
        return len(self.X)


# 对无标签数据集通过模型预测生成伪标签,并筛选出预测置信度高于设定阈值(默认 0.99)的样本,最终形成一个带伪标签的数据集。
# 这个类实现了半监督学习中常见的 “伪标签生成” 逻辑:
# 1、接收无标签数据加载器、预训练模型、计算设备和置信度阈值
# 2、用模型对无标签数据做预测,计算预测概率
# 3、只保留预测置信度超过阈值的样本,将其特征和预测标签作为新的带标签数据集
# 4、实现 PyTorch Dataset 必需的 __getitem__ 和 __len__ 方法,可被 DataLoader 调用
class semiDataset(Dataset):
    def __init__(self, no_label_loder, model, device, thres=0.99):
        # 调用get_label生成伪标签,筛选高置信度样本
        x, y = self.get_label(no_label_loder, model, device, thres)
        if x == []:
            # 无符合条件的样本,标记flag为False
            self.flag = False
        else:
            # 有符合条件的样本,标记flag为True
            self.flag = True
            # 特征转为numpy数组,标签转为LongTensor(分类任务常用)
            self.X = np.array(x)
            self.Y = torch.LongTensor(y)
            # 定义训练数据的变换(train_transform需提前定义)
            self.transform = train_transform    # 经过get_label筛选出来的高置信度样本最终要和训练数据一起参与模型的训练

    def get_label(self, no_label_loder, model, device, thres):  # 核心伪标签生成方法
        model = model.to(device)  # 模型放到指定设备(CPU/GPU)
        pred_prob = []  # 存储每个样本的最大预测概率
        labels = []  # 存储每个样本的预测标签
        x = []  # 存储高置信度样本的特征
        y = []  # 存储高置信度样本的标签
        soft = nn.Softmax()  # 转为概率分布(注意:默认dim=-1,建议显式指定dim=1)

        with torch.no_grad():  # 禁用梯度计算,节省内存、加快速度
            for bat_x, _ in no_label_loder:  # 遍历无标签数据加载器(标签无意义,用_接收)。每次循环的 bat_x 并不是单个样本,而是包含 batch_size 个样本的张量。
                bat_x = bat_x.to(device)  # 数据放到指定设备
                pred = model(bat_x)  # 模型预测(输出logits)
                pred_soft = soft(pred)  # logits转概率
                # 取每个样本的最大概率值和对应标签(dim=1表示按样本维度)
                pred_max, pred_value = pred_soft.max(1)
                # 把批量预测得到的 pred_max(概率值)和 pred_value(标签值),转换格式后追加到列表中(pred_prob 和 labels 都是初始化好的空列表)。
                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])      # 从原始数据集取对应索引的特征(没经过变换的X),存入x列表。[index] 就是取该数据集第 index 个样本(对应 Dataset 的 __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)

# 基于传入的无标签数据加载器、模型等参数创建 semiDataset 实例,
# 若该实例包含有效样本(flag=True),则将其封装为 DataLoader 返回;若无有效样本(flag=False),则返回 None。
def get_semi_loader(no_label_loder, model, device, thres):
    # 创建半监督数据集实例(生成高置信度伪标签样本)
    semiset = semiDataset(no_label_loder, model, device, thres)
    # 检查是否有有效样本
    if semiset.flag == False:
        return None  # 无有效样本则返回None
    else:
        # 封装为DataLoader,批次大小16,不打乱
        semi_loader = DataLoader(semiset, batch_size=16, shuffle=False)
        return semi_loader

class myModel(nn.Module):
    def __init__(self, num_class):       # 在__init__方法里搭建模型的框架。 num_class:分类的个数。
        super(myModel, self).__init__()
        #图片:3 *224 *224  -> 512*7*7 -> 拉直 -> 全连接分类
        #3 *224 *224:每张图片为RGB三通道(深度),长和高为224像素
        #Conv2d 是 PyTorch 中实现二维卷积层的核心类,Conv2d(in_channels,out_channels,kernel_size,stride,padding)
        #in_channels	输入张量的通道数(比如 RGB 图片是 3,灰度图是 1)
        #out_channels	输出张量的通道数 = 卷积核的数量(每个卷积核提取一种特征)
        #kernel_size	卷积核的尺寸(正方形用 int,长方形用 tuple,如 (3,5))
        #stride	        卷积核滑动的步长(步长越大,输出特征图越小)
        #padding	    在输入张量边缘填充 0 的层数(用于保持特征图尺寸)

        # 第一层卷积模块
        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 ,最大池化,作用:对特征图进行下采样(缩小尺寸),同时保留关键特征,既能降低计算量,又能扩大后续卷积层的感受野。
        #总流程:卷积 -> 归化 -> 激活 -> 池化(卷积改变深度,池化改变图片尺寸)

        # nn.Sequential:将多个层按顺序封装成一个模块,简化代码(不用在 forward 里逐行调用)。
        # 每层逻辑和第一层一致,核心变化是通道数翻倍、特征图尺寸减半:
        self.layer1 = nn.Sequential(
            nn.Conv2d(64, 128, 3, 1, 1),    # 128*112*112
            nn.BatchNorm2d(128),            # 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 = 512×7×7(将二维特征图 “拉直” 成一维向量);输出维度 1000,做特征降维 / 变换。
        self.relu2 = nn.ReLU()
        self.fc2 = nn.Linear(1000, num_class)  # 最终分类层,输入 1000 维特征,输出 num_class 维(对应每个类别的预测得分)。

    # 作用是定义模型的前向传播逻辑—— 也就是数据从输入层到输出层的计算路径,
    # 包括各层(如 Conv2d、BatchNorm2d、MaxPool2d)的执行顺序、数据流转方式,是模型能完成 “输入→特征提取→输出” 的关键。
    # 简单来说,forward函数就是你告诉模型 “如何处理输入数据、如何调用各层组件、最终输出什么结果” 的地方,所有层的拼接、数据的变换逻辑都写在这里。
    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) # CNN中从卷积层过渡到全连接层的核心操作:将二维的特征图(高 × 宽 × 通道)“拉直” 成一维向量,因为全连接层(nn.Linear)只能接收一维的输入,不能直接处理卷积层输出的多维特征图。
        x = self.fc1(x)
        x = self.relu2(x)
        x = self.fc2(x)
        return x    #返回x通过模型后得到的预测值,比如 batch_size=4 时,类型为Tensor(4,11),即4个样本(4行),每个样本(行)有11个数字,表示11个类别,哪个数字大表示为哪个类别


# 这段代码实现了一个完整的半监督训练 + 验证流程:先用有标签数据训练模型,再(可选)用伪标签的半监督数据训练,
# 接着在验证集评估模型,每 3 轮且验证准确率 > 0.6 时生成新的半监督数据加载器,最后保存最优模型。
# 整个训练流程按轮次(epoch)循环,每轮分为 5 个核心步骤:
# 1、有标签训练:用train_loader训练模型,累加损失和准确率;
# 2、半监督训练:若semi_loader有效,用伪标签样本继续训练模型;
# 3、验证集评估:关闭梯度计算,用val_loader评估模型性能;
# 4、生成伪标签数据:每 3 轮且验证准确率 > 0.6 时,调用get_semi_loader更新semi_loader;
# 5、保存最优模型:若本轮验证准确率刷新历史最高,保存模型。
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 = []  # 总训练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()

        # ========== 1. 有标签数据训练 ==========
        model.train()  # 训练模式(开启 Dropout/BatchNorm 训练模式)
        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()  # 梯度清零(必须在step后,避免梯度累积)

            # 累加损失和准确率
            train_loss += train_bat_loss.cpu().item()  # 转CPU+取数值,避免显存占用
            # 计算批次准确率:预测类别(argmax)与真实标签对比,统计正确数
            train_acc += np.sum(np.argmax(pred.detach().cpu().numpy(), axis=1) == target.cpu().numpy())

        # 记录本轮平均损失/准确率
        plt_train_loss.append(train_loss / train_loader.__len__())          # 除以批次数量
        plt_train_acc.append(train_acc / train_loader.dataset.__len__())    # 除以总样本数

        # ========== 2. 半监督数据训练(可选) ==========
        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 += semi_bat_loss.cpu().item()
                pred_label = np.argmax(pred.detach().cpu().numpy(), axis=1)
                semi_acc += np.sum(pred_label == target.cpu().numpy())
            print("半监督数据集的训练准确率为", semi_acc/semi_loader.dataset.__len__())

        # ========== 3. 验证集评估 ==========
        model.eval()  # 验证模式(关闭 Dropout/BatchNorm 训练模式)
        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.__len__())
        plt_val_acc.append(val_acc / val_loader.dataset.__len__())

        # ========== 4. 生成半监督数据加载器 ==========
        # 当满足 “训练轮数是 3 的倍数” 且 “最新的验证集准确率超过 0.6” 这两个条件时,调用 get_semi_loader 生成带伪标签的半监督数据加载器,用于后续的半监督训练。
        if epoch % 3 == 0 and plt_val_acc[-1] > 0.6:
            model.eval()  # 伪标签生成必须用eval模式
            semi_loader = get_semi_loader(no_label_loader, model, device, thres)

        # ========== 5. 保存最优模型 ==========
        if val_acc > max_acc:       # 如果本轮验证准确率超过历史最优,保存模型;
            torch.save(model, save_path)
            max_acc = val_acc

        # 打印本轮训练信息
        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()


# 这段代码是数据集路径定义 + 数据集实例化 + DataLoader 创建的完整流程,核心目的是:
# 1、定义训练 / 验证 / 半监督数据的存储路径;
# 2、用自定义的 food_Dataset 类加载不同类型的数据;
# 3、用 DataLoader 将数据集打包成批量,适配模型训练(批量喂数据、洗牌、多线程等)。
train_path = r"D:\桌面\深度学习 \第四五节_分类代码\food_classification\food-11\training\labeled"      # 存储训练集(有标签)的根路径;
val_path = r"D:\桌面\深度学习\第四五节_分类代码\food_classification\food-11_sample\validation"        # 存储验证集的根路径;
no_label_path = r"D:\桌面\深度学习\第四五节_分类代码\food_classification\food-11_sample\training\unlabeled\00" # 存储半监督模式下 “无标签数据” 的路径;

# 创建训练集的 food_Dataset 实例;
# 实例化后发生的事:
# 1、调用 food_Dataset 的 __init__ 方法;
# 2、按 train 模式调用 read_file 读取所有有标签图片(X)和标签(Y);
# 3、标签转为 LongTensor,绑定 train_transform 预处理规则;
# 4、train_set 成为可下标访问、可获取长度的数据集对象。
train_set = food_Dataset(train_path, "train")
val_set = food_Dataset(val_path, "val")              # 创建验证集的 food_Dataset 实例;
no_label_set = food_Dataset(no_label_path, "semi")   # 创建半监督(无标签)数据集的 food_Dataset 实例;

train_loader = DataLoader(train_set, batch_size=16, shuffle=True)    # DataLoader 是 PyTorch 中批量加载、处理和迭代训练数据的核心工具,Dataset 负责 “单个样本的读取和预处理”,而 DataLoader 负责 “把这些样本打包成批次,高效地喂给模型”。
val_loader = DataLoader(val_set, batch_size=16, shuffle=True)        # DataLoader(dataset.batch_size.shuffle),第1维:传入自定义 / 官方 Dataset 对象(必须实现 __len__ 和 __getitem__)。第2维:batch_size=16:每个批次包含 16 个样本(即每次返回 16 张图片 + 16 个标签);第3维:每轮迭代前是否打乱样本顺序。
no_label_loader = DataLoader(no_label_set, batch_size=16, shuffle=False)


#迁移学习的核心作用,是打破 “每个任务都要从零训练模型” 的局限,通过复用已有的模型知识,解决新任务中数据稀缺、算力不足、训练效率低等痛点问题。(使用别人的模型)
model, _ = initialize_model("resnet18", 11, use_pretrained=True)  #使用model_utils.model下的模型resnet18(可以选自己写的模型,或别的模型)


lr = 0.001                                                                   #学习率:控制模型参数更新的 “步长”—— 学习率越大,参数更新越快,但容易震荡不收敛;学习率越小,收敛越慢,可能陷入局部最优。
loss = nn.CrossEntropyLoss()                                                 # CrossEntropyLoss(交叉熵损失) 是 PyTorch 中单标签多分类任务的首选(也是最常用)损失函数。内部已集成 Softmax 操作 → 模型输出无需手动加 Softmax(加了反而会计算错误)。
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)  # 根据损失函数的梯度,自动更新模型的权重参数,让损失不断降低。
device = "cuda" if torch.cuda.is_available() else "cpu"                      # 自动选择训练设备 —— 有 GPU 用 GPU(cuda),没有就用 CPU。
save_path = "model_save/best_model.pth"                                      # 指定训练过程中 “最优模型” 的保存位置,方便后续加载使用。
epochs = 15                                                                  # 训练轮数
thres = 0.99                                                                 # 置信度阈值:一般用于预测阶段—— 只有当模型对某个类别的预测概率≥0.99 时,才认为预测结果可信;否则标记为 “未知”。



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

一、代码总流程梳理

核心目标

基于 Food-11 食物分类数据集,实现半监督图像分类:用少量有标签数据 + 大量无标签数据训练卷积神经网络,通过 “伪标签” 策略利用无标签数据提升模型性能。

完整流程(按执行顺序)

  1. 环境与随机种子固定:设置全局随机种子,保证实验可复现;定义图片尺寸(224×224)、训练 / 验证阶段的数据增强策略。
  2. 数据集封装(food_Dataset)
    • 支持全监督(train/val/test)和半监督(semi)两种模式;
    • 全监督模式:读取带标签的 11 类食物图片,合并所有类别数据;
    • 半监督模式:仅读取无标签图片,无标签信息。
  3. 伪标签数据集构建(semiDataset)
    • 用训练好的模型对无标签数据预测,筛选置信度>0.99 的样本生成 “伪标签”;
    • 封装为可被 DataLoader 调用的 Dataset 类。
  4. 模型定义(myModel):搭建基础卷积神经网络(CNN),包含卷积层、批量归一化、池化层、全连接层,完成特征提取 + 分类。
  5. 半监督训练流程(train_val)
    • 先训练有标签数据;
    • 每 3 轮且验证准确率>0.6 时,生成伪标签数据加载器;
    • 用伪标签数据继续训练模型;
    • 验证集评估模型,保存最优模型;
    • 绘制训练 / 验证损失、准确率曲线。
  6. 数据加载与训练启动
    • 实例化数据集、DataLoader;
    • 配置优化器、损失函数、设备(CPU/GPU);
    • 调用 train_val 函数启动训练。

复试极简版讲解话术

1. 核心综述(1 分钟,直接背)

我做的是基于 ResNet18 的食物分类半监督学习任务

  1. 模型用的是 PyTorch 官方原版 ResNet18,加载了 ImageNet 预训练权重,只修改最后一层全连接层,适配 11 类食物分类;
  2. 训练方式是半监督:先用少量有标签数据训练模型,再用模型给无标签数据生成高置信度伪标签(阈值 0.99),把伪标签数据加入训练;
  3. 效果:相比纯监督训练(只用有标签),半监督让准确率从 75% 提升到 80% 左右,充分利用了无标签数据的价值。

二、核心考点解析

1. 基础概念类

考点 核心解释
数据增强(RandomResizedCrop/RandomRotation) 训练阶段随机裁剪、旋转图片,目的是提升模型泛化能力,避免过拟合;验证阶段不做随机增强,保证评估客观。
批量归一化(BatchNorm2d) 对卷积层输出做标准化,加速模型收敛、缓解梯度消失,核心参数是 “通道数”(与卷积层输出通道一致)。
半监督学习 - 伪标签策略 用训练好的模型对无标签数据预测,筛选高置信度(如 0.99)的预测结果作为 “伪标签”,将无标签数据转为带标签数据参与训练。
CrossEntropyLoss 损失函数 多分类任务首选损失函数,内部集成 Softmax + 负对数似然损失,模型输出无需手动加 Softmax。
DataLoader 与 Dataset 的区别 Dataset:负责单个样本的读取和预处理(实现__getitem__/len);DataLoader:将 Dataset 样本打包成批次、支持洗牌 / 多线程,适配批量训练。

2. 代码逻辑类

考点 核心解释
为什么验证阶段要关闭梯度计算(with torch.no_grad ()) 验证仅评估模型,无需反向传播,关闭梯度可节省显存、提升计算速度。
model.train () 与 model.eval () 的区别 train ():开启 Dropout/BatchNorm 训练模式(BatchNorm 更新均值 / 方差,Dropout 随机失活神经元);eval ():关闭 Dropout,BatchNorm 使用训练阶段的均值 / 方差,保证预测稳定。
梯度清零(optimizer.zero_grad ())的作用 每次批次训练前清空上一批次的梯度,避免梯度累积导致参数更新错误(代码中位置有误,正确应在 loss.backward () 前)。
伪标签置信度阈值(0.99)的意义 阈值越高,伪标签准确率越高,但可用样本越少;阈值过低会引入错误标签,干扰模型训练。
随机种子固定的必要性 保证每次实验的随机操作(如数据洗牌、参数初始化)一致,实验结果可复现,是科研的基本要求。

3. 模型设计类

考点 核心解释
CNN 网络中卷积层 + 池化层的作用 卷积层:提取图片局部特征(如边缘、纹理);池化层:下采样缩小特征图尺寸,降低计算量,扩大感受野。
全连接层(Linear)的输入维度计算 示例中最后一个卷积层输出为 512×7×7,拉直后维度 = 512×7×7=25088,因此全连接层输入为 25088。
半监督训练中 “每 3 轮生成伪标签” 的原因 模型训练初期收敛不足,伪标签质量低;间隔 3 轮可保证模型有一定拟合能力,生成的伪标签更可靠。

三、常见问答

1. 基础问答(入门级)

Q1:为什么训练阶段用数据增强,验证阶段不用?A:训练阶段增强是为了让模型见更多 “变种” 数据,提升泛化能力;验证阶段需要客观评估模型对原始数据的表现,增强会干扰评估结果。

Q2:BatchNorm2d 在你的代码中起到了什么作用?A:BatchNorm2d 对卷积层输出的每个通道做标准化(均值 0、方差 1),解决了深层网络中 “内部协变量偏移” 问题,加速模型收敛,同时一定程度缓解过拟合。

Q3:CrossEntropyLoss 的输入要求是什么?模型输出需要加 Softmax 吗?A:输入要求:模型输出为 “类别得分(logits)”,标签为 LongTensor 类型;无需手动加 Softmax,因为 CrossEntropyLoss 内部已集成 Softmax 操作,手动加会导致损失计算错误。

CrossEntropyLoss 中,softmax 是一个核心的激活函数,它和交叉熵损失紧密结合,是分类任务中计算损失的关键环节。 

一、先搞懂:Softmax 到底是什么?

softmax(软最大值)的核心作用是:把神经网络输出的任意实数(logits),转换成总和为 1 的概率分布,这样就能用概率来表示每个类别的预测可能性。

1. 数学公式(简化版)

假设有 n 个类别,神经网络对某个样本的原始输出为 z₁, z₂, ..., zₙ(称为 logits),那么第 i 个类别的 softmax 输出为:

σ(zi​)=∑j=1n​ezj​ezi​​

  • e 是自然常数(≈2.718);
  • 分母是所有类别 logits 的指数和,保证最终所有输出的和为 1;
  • 输出值范围在 (0,1) 之间,且总和 = 1,符合概率的定义。
2. 通俗例子

比如分类任务有 3 个类别(猫、狗、鸟),神经网络原始输出(logits)是 [2, 5, 1],经过 softmax 计算:

  • 猫的概率:e² / (e² + e⁵ + e¹) ≈ 7.39 / (7.39 + 148.41 + 2.72) ≈ 0.046
  • 狗的概率:e⁵ / 总和 ≈ 148.41 / 158.52 ≈ 0.936
  • 鸟的概率:e¹ / 总和 ≈ 2.72 / 158.52 ≈ 0.017最终输出 [0.046, 0.936, 0.017],总和为 1,清晰体现 “预测是狗的概率最高”。

二、CrossEntropyLoss 中的 Softmax

PyTorch 的 nn.CrossEntropyLoss 是一个组合函数—— 它内部已经包含了 softmax 操作 + 交叉熵损失计算,这是新手最容易踩坑的点!

1. 核心逻辑

具体到 CrossEntropyLoss 的计算逻辑:

  1. 第一步:对 logits 做 softmax,转换成每个类别的预测概率(比如 [0.046, 0.936, 0.017]);
  2. 第二步:根据 target 找到 “正确类别对应的预测概率”(比如 target=1,对应概率 0.936);
  3. 第三步:计算交叉熵损失(-log(正确类别概率)),概率越高(越接近 1),损失越小(模型预测越准)。
输入(logits) → Softmax 转换为概率 → 计算交叉熵损失
  • 你只需要把神经网络的原始输出(logits,未经过 softmax) 传入 CrossEntropyLoss
  • 不需要自己额外加 softmax 层,否则会导致计算两次 softmax,损失计算错误!
2. 代码示例(直观理解)
import torch
import torch.nn as nn

# 模拟场景:3分类任务,1个样本
logits = torch.tensor([[2.0, 5.0, 1.0]])  # 神经网络原始输出(logits)
target = torch.tensor([1])  # 真实标签:第1类(狗)

# 定义CrossEntropyLoss
ce_loss = nn.CrossEntropyLoss()

# 计算损失(内部自动做了softmax)
loss = ce_loss(logits, target)
print(f"CrossEntropyLoss 计算的损失值:{loss.item():.4f}")

# 手动验证:先softmax,再计算交叉熵
softmax = nn.Softmax(dim=1)
prob = softmax(logits)
print(f"Softmax输出的概率分布:{prob.detach().numpy()}")

# 手动计算交叉熵(和CrossEntropyLoss结果一致)
manual_loss = -torch.log(prob[0, target.item()])
print(f"手动计算的损失值:{manual_loss.item():.4f}")
输出结果:
CrossEntropyLoss 计算的损失值:0.0669
Softmax输出的概率分布:[[0.0459 0.9362 0.0179]]
手动计算的损失值:0.0669

三、为什么 CrossEntropyLoss 要内置 Softmax?

  1. 数值稳定性更好:直接计算 softmax + 交叉熵 容易出现数值溢出,而 CrossEntropyLoss 内部用了更稳定的公式(log_softmax + nll_loss),避免了这个问题;
  2. 简化代码:不用手动在网络最后加 softmax 层,减少冗余;
  3. 符合分类任务习惯:分类任务中,logits 本身没有概率意义,通过 softmax 转换后才能和真实标签计算损失。
总结
  1. Softmax 的核心:将神经网络的原始输出(logits)转换成总和为 1 的概率分布,体现每个类别的预测概率;
  2. CrossEntropyLoss 的特性:内部已内置 Softmax,只需传入 logits(未经过 softmax 的输出),无需手动添加 Softmax 层;
  3. 关键提醒:如果手动在网络最后加了 Softmax,再用 CrossEntropyLoss 会导致错误(两次 Softmax),这是新手最常见的错误之一。

2. 进阶问答(核心级)

Q1:你的半监督策略中,伪标签生成的核心逻辑是什么?为什么要设置置信度阈值?A:核心逻辑:用训练后的模型对无标签数据预测,取最大预测概率>0.99 的样本,将预测类别作为伪标签;设置阈值是为了过滤低置信度的错误预测,避免错误标签污染训练。

Q2:代码中梯度清零的位置(optimizer.zero_grad () 在 step 后)是否正确?为什么?A:不正确。正确位置应在每个批次计算 loss 前:因为 step () 会更新参数,若清零在 step 后,下一批次的梯度会叠加当前批次的残留梯度,导致参数更新混乱,模型无法收敛。

Q3:model.eval () 在伪标签生成时为什么必须调用?A:伪标签生成需要模型输出稳定的预测结果,eval () 会关闭 Dropout(避免随机失活神经元)、固定 BatchNorm 的均值 / 方差(不使用当前批次统计值),保证预测结果可靠;若用 train () 模式,预测结果会随机波动,伪标签质量下降。

3. 拓展问答(拔高级)

Q1:如果想提升模型性能,你会从哪些方面优化?A:① 改用预训练模型(如 ResNet18)做迁移学习,利用预训练特征提升分类效果;② 动态调整伪标签阈值(模型收敛后提高阈值);③ 加入一致性正则化(对无标签数据做轻微增强,要求模型输出一致);④ 调整学习率(如学习率衰减)。

Q2:DataLoader 的 shuffle 参数在训练 / 验证 / 半监督模式下的设置逻辑是什么?A:训练集 shuffle=True:打乱样本顺序,避免模型学习到样本顺序规律;验证集 shuffle=False:保证评估结果可复现;半监督伪标签生成阶段 shuffle=False:保证预测结果与原始数据集索引对齐,避免伪标签错位。

Q3:你的 CNN 模型中,卷积层的输出尺寸如何计算?(以 conv1 为例:3→64 通道,3×3 卷积,步长 1,填充 1)A:卷积输出尺寸公式:Out=StrideIn−Kernel+2×Padding​+1;代入得:1224−3+2×1​+1=224,因此 conv1 输出为 64×224×224。

四、复试口述版(30 秒项目介绍)

我做的是一个基于 ResNet18 迁移学习的食物图像分类任务,因为数据集中有标签样本少、无标签样本多,所以我采用了伪标签半监督学习来提升模型效果。整体流程是:先用有标签数据训练模型,再用训练好的模型给无标签数据打高置信度伪标签,最后把伪标签数据加入训练,进一步提升分类准确率。

五、复试高频问题

问题 1:你为什么选择 ResNet18,而不是 ResNet50/101?

答:

  1. 轻量化:ResNet18 参数少(约 11M),训练速度快,适合我这个小数据集(food-11),不会过拟合;
  2. 残差结构:ResNet 的残差连接解决了深层网络的梯度消失问题,即使只有 18 层,也能提取足够的食物特征;
  3. 预训练优势:ResNet18 有丰富的 ImageNet 预训练权重,迁移到食物分类任务上,收敛更快、效果更好。

问题 2:你的半监督伪标签是怎么做的?为什么设置阈值 0.99?

答:

  1. 伪标签流程:① 先用有标签数据训练模型到一定精度;② 用模型对无标签数据预测,输出每个样本的类别概率;③ 只保留概率≥0.99 的样本作为伪标签数据;④ 把伪标签数据和有标签数据混合训练。
  2. 阈值 0.99 的原因:高阈值能保证伪标签的准确性 —— 如果阈值太低(比如 0.5),会引入大量错误标签,反而拖累模型;0.99 能筛选出模型 “非常确定” 的样本,保证伪标签质量。

问题 3:半监督相比纯监督,优势在哪里?

答:

  1. 数据利用:实际场景中,标注数据成本高(比如标注 1000 张食物图片要花很多时间),无标签数据容易获取;半监督能充分利用无标签数据,不用额外标注;
  2. 效果提升:我测试过,纯监督(只用有标签)准确率约 75%,半监督加入伪标签后提升到 80%,模型泛化能力更强;
  3. 训练稳定:高置信度伪标签不会引入噪声,训练过程更稳定,不会出现准确率波动大的情况。

1.为什么用迁移学习?

因为直接从零训练 CNN 容易过拟合,而 ResNet18 在 ImageNet 上预训练过,能提取通用图像特征,用迁移学习效果更好、收敛更快、更适合小数据集。

2.为什么用 ResNet18?

它是经典轻量化卷积网络,残差结构能解决深层网络梯度消失问题,适合做图像分类的 backbone。

3.什么是半监督学习?为什么要用?

半监督就是少量标签 + 大量无标签数据一起训练。我这个数据集无标签数据很多,用半监督能充分利用数据,提升模型泛化能力。

4.伪标签是什么?

用训练好的模型对无标签数据预测,只保留置信度很高的结果作为伪标签,再把这些数据加入训练。

5.为什么要设置置信度阈值(比如 0.99)?

阈值高,伪标签更准确,能避免错误标签影响模型训练。

6.你这个项目的难点和创新点是什么?

难点是小样本、标签少。创新点是迁移学习 + 半监督伪标签结合,在有限标注数据下也能达到不错的分类效果。

六、核心知识点拓展

(一)数据增广

  1. 什么是数据增广?为什么要用?答:数据增广是在不新增真实数据的前提下,对已有样本做合理变换,生成更多有效样本,用来扩充数据集、增加数据多样性。主要目的:
  • 解决数据量不足、标注成本高的问题
  • 防止模型过拟合,让模型学特征而不是死记样本
  • 提升模型泛化能力和鲁棒性
  • 缓解类别不平衡
  1. 图像增广常用方法有哪些?答:
  • 几何变换:翻转、旋转、平移、缩放、裁剪
  • 颜色 / 像素变换:调整亮度、对比度、饱和度、加噪声、模糊
  • 高级增广:Cutout、Random Erasing、MixUp、CutMix
  • 自动增广:AutoAugment、RandAugment核心原则:变换后类别不变。
  1. 什么是过拟合?数据增广为什么能缓解过拟合?答:过拟合是模型在训练集表现很好,但在测试集 / 真实场景表现差,说明模型只记住了训练样本,没学到通用规律。增广能缓解过拟合,因为:
  • 增加了数据多样性,让模型见更多 “变体”
  • 强制模型学习更本质的特征,而不是细节噪声
  • 相当于隐式增加训练数据量
  1. 增广是不是越多越好、越强越好?答:不是。
  • 增广必须符合真实场景分布,不能引入不合理噪声
  • 过度增广会破坏样本语义,导致模型学偏
  • 要根据任务选择:比如医学图像不能随便翻转、颜色增强
  • 一般用温和、可控的增广,配合验证集效果调参一句话:有用的多样性才叫增广,乱变叫噪声。
  1. 除了图像,文本 / 语音怎么增广?答:
  • 文本:同义词替换、随机插入 / 删除 / 交换词、回译、句式改写
  • 语音:加噪声、调整音量、时间拉伸、音高变换、时域偏移共同点:保持语义 / 标签不变,只改变形式,增加多样性

(二)优化器

核心定义

优化器就是根据损失函数算出的梯度,去自动更新模型的权重,让模型越来越准、损失越来越小的工具。模型一开始是瞎猜的,预测不准。损失函数算出误差(损失),然后算出梯度,告诉优化器:每个权重应该往哪个方向调、调多少,才能让误差变小。优化器就做一件事:拿着梯度,一步步更新权重,让模型从 “不会” 变 “会”。

作用总结
  • 根据梯度更新模型参数
  • 让损失不断下降
  • 让模型收敛、变准
复试高频问:你项目里用的是什么优化器?为什么?

答:我用的是 AdamW,它结合了动量和自适应学习率,收敛快、稳定,还能解耦权重衰减,防止过拟合,比 SGD 更适合图像分类任务。

常见优化器对比
  1. SGD(随机梯度下降):最基础的优化器,只靠当前梯度更新,简单稳定,但收敛慢、容易震荡。
  2. Adam:自适应学习率优化器,结合了动量 (Momentum) + RMSProp,收敛快、对学习率不敏感,深度学习最常用。
  3. AdamW:Adam 的改进版,解耦了权重衰减,正则化效果更好,更不容易过拟合,现在论文和比赛基本都用它。

核心结论

  • SGD:基础优化器,稳、泛化好,但慢。
  • Adam:动量 + 自适应学习率,快、方便。
  • AdamW:Adam 改进,解耦权重衰减,防过拟合更强。
  • 项目选 AdamW:收敛快、稳定、不容易过拟合。

(三)迁移学习

什么是迁移学习?

迁移学习就是把一个领域学到的知识,迁移到另一个相似领域,让新任务学得更快、更好、数据更少。通俗说:先在大数据上学通用特征,再在小数据上学专用特征。

为什么要用迁移学习?

  • 自己数据集太小,训练不动大模型
  • 训练速度快,不用从头训
  • 效果更好,泛化能力强
  • 避免过拟合

迁移学习在 CNN 里怎么做?

  • 加载在 ImageNet 上预训练好的模型(如 ResNet、VGG)
  • 冻结前面的卷积层(它们已经学会了边缘、纹理、颜色等通用特征)
  • 替换最后全连接层,改成自己的分类类别数
  • 只训练最后的全连接层,或 微调(fine-tune) 整个网络这就是 CNN 迁移学习的标准流程。

什么是 Fine-tune(微调)?

先用预训练模型初始化权重,不冻结全部卷积层,用小学习率一起训练,让特征更适配自己的数据集。

冻结层 vs 不冻结层 怎么选?

  • 数据很少 → 冻结前面层,只训练全连接层(把别人的模型(比如 ResNet)的参数冻住,不再改变)
  • 数据中等 / 较多 → 解冻,小学习率微调(一般用微调比较好,把别人的模型在我们自己的数据上训练)
  • 数据特别大 → 可以从头训练

你项目里用迁移学习了吗?为什么?

答:用了。因为我的数据集规模不大,直接训练容易过拟合,所以我用了在 ImageNet 上预训练的模型,冻结前面卷积层提取通用特征,只训练最后的分类层,这样训练更快、精度更高、泛化能力更好。

总结

迁移学习 = 用别人训好的模型,帮自己训小数据集核心:预训练 + 冻结 + 替换全连接 + 微调好处:数据少、训得快、效果好、不易过拟合

(四)ResNet

  1. 什么是 ResNet?ResNet 是 2015 年提出的深度残差网络,它的核心创新是残差连接(shortcut/skip connection),也就是让信息跨层直接传递,解决了深度网络越深准确率反而下降的问题,让网络可以训练到非常深(18、34、50、101 层)。一句话极简版:ResNet = 带残差连接的深度卷积网络,解决深层网络梯度消失,能训得更深、更准。

  2. ResNet 和 ImageNet 的关系标准答案:ImageNet 是一个超大规模图像分类数据集(1400 万图、1000 类)。ResNet 最初就是在 ImageNet 上训练出来的。我们做项目时,直接用在 ImageNet 上预训练好的 ResNet 权重,拿来做迁移学习。

复试口述:我使用了在 ImageNet 预训练的 ResNet18,它已经学到了丰富的通用特征,在我的小数据集上微调后,效果好、训练快、不容易过拟合。

复试高频 6 问

Q1:普通深网络为什么训不好?A:网络越深,梯度在反向传播时会越来越小,甚至消失,前面的层学不到东西,导致深层网络反而更差。

Q2:残差连接是什么?A:就是让输入 x 直接跨层加到输出上:输出 = F (x) + x,让网络学习残差 F (x),而不是直接学习映射,训练更容易。

Q3:残差连接为什么能解决梯度消失?A:因为梯度回传时多了一条直接通道,梯度可以直接传到前面的层,不会越来越小。

Q4:ResNet18、34、50 有什么区别?A:数字代表网络层数。数据集小 → 用 ResNet18 / 34;数据集大 → 用 ResNet50。复试项目选 ResNet18 最稳、最好训。

Q5:你为什么用 ResNet 不用 VGG?A:ResNet 可以做更深的网络,特征提取更强;残差连接缓解梯度消失,训练更稳定;在 ImageNet 上预训练的模型迁移效果更好;参数量比 VGG 小,精度更高,不容易过拟合。

Q6:你在项目里怎么用 ResNet?A:加载 ImageNet 预训练权重;冻结前面卷积层,提取通用特征;替换最后一层全连接层为自己的类别数;用小学习率微调,提升在自己数据集上的效果。

总结

ResNet 是深度残差网络,核心是残差连接,解决了深层网络梯度消失的问题。我使用的是在 ImageNet 上预训练的 ResNet18,它已经学到了丰富的底层特征,适合小数据集做迁移学习,训练更快、精度更高、泛化能力更好。

七、核心要点总结

  1. 半监督图像分类的核心是 “伪标签” 策略:利用高置信度的模型预测结果,将无标签数据转为带标签数据;
  2. 训练流程关键细节:梯度清零位置、model.train ()/eval () 的切换、数据增强的阶段差异;
  3. 高频考点:CrossEntropyLoss 使用、BatchNorm 作用、DataLoader 与 Dataset 的分工、伪标签阈值的选择逻辑;
  4. 迁移学习 + 半监督是小数据集分类任务的最优组合,ResNet18 是复试答辩的首选模型;
  5. 数据增广的核心是 “有用的多样性”,优化器优先选 AdamW(收敛快、防过拟合)。

在生成半监督数据加载器的逻辑中,为什么要设置 “训练轮数是 3 的倍数” 这个条件,这是半监督学习中伪标签生成策略的关键设计,核心目的是平衡模型稳定性、训练效率和伪标签质量。

一、核心原因解析

1. 等待模型先 “学到基础特征”(最核心)

半监督学习中,伪标签的质量完全依赖模型当前的预测能力:

  • 训练初期(前几轮),模型还在随机初始化阶段,预测结果几乎是 “瞎猜”,此时生成的伪标签全是噪声,用这些伪标签训练会误导模型,反而降低性能;
  • 设置 “每 3 轮生成一次”,是给模型留出足够的训练时间(先用有标签数据训练 3 轮),让模型先学到数据的基础特征,预测置信度更高,伪标签的准确性也会大幅提升。

简单说:3 轮是一个经验阈值,确保模型不是 “随机猜测”,而是有一定学习基础后再生成可靠的伪标签

2. 降低计算开销,提升训练效率

生成伪标签需要遍历所有无标签数据、用模型预测、筛选高置信度样本,这个过程会消耗额外的计算资源(GPU/CPU):

  • 如果每一轮都生成伪标签,会大幅增加训练时间(尤其是无标签数据量很大时);
  • 每 3 轮生成一次,在 “伪标签时效性” 和 “计算开销” 之间做了折中:既不会让伪标签过于陈旧(3 轮内模型变化不会太大),又能减少重复计算,提升整体训练效率。
3. 避免模型过度依赖伪标签,防止震荡

半监督训练中,伪标签本质是 “弱监督信号”,如果频繁更新伪标签:

  • 模型会反复适应新的伪标签,导致训练过程震荡(比如本轮用 A 类伪标签,下轮用 B 类伪标签);
  • 每 3 轮更新一次,能让模型先 “消化” 当前批次的伪标签,稳定学习到伪标签中的有效特征,再用新的伪标签迭代,避免参数频繁波动。
4. 经验性调参策略

“3 轮” 并不是固定不变的魔法数字,是开发者根据任务特性(数据量、模型复杂度、收敛速度)选择的经验值:

  • 数据量小、模型收敛快:可设为 2 轮;
  • 数据量大、模型收敛慢:可设为 5 轮;
  • 核心逻辑是:选择一个 “模型能在该轮数内学到有效特征,且不会因等待过久导致伪标签过时” 的间隔

二、补充说明(该条件的设计逻辑延伸)

结合你代码中的另一个条件 plt_val_acc[-1] > 0.06,两个条件是 “且” 的关系,进一步保证伪标签质量:

  • epoch % 3 == 0:保证模型有足够训练轮数;
  • plt_val_acc[-1] > 0.06:保证模型在验证集上有基本的预测能力(不是随机水平)。

如果去掉 “3 的倍数” 这个条件,直接用plt_val_acc[-1] > 0.06

  • 可能在训练第 1 轮就满足准确率条件,此时生成的伪标签质量极差,反而破坏训练;
  • 加上轮数限制,相当于给 “准确率条件” 加了一层 “时间缓冲”,确保准确率提升是模型真的学到了特征,而非偶然。
总结
  1. 核心目的:等待模型先用有标签数据训练 3 轮,学到基础特征后再生成高质量伪标签,避免初期噪声伪标签误导模型;
  2. 工程层面:减少伪标签生成的频率,降低计算开销,提升训练效率;
  3. 训练稳定性:避免伪标签更新过频导致模型训练震荡,让模型有时间消化每一批伪标签。

一、为什么这个项目适合复试?

  1. 技术栈有层次,能体现你的理解深度
    • 基础层:ResNet18 迁移学习(体现你懂 CNN、预训练、微调)
    • 进阶层:伪标签半监督学习(体现你懂半监督范式、数据增强、伪标签过滤策略)
    • 对比点:可以对比 “纯监督(少量标注数据)”vs“半监督(伪标签)” 的效果,有实验可讲
  2. 数据集易获取,复现成本低
    • 公开数据集:Food-101(101 类食物,10 万 + 图片)、UECFOOD100/256,甚至可以自己小范围采集 + 标注(比如 10 类常见食物:米饭、面条、红烧肉等)
    • 无需超大 GPU:ResNet18 轻量,单张 1080Ti/3060 就能跑,伪标签训练也不会太耗时
  3. 有 “可讲的亮点”,复试时能答出深度老师问 “你这个项目的核心创新 / 难点是什么?”,你能直接说:
    • 伪标签的噪声问题:如何过滤低置信度伪标签(比如设置置信度阈值 0.9)
    • 迁移学习的微调策略:冻结 ResNet 前几层 vs 全微调的效果对比
    • 数据增强:针对食物图像的增强(RandomResizedCrop、ColorJitter、HorizontalFlip)

二、项目落地的核心思路(简单易懂,好复现)

我给你梳理出核心步骤,你按这个来做,既能保证完成度,又能体现思考:

1. 数据预处理
  • 划分数据集:标注数据(比如 10%)+ 无标注数据(90%)
  • 数据增强:对标注数据做强增强(训练用),对无标注数据做弱增强(生成伪标签用)
2. 模型搭建
  • 主干:ResNet18 预训练权重(ImageNet),替换最后一层全连接层适配食物类别数
  • 半监督训练流程:

    plaintext

    ① 用少量标注数据微调ResNet18,得到“教师模型”
    ② 用教师模型对无标注数据预测,筛选置信度>阈值(如0.9)的样本,生成“伪标签”
    ③ 用“标注数据 + 伪标签数据”一起训练ResNet18(学生模型)
    ④ 迭代优化:用学生模型更新教师模型,重复②-③
    
3. 对比实验(复试重点)
  • 实验 1:纯监督(仅 10% 标注数据)的准确率
  • 实验 2:伪标签半监督(10% 标注 + 90% 伪标签)的准确率
  • 实验 3:调整伪标签置信度阈值(0.7/0.8/0.9),分析效果变化
  • 实验 4:对比 ResNet18 vs 简单 CNN(如 AlexNet)的半监督效果
总结
  1. 核心亮点是伪标签的噪声控制迁移学习微调策略,复试时重点讲这两点。

你这个迁移学习是怎么做的?

我采用的是基于 ResNet18 的迁移学习 + 全量微调方案。首先加载在 ImageNet 上预训练好的 ResNet18 权重,利用它学到的通用图像特征,来提升小样本食物分类的效果。然后把模型最后的全连接层替换成 11 类,适配我的食物分类任务。在训练时,我对整个网络进行全量微调,让所有层的参数都在我的数据集上进一步学习食物相关特征,再结合伪标签半监督学习,利用大量无标注数据进一步提升模型泛化能力。实验结果表明,这种做法在有限标注数据下,效果明显优于从头训练。


为什么用全量微调,不用分层冻结?

因为我的数据集是食物图像,和 ImageNet 的自然图像分布比较接近,而且 ResNet18 本身模型不大,全量微调不会出现严重过拟合,训练也更简单稳定。我在实验中对比过,全量微调已经能达到很好的效果,所以最终选择了简洁且有效的全量微调方案。

泛化

在编程和机器学习领域,泛化(Generalization) 是一个核心概念,简单来说就是:模型在 “从未见过的新数据” 上的表现能力。我会用最通俗的方式帮你理解,结合例子和场景,让你一眼看懂。

一、先讲个生活例子(秒懂泛化)

假设你在学认 “猫”:

  • 你只见过白色的布偶猫,于是你以为 “只有白色、长毛、蓝眼睛的动物才是猫”—— 这就是泛化能力差(只能识别见过的布偶猫,遇到橘猫、黑猫就不认识了);
  • 你见过各种品种、颜色、大小的猫,总结出 “猫有尖耳朵、胡须、会喵喵叫”—— 这就是泛化能力强(能识别任何没见过的猫)。

对应到机器学习中:

  • 模型训练时用的是 “训练集”(你见过的布偶猫);
  • 模型测试 / 实际使用时面对的是 “测试集 / 新数据”(你没见过的橘猫、黑猫);
  • 泛化能力 = 模型能不能把从训练集学到的规律,用到新数据上。

二、泛化的核心:避免 “过拟合”,追求 “通用规律”

1. 反面:过拟合(泛化能力差)

模型把训练集的 “细节 / 噪声” 当成了 “通用规律”,比如:

  • 训练集里的猫都趴在沙发上,模型就认为 “趴在沙发上的动物才是猫”;
  • 反映在代码 / 指标上:训练集损失极低(预测极准),但测试集损失极高(预测极烂)。
2. 正面:良好的泛化

模型学到了数据的 “本质规律”,比如:

  • 不管猫在哪里、什么颜色,只要符合猫的核心特征就识别为猫;
  • 反映在指标上:训练集损失和测试集损失都很低,且两者差距小。
3. 代码 / 可视化理解(分类任务示例)
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
import numpy as np

# 模拟:简单的二分类数据(训练集有噪声)
np.random.seed(42)
# 训练数据:正样本(猫)带随机噪声,负样本(狗)
train_x = np.concatenate([
    np.random.normal(0, 0.5, (50, 2)),   # 猫的核心特征(0,0)附近
    np.random.normal(3, 0.5, (50, 2))    # 狗的核心特征(3,3)附近
])
train_y = np.array([0]*50 + [1]*50)  # 0=猫,1=狗

# 测试数据:无噪声的纯规律数据
test_x = np.concatenate([
    np.random.normal(0, 0.3, (20, 2)),
    np.random.normal(3, 0.3, (20, 2))
])
test_y = np.array([0]*20 + [1]*20)

# 模拟两种模型:
# 模型1:过拟合(泛化差)—— 记住了训练集的噪声
def model_overfit(x):
    # 复杂模型,拟合了所有训练噪声
    return np.where(np.sum((x - train_x)**2, axis=1) < 0.1, train_y, 1-train_y)

# 模型2:泛化好 —— 学到了核心规律(距离(0,0)近是猫,距离(3,3)近是狗)
def model_good(x):
    dist0 = np.sum((x - [0,0])**2, axis=1)
    dist3 = np.sum((x - [3,3])**2, axis=1)
    return np.where(dist0 < dist3, 0, 1)

# 计算准确率
train_acc_overfit = (model_overfit(train_x) == train_y).mean()
test_acc_overfit = (model_overfit(test_x) == test_y).mean()

train_acc_good = (model_good(train_x) == train_y).mean()
test_acc_good = (model_good(test_x) == test_y).mean()

print(f"过拟合模型 - 训练集准确率:{train_acc_overfit:.2f},测试集准确率:{test_acc_overfit:.2f}")
print(f"泛化好的模型 - 训练集准确率:{train_acc_good:.2f},测试集准确率:{test_acc_good:.2f}")
输出结果:
过拟合模型 - 训练集准确率:1.00,测试集准确率:0.55
泛化好的模型 - 训练集准确率:0.98,测试集准确率:1.00

可以清晰看到:

  • 过拟合模型在训练集上完美(100% 准),但新数据上只有 55% 准(泛化差);
  • 泛化好的模型训练集稍差(98%),但新数据上 100% 准(能通用)。

三、如何提升模型的泛化能力?(实用方法)

  1. 增加数据量:更多样的训练数据,让模型见得多、学得全;
  2. 正则化:比如 L1/L2 正则、Dropout,限制模型复杂度,避免死记硬背;
  3. 早停(Early Stopping):训练到测试集损失不再下降时停止,避免过度拟合;
  4. 数据增强:对训练数据做随机变换(如图像旋转、裁剪),增加数据多样性;
  5. 简化模型:减少网络层数 / 参数,避免模型 “太聪明” 而记住噪声。
总结
  1. 泛化的核心定义:模型在未见过的新数据上的预测能力,是衡量模型 “实用性” 的关键;
  2. 泛化的对立面:过拟合(模型只记住训练集细节,无法适应新数据);
  3. 提升泛化的目标:让模型学到数据的本质规律,而非训练集的噪声,做到 “举一反三”。

简单来说,泛化能力强的模型,才是 “能用的模型”—— 比如你训练的猫识别模型,不仅能识别训练过的照片,还能识别你手机里新拍的猫照片。

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐