半监督图像分类模型

发布时间:2026/7/24 21:15:40

半监督图像分类模型 import random import torch # PyTorch核心库张量、模型、梯度计算 import torch.nn as nn # 神经网络层卷积、全连接、损失函数等 import numpy as np # 数值计算数组处理、矩阵运算 import os # 系统操作文件/文件夹路径、遍历 from torch.utils.data import Dataset,DataLoader # 数据集/数据加载器批处理、打乱 from PIL import Image # 图片读取/处理PIL库 from torchvision import transforms # 图片预处理裁剪、旋转、转张量等 import time # 计时统计每轮训练耗时 import matplotlib.pyplot as plt # 绘图损失/准确率曲线 from model_utils.model import initialize_model # 自定义工具初始化预训练模型如ResNet18 #作用深度学习中随机操作权重初始化、数据打乱、Dropout 等会导致结果波动固定所有随机种子后每次运行代码的结果完全一致方便调试和对比。 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 transforms.Compose( [ transforms.ToPILImage(), # 把numpy数组转PIL图片适配后续操作 transforms.RandomResizedCrop(224), # 随机裁剪缩放模拟不同视角 transforms.RandomRotation(50), # 随机旋转±50度增强鲁棒性 transforms.ToTensor() # 转张量 ] ) # 验证/测试集预处理无增强用原图评估 val_transform transforms.Compose( [ transforms.ToPILImage(),# 转PIL图片 transforms.ToTensor()# 转张量 ] ) class food_Dataset(Dataset): def __init__(self,path,modetrain):#根据模式和文件地址把需要的X、Y读出来 self.mode mode if mode semi: self.Xself.read_file(path)# 无标注数据只加载图片X else : self.X,self.Y self.read_file(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):#具体执行读取X、Y的函数 if self.mode semi: file_list os.listdir(path) #遍历无标签文件夹下所有图片 xi np.zeros((len(file_list), HW, HW, 3), dtypenp.uint8) # 创建了file_list个格子每个格子都是224*224*3的数组 for j, img_name in enumerate(file_list): img_path os.path.join(path, img_name) # 拼接图片完整路径 img Image.open(img_path) # 打开图片 img img.resize((HW, HW)) # 缩放到224×224 xi[j, ...] img # 把第j张图片存入数组 print(读到了%d个数据 % len(xi)) return xi #返回读到的图片即X数组 # 处理有标注数据train/val模式按类别文件夹读取自动生成标签 else: for i in range(11): # 遍历00~10共11个类别文件夹对应11类食物 file_dir path /%02d % i # 类别文件夹路径00、01...10 file_list os.listdir(file_dir) # 遍历该类别下所有图片 xi np.zeros((len(file_list), HW, HW, 3), dtypenp.uint8) # 创建了file_list个格子每个格子都是224*224*3的数组 yi np.zeros(len(file_list), dtypenp.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 # 把图片放在第j个 yi[j] i # 存标签当前文件夹对应类别i # 合并所有类别的数据 if i 0: X xi Y yi else: X np.concatenate((X, xi), axis0) # 按纵向拼接图片 Y np.concatenate((Y, yi), axis0) # 按纵向拼接标签 print(读到了%d个数据 % len(Y)) return X, Y # xxxx_loader 正是通过调用你自定义的 __getitem__() 方法来获取批量数据的 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]# 有标注数据返回「预处理后的图片 标签」 def __len__(self):#返回数据集总长度DataLoader需要知道总样本数 return len(self.X) class semiDataset(Dataset):#半监督数据集类semiDataset给无标注数据 “伪标签” def __init__(self,no_label_loader,model,device,thres0.99): # 调用get_label方法用模型给无标注数据生成“高置信度伪标签” x,y self.get_label(no_label_loader,model,device,thres) if x[]:#没有数据符合要求 self.flagFalse else: self.flagTrue self.X np.array(x) # 高置信度无标注图片 self.Y torch.LongTensor(y) # 对应的伪标签 self.transform train_transform # 用训练集的增强方式 def get_label(self,no_label_loader,model,device,thres): model model.to(device) pred_prob [] # 存储每个样本的预测最大概率 labels [] # 存储每个样本的预测类别 x [] # 存储满足置信度的图片 y [] # 存储对应的伪标签 soft nn.Softmax() # 将原始输出转换为和为 1 的概率分布 with torch.no_grad(): # 关闭梯度计算仅预测不训练 for bat_x,_ in no_label_loader:# 遍历无标注数据加载器 bat_x bat_x.to(device)# 图片放到设备上 pred model(bat_x) # 模型预测 pred_soft soft(pred) # 转概率分布 pred_max,pred_value pred_soft.max(1)# 取每个样本的最大概率pred_max和对应类别pred_value#1表示横向 # 把结果转numpy并存入列表 pred_prob.extend(pred_max.cpu().numpy().tolist()) labels.extend(pred_value.cpu().numpy().tolist()) # 筛选高置信度样本概率阈值thres才保留避免伪标签错误 for index,prob in enumerate(pred_prob): if prob thres: x.append(no_label_loader.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) def get_semi_loader(no_label_loader,model,device,thres): semiset semiDataset(no_label_loader,model,device,thres) if semiset.flag False: return None else : semi_loader DataLoader(semiset,batch_size4,shuffleFalse) ######################模型部分 class myModel(nn.Module): def __init__(self,num_class): super(myModel, self).__init__() # 3*224*224 -- 512*7*7 -- flatten -- 全连接 self.layer1 nn.Sequential( nn.Conv2d(3, 64, 3, 1, 1), # ➡ 64*224*224 nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2) # ➡ 64*112*112 ) self.layer2 nn.Sequential( nn.Conv2d(64, 128, 3, 1, 1), # ➡ 128*112*112 nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2) # ➡ 128*56*56 ) self.layer3 nn.Sequential( nn.Conv2d(128, 256, 3, 1, 1), # ➡ 256*112*112 nn.BatchNorm2d(256), nn.ReLU(), nn.MaxPool2d(2) # ➡ 256*28*28 ) self.layer4 nn.Sequential( nn.Conv2d(256, 512, 3, 1, 1), # ➡ 512*112*112 nn.BatchNorm2d(512), nn.ReLU(), nn.MaxPool2d(2) # ➡ 512*14*14 ) self.pool1 nn.MaxPool2d(2) # ➡ 512*7*7 self.fc1 nn.Linear(25088,1000) # 25088➡1000 self.relu1 nn.ReLU() self.fc2 nn.Linear(1000,num_class) # 1000➡11 def forward(self,x): x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) x self.pool1(x) x torch.flatten(x,1) x self.fc1(x) x self.relu1(x) x self.fc2(x) return x ################################定义训练与验证函数 # model定义好的模型train_loader/val_loader训练/验证数据加载器lr学习率 # optimizer优化器device训练设备cpu/gpuepochs训练轮数save_path最优模型保存路径 #实现「训练梯度下降→验证评估效果→保存最优模型→绘制损失曲线」全流程 def train_val(model,train_loader,val_loader,no_label_loader,lr,optimizer,device,epochs,thres,save_path): model model.to(device) #即插即用所以为了防止意外再放一次 semi_loader None # 初始化半监督数据加载器初始无 plt_train_loss [] # 记录每轮训练的平均损失用于画图 plt_val_loss [] # 记录每轮验证的平均损失用于画图 plt_train_acc [] plt_val_acc [] max_acc0.0 #初始化最大准确率用于保存最优模型 ######训练过程 for epoch in range(epochs): #发枪指令模型训练的开始 model.train() #模型切换为训练模式启用梯度计算、Dropout等 start_time time.time() #记录本轮训练开始时间计算耗时 train_loss 0.0 train_acc 0.0 val_loss 0.0 val_acc 0.0 semi_loss 0.0 semi_acc 0.0 # 1. 有标注数据训练核心监督训练 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() # 累加批次损失转cpu取数值 # 累计批次准确次数预测类别argmax和真实标签对比 train_acc np.sum(np.argmax(pred.detach().cpu().numpy(), axis1) 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() # 累加半监督损失 #计算半监督训练的批次准确次数 semi_acc np.sum(np.argmax(pred.detach().cpu().numpy(), axis1) target.cpu().numpy()) print(半监督数据集的训练准确率为,semi_acc/semi_loader.dataset.__len__()) ######验证过程每一轮的流程监督与半监督集训练➡验证集评估模型不更新参数 model.eval() #模型切换为验证模式禁用梯度计算、Dropout等 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(), axis1) target.cpu().numpy()) plt_val_loss.append(val_loss / val_loader.__len__()) # 计算本轮验证集的平均损失存入列表 plt_val_acc.append(val_acc / val_loader.dataset.__len__())# 记录本轮验证集准确率到列表 # 3. 每3轮生成一次半监督数据验证集准确率0.05才生成避免初期模型差导致伪标签错误 if epoch%30 and plt_val_acc[-1]0.05: semi_loader get_semi_loader(no_label_loader,model,device,thres) # 保存最优模型验证准确率更高则覆盖保存 if val_accmax_acc: max_acc val_acc torch.save(model,save_path) # 保存整个模型到指定路径 # 打印本轮训练信息轮数、耗时、训练损失、验证损失 print([%03d/%03d] %2.2f sec(s) train_loss: %.6f val_loss:%.6f train_acc: %.6f val_acc:%.6f% \ (epoch, epochs, time.time()-start_time, plt_train_loss[-1], plt_val_loss[-1],plt_train_acc[-1],plt_val_acc[-1])) # 训练结束后绘制训练/验证损失曲线 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 rD:\武大考研\Pycharm与复试项目\复试项目课件\第④、⑤节图片分类知识点代码\第四五节_分类代码2\food_classification\food-11_sample\training\labeled val_path rD:\武大考研\Pycharm与复试项目\复试项目课件\第④、⑤节图片分类知识点代码\第四五节_分类代码2\food_classification\food-11_sample\validation no_label_path rD:\武大考研\Pycharm与复试项目\复试项目课件\第④、⑤节图片分类知识点代码\第四五节_分类代码2\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_size4,shuffleTrue) # 打乱 并 每4个成一批 val_loader DataLoader(val_set,batch_size4,shuffleTrue) no_label_loader DataLoader(no_label_set,batch_size4,shuffleFalse) # from torchvision.models import resnet18 # model resnet18(pretrainedTrue) #不仅用模型还用参数 # in_fetures model.fc.in_features # model.fc nn.Linear(in_fetures, 11) # model myModel(11) #这是用自己写的模型 model,_ initialize_model(resnet18,11,use_pretrainedTrue) lr 0.001 loss nn.CrossEntropyLoss()# 分类任务损失函数自带Softmax # Adam①会综合之前的梯度与现在的梯度 ②会自适应改变lr AdamW就是Adam加上权重衰减 optimizer torch.optim.AdamW(model.parameters(),lrlr,weight_decay1e-4) device cuda if torch.cuda.is_available() else cpu save_path model_save/best_model.pth # 最优模型保存路径 epochs 15 # 训练轮数 thres 0.1 # 半监督伪标签置信度阈值0.1为测试用实际建议0.9 train_val(model,train_loader,val_loader,no_label_loader,lr,optimizer,device,epochs,thres,save_path) 下图是我对整个代码流程的理解如有错误欢迎批评指正

相关新闻