【深度学习】卷积神经网络 数据增强、保存最优模型实现,详细解读 文章目录一、数据增强详解1. 什么是数据增强2. 核心目标3. 常用数据增强方法4. 数据预处理与增强的代码实现5. 自定义数据集类与增强集成二、训练过程中保存最优模型1. 为什么要保存最优模型2. 定义 CNN 模型用于图像分类3. 训练函数4. 测试函数与最优模型保存两种方式5. 训练主循环6. 模型文件说明一、数据增强详解1. 什么是数据增强数据增强Data Augmentation是指在不改变原始数据语义的前提下通过一系列随机变换和组合操作对已有训练样本进行扩展生成大量“新”样本的过程。其本质是人为增加训练集的规模和多样性从而使深度学习模型在面对实际场景中的各种变化如光照、角度、遮挡时具备更强的适应能力和稳定性。2. 核心目标核心目标是模拟现实世界的复杂多变环境迫使模型学习到更抽象、更鲁棒的特征表示而非仅仅记住训练集的具体样本从而有效降低过拟合风险提升模型的泛化性能。3. 常用数据增强方法方法描述随机旋转将图像绕中心旋转一定角度如 -45°~45°水平/垂直翻转沿水平或垂直轴镜像翻转图像随机缩放按比例放大或缩小图像尺寸随机平移沿水平或垂直方向移动若干像素随机裁剪从原图中截取部分区域亮度/对比度/饱和度调整改变颜色空间的数值添加噪声叠加高斯、椒盐等噪声几何扭曲仿射变换、弹性变形等4. 数据预处理与增强的代码实现在 PyTorch 中通常使用 torchvision.transforms 组合多种操作并分别对训练集和验证集设置不同的处理流水线。import torch from torch.utils.data import DataLoader,Dataset import numpy as np from PIL import Image from torchvision import transforms # 定义训练和验证阶段的不同预处理 data_transforms{train:transforms.Compose([transforms.Resize([300,300]),# 统一尺寸 transforms.RandomRotation(degrees45),# 随机旋转[-45°,45°]transforms.CenterCrop(256),# 中心裁剪为256x256 transforms.RandomHorizontalFlip(p0.5),# 水平翻转概率50%transforms.RandomVerticalFlip(p0.5),# 垂直翻转概率50%transforms.ColorJitter(brightness0.2,contrast0.1,saturation0.1,hue0.1),# 颜色抖动 transforms.RandomGrayscale(p0.1),#10%概率转为灰度图 transforms.ToTensor(),# 转为 Tensor 并归一化到[0,1]transforms.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225])# ImageNet 标准化]),valid:transforms.Compose([transforms.Resize([256,256]),transforms.ToTensor(),transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])])}提醒数据增强并非总能提升效果需根据具体任务和数据集进行调优但通常情况下会带来正向收益。5. 自定义数据集类与增强集成我们通过继承 Dataset 类在getitem中应用上述变换从而在每次读取样本时动态生成增强后的图像。classFoodDataset(Dataset):自定义食物图片数据集def__init__(self,file_path,transformNone):self.file_pathfile_path self.transformtransform self.image_paths[]self.labels[]# 解析文件每行格式图片路径 类别标签 withopen(file_path)as f:lines[line.strip().split()forline in f.readlines()]forimg_path,label in lines:self.image_paths.append(img_path)self.labels.append(label)def__len__(self):returnlen(self.image_paths)def__getitem__(self,idx):imgImage.open(self.image_paths[idx])ifself.transform:imgself.transform(img)# 标签转换为 Tensor labeltorch.from_numpy(np.array(int(self.labels[idx]),dtypenp.int64))returnimg,label # 实例化训练集和验证集 train_datasetFoodDataset(file_path./trainda.txt,transformdata_transforms[train])valid_datasetFoodDataset(file_path./testda.txt,transformdata_transforms[valid])# 设备选择 devicecudaiftorch.cuda.is_available()elsempsiftorch.backends.mps.is_available()elsecpuprint(f当前使用的设备: {device})其中 trainda.txt 和 testda.txt 的内容格式为…每行一个样本路径与标签用空格分隔二、训练过程中保存最优模型1. 为什么要保存最优模型在深度学习的迭代训练中模型参数会随着优化步骤不断更新。通常我们会在每个 Epoch 结束时在验证集上评估性能并将当前验证集上表现最好的模型参数持久化到磁盘常见扩展名为 .pt、.pth 或 .t7。这样做可以避免因过拟合或训练后期震荡而丢失最佳状态也便于后续部署或继续微调。2. 定义 CNN 模型用于图像分类这里构建一个三层卷积 全连接的简单 CNN输入为 3×256×256 的 RGB 图像。from torch import nn classCNN(nn.Module):def__init__(self):super(CNN,self).__init__()self.conv1nn.Sequential(nn.Conv2d(in_channels3,out_channels16,kernel_size5,stride1,padding2),#-(16,256,256)nn.ReLU(),nn.MaxPool2d(kernel_size2)#-(16,128,128))self.conv2nn.Sequential(nn.Conv2d(16,32,kernel_size5,stride1,padding2),#-(32,128,128)nn.ReLU(),nn.MaxPool2d(2)#-(32,64,64))self.conv3nn.Sequential(nn.Conv2d(32,128,kernel_size5,stride1,padding2),#-(128,64,64)nn.ReLU())self.fcnn.Linear(128*64*64,20)# 假设类别数为20defforward(self,x):xself.conv1(x)xself.conv2(x)xself.conv3(x)xx.view(x.size(0),-1)outself.fc(x)returnout modelCNN().to(device)print(model)3. 训练函数deftrain_one_epoch(dataloader,model,loss_fn,optimizer):model.train()batch_idx1forX,y in dataloader:X,yX.to(device),y.to(device)predmodel(X)# 前向传播 lossloss_fn(pred,y)# 计算损失 optimizer.zero_grad()# 清零梯度 loss.backward()# 反向传播 optimizer.step()# 更新参数ifbatch_idx%1000:print(f 批次 {batch_idx} 损失: {loss.item():.4f})batch_idx14. 测试函数与最优模型保存两种方式定义全局变量 best_acc 跟踪最高准确率若当前验证准确率更高则保存模型。best_acc0.0defevaluate_and_save(dataloader,model,loss_fn):global best_acc sizelen(dataloader.dataset)num_batcheslen(dataloader)model.eval()test_loss,correct0.0,0with torch.no_grad():forX,y in dataloader:X,yX.to(device),y.to(device)predmodel(X)test_lossloss_fn(pred,y).item()correct(pred.argmax(1)y).type(torch.float).sum().item()test_loss/num_batches accuracycorrect/sizeprint(f验证结果: 准确率 {accuracy:.2%}, 平均损失 {test_loss:.4f})# 保存最优模型ifaccuracybest_acc:best_accaccuracy # 方式一仅保存模型参数推荐占用空间小#torch.save(model.state_dict(),best_params.pth)# 方式二保存完整模型包含架构和参数 torch.save(model,best_model.pt)print(f模型已更新保存当前最佳准确率: {best_acc:.2%})5. 训练主循环loss_fnnn.CrossEntropyLoss()optimizertorch.optim.Adam(model.parameters(),lr0.001)train_loaderDataLoader(train_dataset,batch_size64,shuffleTrue)valid_loaderDataLoader(valid_dataset,batch_size64,shuffleFalse)epochs150forepoch inrange(epochs):print(f\nEpoch {epoch1}/{epochs})train_one_epoch(train_loader,model,loss_fn,optimizer)evaluate_and_save(valid_loader,model,loss_fn)运行过程会逐轮输出训练损失和验证准确率当验证准确率超过历史最佳时自动保存新模型。6. 模型文件说明训练结束后best_model.pt或 best_params.pth即为最优模型文件。加载方法若保存的是完整模型modeltorch.load(best_model.pt)若保存的只是状态字典先实例化模型结构再model.load_state_dict(torch.load(best_params.pth))

