ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

初识深度学习——数据增强与模型保存

初识深度学习——数据增强与模型保存 一、引言让模型长见识让成果留下来在前两篇博客中我们完成了从自定义数据集到CNN模型训练的完整流程。但如果你仔细回顾会发现一个潜在的问题训练数据太单一。模型只见过正着摆放的、亮度固定的、同一角度的物品。一旦测试图片稍有旋转、翻转或颜色变化模型就可能傻眼——这就是所谓的过拟合。解决这个问题的利器就是数据增强Data Augmentation。它通过对训练图片进行随机变换旋转、翻转、调色等人为地制造出更多样化的训练样本让模型学会忽略这些无关变化专注于真正的类别特征。与此同时训练了若干轮之后我们得到了一个不错的模型——但如果没有保存下次就得从头再来。模型保存让训练成果得以持久化随时可以加载使用。本篇博客将围绕这两大主题展开基于完整代码讲解数据增强、标准化、最优模型保存三大核心知识点。二、数据增强2.1 训练集 vs 验证集两套不同的变换策略代码中最醒目的设计是定义了两套变换流程data_transforms { trainda: transforms.Compose([ transforms.RandomRotation(45), transforms.CenterCrop(256), transforms.RandomHorizontalFlip(p0.5), transforms.RandomVerticalFlip(p0.5), transforms.ColorJitter(brightness0.2, contrast0.1, saturation0.1, hue0.1), transforms.RandomGrayscale(p0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), valid: transforms.Compose([ transforms.Resize([256, 256]), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]), }核心原则训练集使用随机变换数据增强让每个epoch看到的图片都略有不同提升泛化能力。验证集/测试集只做必要的尺寸统一和标准化不能加入随机性否则评估结果不稳定。2.2 常用数据增强方法详解1RandomRotation——随机旋转transforms.RandomRotation(45)在-45°到45°之间随机旋转图片。这模拟了拍摄角度不同的情况让模型学会识别旋转后的物体。2CenterCrop——中心裁剪transforms.CenterCrop(256)从图像中心裁剪出256×256的区域。配合RandomRotation使用可以裁掉旋转后产生的黑边保证输入尺寸一致。3RandomHorizontalFlip / RandomVerticalFlip——随机翻转transforms.RandomHorizontalFlip(p0.5) # 水平翻转50%概率 transforms.RandomVerticalFlip(p0.5) # 垂直翻转50%概率以指定概率对图片进行翻转。水平翻转适合大多数场景如动物、车辆垂直翻转则要谨慎使用对于人脸等有方向性的物体可能不合适。4ColorJitter——颜色抖动transforms.ColorJitter(brightness0.2, contrast0.1, saturation0.1, hue0.1)随机调整图像的亮度、对比度、饱和度、色相。这模拟了不同光照条件下的拍摄效果提升模型对光照变化的鲁棒性。参数含义取值范围brightness亮度0.2表示在[0.8, 1.2]倍之间随机调整contrast对比度同上saturation饱和度同上hue色相0.1表示在[-0.1, 0.1]之间偏移5RandomGrayscale——随机灰度化transforms.RandomGrayscale(p0.1)以10%的概率将彩色图片转为灰度图RGB。这强制模型不依赖颜色信息学习更本质的形状特征。2.3 ToTensor 与 Normalize——标准化的两步ToTensor从PIL到张量transforms.ToTensor()作用将PIL图像或NumPy数组转为PyTorch张量将像素值从0-255缩放到0-1将通道维度从 HWC 转为 CHWPyTorch要求Normalize标准化到标准正态分布transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])计算方式为什么用这组特定的均值和标准差它们是ImageNet数据集上统计出来的RGB三通道均值和标准差。由于大多数预训练模型都是在ImageNet上训练的使用相同的标准化参数可以保持数据分布一致。标准化的意义让数据分布接近标准正态分布加速梯度下降收敛消除不同通道之间的量纲差异是迁移学习中使用预训练模型时的必要步骤三、自定义数据集回顾数据集类与上一篇博客一致class food_dataset(Dataset): def __init__(self, file_path, transformNone): self.file_path file_path self.imgs [] self.labels [] self.transform transform with open(self.file_path) as f: samples [x.strip().split( ) for x in f.readlines()] for img_path, label in samples: self.imgs.append(img_path) self.labels.append(label) def __len__(self): return len(self.imgs) def __getitem__(self, idx): image Image.open(self.imgs[idx]) if self.transform: image self.transform(image) label torch.from_numpy(np.array(self.labels[idx], dtypenp.int64)) return image, label然后分别用训练变换和验证变换创建数据集training_data food_dataset(file_path./train.txt, transformdata_transforms[trainda]) test_data food_dataset(file_path./test.txt, transformdata_transforms[valid]) train_dataloader DataLoader(training_data, batch_size64, shuffleTrue) test_dataloader DataLoader(test_data, batch_size64, shuffleTrue)四、CNN模型结构模型与上一篇相同针对 3×256×256 彩色输入输出20个类别class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() self.conv1 nn.Sequential( nn.Conv2d(3, 16, 5, 1, 2), nn.ReLU(), nn.MaxPool2d(2), ) self.conv2 nn.Sequential( nn.Conv2d(16, 32, 5, 1, 2), nn.ReLU(), nn.Conv2d(32, 64, 5, 1, 2), nn.ReLU(), nn.MaxPool2d(2), ) self.conv3 nn.Sequential( nn.Conv2d(64, 128, 5, 1, 2), nn.ReLU(), ) self.out nn.Linear(128*64*64, 20) def forward(self, x): x self.conv1(x) x self.conv2(x) x self.conv3(x) x x.view(x.size(0), -1) output self.out(x) return output尺寸变化输入3×256×256conv1后16×128×128conv2后64×64×64conv3后128×64×64展平128×64×64 524288 维输出20类五、训练函数def train(dataloader, model, loss_fn, optimizer): model.train() batch_size_num 1 for x, y in dataloader: x, y x.to(device), y.to(device) pred model.forward(x) loss loss_fn(pred, y) optimizer.zero_grad() loss.backward() optimizer.step() loss_value loss.item() if batch_size_num % 1 0: print(floss:{loss_value:7f} [number:{batch_size_num}]) batch_size_num 1训练过程与之前一致前向传播→计算损失→梯度清零→反向传播→更新参数。六、模型保存这是本篇博客的重点。测试函数中当模型准确率创新高时会保存模型best_acc 0 def test(dataloader, model, loss_fn): global best_acc size len(dataloader.dataset) num_batches len(dataloader) model.eval() test_loss, correct 0, 0 with torch.no_grad(): for X, y in dataloader: X, y X.to(device), y.to(device) pred model.forward(X) test_loss loss_fn(pred, y).item() correct (pred.argmax(1) y).type(torch.float).sum().item() test_loss / num_batches correct / size print(fTest result: \n Accuracy: {(100*correct)}%, Avg loss: {test_loss}) # 保存最优模型 if correct best_acc: best_acc correct print(model.state_dict().keys()) torch.save(model.state_dict(), fxxxxxx.pth) script_model torch.jit.script(model) torch.jit.save(script_model, fxxxxxxx.pth)6.1 两种保存方式的对比方式一保存模型参数state_dicttorch.save(model.state_dict(), xxxxxx.pth)保存内容仅保存模型的参数权重w和偏置b不包含模型结构。加载方式model CNN() # 先定义模型结构 model.load_state_dict(torch.load(xxxxxx.pth)) model.eval()优点文件小只保存参数灵活可以加载到不同但结构相同的模型是PyTorch推荐的方式缺点加载时需要先定义模型结构方式二保存完整模型TorchScriptscript_model torch.jit.script(model) torch.jit.save(script_model, xxxxxxx.pth)保存内容模型结构 参数 计算图是一个独立可执行的文件。加载方式model torch.jit.load(xxxxxxx.pth) model.eval()优点无需定义模型结构直接加载即可用可以跨平台部署C、移动端等适合生产环境缺点文件较大某些复杂动态结构可能无法脚本化6.2 模型文件扩展名扩展名说明.pt/.pthPyTorch通用模型文件.t7Torch7格式旧版.onnx开放神经网络交换格式6.3 best_acc 的作用if correct best_acc: best_acc correct # 保存模型通过维护一个全局的best_acc只有当当前epoch的准确率超过历史最优时才保存。这样可以避免保存效果较差的模型确保最终保存的是训练过程中表现最好的版本。七、完整训练流程loss_fn nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) epochs 10 for t in range(epochs): print(fEpoch {t1}\n-----------------------------------) train(train_dataloader, model, loss_fn, optimizer) print(Done!) test(test_dataloader, model, loss_fn)注意test()只在训练结束后调用一次。如果希望在每个epoch后都评估并保存最优模型可以在训练循环内调用test()。八、数据增强的效果分析增强方法模拟的现实变化对模型的影响RandomRotation拍摄角度不同提升旋转不变性RandomFlip镜像拍摄提升翻转不变性ColorJitter光照条件不同提升光照鲁棒性RandomGrayscale黑白照片减少对颜色的依赖Normalize数据分布统一加速收敛提升稳定性实践建议数据增强不是越多越好要根据任务特点选择对于人脸识别垂直翻转通常不合适人脸有方向性对于食物分类旋转、翻转、颜色抖动都很合适验证集必须使用与测试集相同的变换不能加入随机性九、总结本篇博客围绕数据增强和模型保存两大主题系统讲解了知识点核心内容数据增强RandomRotation、RandomFlip、ColorJitter、RandomGrayscale标准化ToTensor Normalize使用ImageNet统计参数训练/验证变换训练集用增强验证集只用必要变换模型保存方式一torch.save(model.state_dict())保存参数模型保存方式二torch.jit.script()torch.jit.save()保存完整模型最优模型保存用best_acc追踪只保存最好的版本核心收获数据增强是提升模型泛化能力的关键手段相当于“免费”扩充数据集。标准化是深度学习训练的标准步骤不可省略。模型保存让训练成果可复用是工程落地的必要环节。两种保存方式各有优劣根据部署需求选择。
返回列表