相关新闻

最新新闻

线下销售过程管理缺数据支撑,2026 智能工牌硬件部署完整攻略

线下销售过程管理缺数据支撑,2026 智能工牌硬件部署完整攻略

企业评估智能工牌时,容易陷入只看硬件参数或录音时长的误区。实际上,一套能真正产生业务价值的方案,不仅需要稳定的线下采集能力,更考验其在复杂环境下的AI分析深度、行业规模化验证以及后续交付运营服务。 线下销售团队选智能工…

2026/7/22 14:12:45
C++ STL核心价值解析:从性能优化到工程实践

C++ STL核心价值解析:从性能优化到工程实践

1. 从“造轮子”到“用轮子”:一个C新手的必经之路 刚接触C那会儿,我特别痴迷于“自己动手,丰衣足食”。老师讲链表,我就吭哧吭哧写一个 MyList ;讲到动态数组,我又去实现一个 MyVector 。看着自己写的…

2026/7/22 14:12:45
中小企业建站平台怎么选?预算、上线效率与SEO基础更值得比较

中小企业建站平台怎么选?预算、上线效率与SEO基础更值得比较

“中小企业建站平台首选”没有适用于所有企业的固定答案。展示官网、营销获客、多语言网站和复杂业务系统,对预算、功能和维护人员的要求并不相同。更稳妥的做法是先确定网站承担什么任务,再比较年费、上线周期、后台操作、搜索配置和扩展边界&#xff0…

2026/7/22 14:12:45
Copilot与pytest结合:AI如何提升自动化测试开发效率与质量

Copilot与pytest结合:AI如何提升自动化测试开发效率与质量

1. 项目概述:当AI副驾驶遇上自动化测试 最近半年,我团队内部一直在推动测试左移和自动化测试覆盖率的提升。在这个过程中,一个绕不开的痛点就是编写和维护自动化测试脚本的投入产出比。测试工程师不仅要懂业务,还要精通编程和测试…

2026/7/22 14:12:45
关于一只鸟的笑话

关于一只鸟的笑话

一只麻雀去参加鸟类选美大赛,结果第一轮就被淘汰了。它委屈地问评委:"为什么?我羽毛也很漂亮啊!"评委说:"你确实不错,但本次大赛要求选手必须会鸟语。"麻雀愣住了:"可…

2026/7/22 14:12:45
大学生AI预测引擎项目解析:从Vibe Coding到Streamlit部署

大学生AI预测引擎项目解析:从Vibe Coding到Streamlit部署

1. 项目背景与现象解析2023年GitHub全球趋势榜上出现了一个令人惊讶的现象:一个名为"AI预测引擎"的大学生项目,在短短十天内两次登顶全球热门榜首。这个由单人开发的项目Star数在48小时内突破5000,Forks超过800次,成为当…

2026/7/22 14:07:45

月新闻