ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch Transforms全解析:从数据预处理到图像增强实战

PyTorch Transforms全解析:从数据预处理到图像增强实战 系列第三篇写Transforms我酝酿了挺久。前面两篇咱们把Pytorch环境搭好了、张量基础也过了一遍但很多人走到这一步会卡住明明Dataset已经把图片读出来了怎么喂给模型就报错或者训练了半天、loss曲线像心电图一样问题往往不在网络结构而在数据预处理这一步。Pytorch里Transforms就是干这个的把乱七八糟的原始数据洗干净、切好、摆盘变成模型能直接下嘴的样子。这篇我用大量可以粘贴运行的代码把Transforms从原理到实战一次讲透新手照着敲就行。1. Transforms到底在解决什么问题1.1 模型不能直接吃原始图片先说一个很多入门教程没讲清楚的点你从网上随便下载的图片、用cv2或PIL读进来的图片和模型真正需要的数据格式之间至少有四层差距。第一是类型。PIL读进来是PIL.Image.Image对象OpenCV读进来是numpy.ndarray但Pytorch的模型只认torch.Tensor。第二是维度顺序。图片在numpy里通常是[H, W, C]高、宽、通道RGB三通道在最后一维可Pytorch的卷积层要求[C, H, W]通道维必须放在最前面。第三是数值范围。普通图片的像素值是0到255的整数而神经网络内部用的是浮点数计算梯度计算对数值范围很敏感大范围输入很容易让loss震荡甚至发散。第四是数据分布。如果每张图片的亮度、对比度天差地别模型要花很大精力去适应这种差异训练速度会明显变慢。Transforms就是专门在这层差距上搭桥的工具。它是一系列可组合的数据处理操作可以逐个对样本做处理也可以通过Compose把它们串成一条流水线。我习惯把这个过程类比成做饭Dataset负责把食材从菜市场买回来读原始数据Transforms负责洗菜、切菜、配菜格式转换、去噪、增强模型就是大厨只负责炒菜你给他一块没洗的带泥萝卜他再厉害也没法下锅。1.2 训练集和测试集要用不同的Transforms这是新手最容易忽略的一个问题。很多人写了transform transforms.Compose([...])之后训练和测试全用同一套结果就是模型在训练集上效果不错一换到新图片上立刻崩盘。原因在于训练集需要数据增强。所谓增强就是给原始图片做随机的翻转、裁剪、调色人为制造出“看起来差不多但细节不同”的样本。这样模型学到的就不是某一张图的固定特征而是一类图的内在规律。相当于学生备考要做大量变换花样的模拟题而不是把某一道题的答案背下来。测试集则恰恰相反必须用固定的处理方式——只做ToTensor和Normalize不做任何随机增强才能公平地衡量模型在真实数据上的表现。所以标准写法是定义两个Composefrom torchvision import transforms transform_train transforms.Compose([ transforms.RandomHorizontalFlip(p0.5), transforms.RandomCrop(32, padding4), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])训练时用transform_train验证、测试时用transform_test。我见过很多人图省事测试集也加了RandomHorizontalFlip结果同一张图每次预测结果都不一样最后排查半天才发现是这个原因。这一点请务必记住。2. 逐个拆解最常用的Transforms2.1 拼装管道Compose的用法和顺序transforms.Compose是最核心的组织方式。它的用法很简单传入一个列表列表里的操作按照顺序依次执行。但顺序里面有讲究我踩过坑之后总结了几条原则。第一条所有需要PIL Image格式的操作比如Resize、RandomCrop、RandomHorizontalFlip要放在ToTensor之前。原因是ToTensor会把PIL Image转成Tensor而很多几何变换操作老版本只支持PIL输入传给Tensor直接报错。虽然新版torchvision对Tensor的支持越来越好了但为了兼容性我仍然习惯把几何变换放在前面。第二条Normalize必须放在ToTensor之后。因为ToTensor会把像素值缩放到0到1之间Normalize按通道减均值除标准差才能得到合理的结果。你要是把顺序倒过来对着0到255的整数做归一化那计算出来的东西完全是乱的。第三条Compose里每个操作的对象必须匹配。你可以在一个Compose里混用不同类别的操作但要保证上一个操作的输出类型是下一个操作能接受的输入类型。transform transforms.Compose([ transforms.Resize((224, 224)), # PIL Image - PIL Image transforms.RandomHorizontalFlip(), # PIL Image - PIL Image transforms.ToTensor(), # PIL Image - Tensor [0,1] transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), # Tensor - Tensor ])2.2 最核心的两个ToTensor和NormalizeToTensor做三件事把数据转成torch.Tensor把维度从HWC调成CHW把像素值从0到255缩放到0到1。这三件事一步到位是任何图像任务的必经之路。Normalize是我们的老朋友了公式很简单output (input - mean) / std。它的作用是把数据分布拉到一个均值接近0、标准差接近1的状态。这样做的好处是让不同维度的特征处在相似的数量级上梯度更新更稳定模型收敛更快。很多新手问Normalize的参数为什么是0.485, 0.456, 0.406和0.229, 0.224, 0.225这组数字是ImageNet数据集的RGB三通道均值和标准差。如果做迁移学习用的预训练模型是在ImageNet上训的那你就必须用这组参数因为预训练模型“习惯”这种数据分布。如果你自己从零训练一个模型严格来说应该统计自己数据集的均值和标准差。import torch import numpy as np from PIL import Image img Image.open(demo.jpg).convert(RGB) print(type(img)) # class PIL.JpegImagePlugin.JpegImageFile print(img.size) # 宽和高比如 (640, 480) tensor transforms.ToTensor()(img) print(tensor.shape) # torch.Size([3, 480, 640]) print(tensor.min().item(), tensor.max().item()) # 0.0, 1.02.3 几何变换Resize、Crop、Flip、Rotate这一组操作主要解决图片尺寸不统一的问题同时通过随机几何变化做数据增强。Resize就是把图片缩放到指定尺寸注意参数是(w, h)还是(h, w)torchvision里Resize接受的是(h, w)别传反了。RandomCrop在图片上随机裁一块配合padding使用可以先填充再裁剪等于在小范围内制造平移效果。RandomHorizontalFlip按概率p水平翻转一张图p默认0.5。RandomRotation随机旋转一定角度配合expand参数可以控制是否放大画布防止边缘被切掉。对于小数据集这些几何变换几乎是免费的增广手段。我做过一个实验同样一个小分类模型加上RandomHorizontalFlip和RandomCrop之后验证集准确率从72%提到了81%白捡了将近10个点这就是数据增强的威力。2.4 颜色增强与高级组合ColorJitter可以随机调整亮度、对比度、饱和度、色调参数分别是brightness、contrast、saturation、hue。比如ColorJitter(brightness0.2, contrast0.2)就表示在0.8到1.2倍之间随机调整亮度。RandomResizedCrop很有意思先随机选一块区域裁剪下来再缩放成指定尺寸。这比单纯的RandomCrop更激进模拟了“从不同距离、不同角度观察物体”的效果是目前分类任务里最常用的增强方式之一。Lambda则允许你写自定义函数灵活度最大。比如做一个随机擦除transform transforms.Compose([ transforms.ToTensor(), transforms.Lambda(lambda x: x * torch.rand_like(x).gt(0.9).float()), ])这段代码会随机把大约10%的像素直接置0效果类似RandomErasing逼迫模型去学更鲁棒的特征。当然实战中直接用torchvision.transforms.v2或RandomErasing更省事但Lambda能让你理解Transforms底层就是个“函数套函数”的机制。3. 实战一个完整可运行的图像分类流程3.1 CIFAR-10加载与Transforms组合理论说了这么多不跑代码都是纸上谈兵。我们用CIFAR-10这个经典数据集走一遍完整流程。第一步定义训练和测试的Transforms并加载数据。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms transform_train transforms.Compose([ transforms.RandomHorizontalFlip(p0.5), transforms.RandomCrop(32, padding4), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616]), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616]), ]) train_set datasets.CIFAR10( root./data, trainTrue, downloadTrue, transformtransform_train ) test_set datasets.CIFAR10( root./data, trainFalse, downloadTrue, transformtransform_test ) train_loader DataLoader(train_set, batch_size64, shuffleTrue, num_workers2) test_loader DataLoader(test_set, batch_size64, shuffleFalse, num_workers2)注意这组Normalize参数是CIFAR-10自己的统计值不是ImageNet那组。CIFAR-10图片是32x32的小图和ImageNet那种大图分布不一样如果照搬ImageNet的参数效果会差一点。3.2 用可视化脚本验证Transforms效果数据加载写完后我强烈建议你先别急着写模型先做一个可视化检查。这一步能帮你确认Transforms到底把图片变成什么样了有没有裁坏有没有颜色失真import matplotlib.pyplot as plt import numpy as np def imshow_tensor(img_tensor, mean(0.4914, 0.4822, 0.4465), std(0.2470, 0.2435, 0.2616)): img img_tensor.clone() img img.permute(1, 2, 0).numpy() # CHW - HWC img img * np.array(std) np.array(mean) # 反归一化 img np.clip(img, 0, 1) plt.imshow(img) plt.show() dataiter iter(train_loader) images, labels next(dataiter) # 打印batch中第一张图的信息 print(images.shape) # torch.Size([4, 3, 32, 32]) print(images[0].shape) # torch.Size([3, 32, 32]) print(images[0].mean(), images[0].std()) imshow_tensor(images[0])这个脚本里有几个关键点要解释。images打印出来是torch.Size([4, 3, 32, 32])前面那个4是batch size后面才是通道和高宽。因为做了Normalizetensor里的值有正有负、均值也不为0直接去显示会看到一片灰蒙蒙的图所以需要先把均值加回来、标准差乘回来这叫反归一化最后clip到0到1之间。我在实际排查中多次通过这个可视化脚本发现了问题。有一次写自定义Dataset时__getitem__里忘了写self.transform(image)结果模型训练时维度对不上有一次ColorJitter参数设太大图片色彩完全漂移看起来像老化的胶卷还有一次RandomCrop把目标物体裁掉了一半导致模型学了一堆背景噪声。这些光看数字是发现不了的一定要把图打出来看。3.3 训练循环让Transforms真正跑起来接下来我们定义一个极简的小模型跑几个epoch把整个链路打通。class SimpleCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 8 * 8, 128), nn.ReLU(), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x)) model SimpleCNN() criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) for epoch in range(3): model.train() running_loss 0.0 for i, (inputs, labels) in enumerate(train_loader): optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() # 验证阶段使用transform_test的数据 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() avg_loss running_loss / len(train_loader) acc 100.0 * correct / total print(fEpoch {epoch1}: loss{avg_loss:.4f}, accuracy{acc:.2f}%)这个网络很小跑3个epoch大概一两分钟就完事准确率能到60%左右就说明整个链路没问题。这里有几个容易踩的坑值得说一下。model.train()和model.eval()一定要成对出现训练模式下BatchNorm会更新统计量并且Dropout会生效切换到测试模式后这些行为要关掉不然用torch.no_grad()算出来的结果也是不准的。torch.no_grad()能省显存和加速推理因为推理时不需要保存梯度。3.4 自定义Dataset时如何接Transforms很多人的实际任务不是CIFAR-10而是自己文件夹里的图片。这种情况下你必须自己写Dataset并且要记得预留transform参数。我见过不少新手在自定义Dataset里把transform写死导致训练和测试想用不同预处理时只能复制粘贴类非常痛苦。import os from PIL import Image from torch.utils.data import Dataset class MyImageDataset(Dataset): def __init__(self, image_dir, labels, transformNone): self.image_paths image_dir self.labels labels self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image Image.open(self.image_paths[idx]).convert(RGB) label self.labels[idx] if self.transform: image self.transform(image) return image, label然后这样用train_dataset MyImageDataset(train_img_paths, train_labels, transformtransform_train) test_dataset MyImageDataset(test_img_paths, test_labels, transformtransform_test)这里有一个细节Image.open默认不会把图片转成RGB如果是带透明通道的PNG图读进来是四通道直接喂给Resize和ToTensor没问题但后面卷积层输入通道数就得改成4。所以在你用Image.open之后养成加.convert(RGB)的习惯把通道统一成3。4. 高频报错与排查实录4.1 ImportError: cannot import name transforms from albumentations.augmentations这个报错的热度近年非常高很多人都栽在这里。常见的情况是你写了这种代码# 错误写法 from albumentations.augmentations import transforms或者有的人项目里同时装了torchvision和albumentations然后混淆了两个库的导入路径。albumentations是一个很优秀的第三方数据增强库但它的内部结构是albumentations.augmentations下面直接放各种增强类的根本没有一个叫transforms的子模块。正确的写法是from torchvision import transforms # 如果你想用Pytorch官方的那套 import albumentations as A # 如果你想用albumentations那套顺便说一句如果用albumentations它返回的结果需要你自己转成numpy格式通常不能直接和Dataset返回协议无缝衔接需要一点额外封装这里不展开讲了。建议零基础的读者先老老实实用torchvision.transforms等熟悉了数据流再考虑切换。4.2 RuntimeError: pic should be PIL Image or ndarray这个报错出现在调用ToTensor()时。原因很简单ToTensor只接受PIL.Image或numpy.ndarray作为输入如果你传进去一个torch.Tensor它就直接报错。什么情况下会传成Tensor最常见的是你在Compose里写了两次transforms.ToTensor()第一次已经把PIL Image转成了Tensor第二次再执行就炸了。另一个常见场景是自定义Dataset的__getitem__里自己手动转了一次Tensor外面又套了含有ToTensor的Compose。检查方法很简单在__getitem__里打印一下image.type()看看到底是什么类型再决定要不要继续转。4.3 图片显示出来全黑或全白该怎么排查这个我在3.2节提过单独拎出来再说一次是因为真的太常见了。全黑通常是因为你直接显示了Normalize之后的tensor。Normalize之后很多像素值是负数matplotlib遇到负数要么显示成黑色要么报值域警告。全白则是另一个极端常见于忘了做ToTensor还想直接显示没缩放的uint8数值。排查步骤我总结成口诀先permute再反归一化最后clip。任何一张Tensor图片要显示先permute(1, 2, 0)把通道调回来然后如果是归一化过的数据就乘std加mean最后np.clip到0到1。记住这三步图片显示就不会有玄学问题。4.4 维度对不上和诡异的shape新手的另一个高频事故是模型输入维度不对。我见过有人把[64, 3, 224, 224]这种batch数据误当成单张图处理或者忘了CIFAR-10的图本身就是32x32就不再Resize结果模型里Linear层输入维度算错了报mat1 and mat2 shapes cannot be multiplied。遇到维度报错第一反应不是改代码而是打印每一层前后的shape。养成在小模型里随手加注释的习惯x self.conv1(x) # [B, 3, 32, 32] - [B, 32, 16, 16] x self.conv2(x) # [B, 32, 16, 16] - [B, 64, 8, 8] x x.view(x.size(0), -1) # [B, 64, 8, 8] - [B, 4096]Transforms的维度只是数据流的一部分但从这里排查往往最快。5. 一些容易踩但没人提醒的细节写到这里我再补充几个贯穿整个Transforms使用过程中的小细节都是我在项目和带新人的过程中反复验证过的经验。关于Resize和RandomCrop的参数Resize接收的是(height, width)不是(width, height)。很多人写习惯了OpenCV的(w, h)参数顺序到了torchvision里直接照搬结果所有训练图片都被拉伸变形了还浑然不知。我自己为了这个事专门在代码里加了命名参数写成transforms.Resize(size(224, 224))防止哪天脑子抽风传反。关于随机种子数据增强里用到的随机操作如果不在训练前固定随机种子每次跑结果都不一样。这在调参时很讨厌——你以为这次改了个参数效果好其实是随机种子变了而已。在训练脚本开头加上import random import numpy as np import torch seed 42 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed)如果用了CUDA最好再加上torch.cuda.manual_seed_all(seed)。这样相同代码相同参数一定能复现结果。关于num_workersDataLoader里的num_workers不要调太大。在Windows上用多进程加载数据经常报错在Linux上设成2到4就够了设成8反而可能因为频繁切换进程而变慢。如果你在调试阶段建议直接设成0此时数据在主进程里加载能看到更多报错细节。关于性能Transforms是逐样本处理的处理大量数据时可能成为瓶颈。2023年之后的torchvision推出了transforms.v2可以处理批量数据、支持GPU加速性能比v1提升很多。初学者不需要立刻学v2但知道有这个东西以后遇到大数据量时能有个方向。关于和Albumentations冲突的问题如果你的项目里既想用torchvision的ToTensor和Normalize又想用Albumentations做增强我见过不少人直接写from albumentations import transforms然后把torchvision.transforms也引进来结果命名空间互相覆盖怎么报错都不知道。建议给它们起别名from torchvision import transforms as tv_transforms。或者在同一个工程文件里只用其中一个不混合使用能避免大量烦恼。关于Normalize用的是global mean还是per-image mean很多人会问Normalize里那个mean到底该用所有训练图片算出来的全局mean还是每张图片自己的mean答案是全局mean只要统计一次整个训练集的均值即可。但注意要按通道统计RGB三通道各有一个mean和一个std。如果你想偷懒直接用ImageNet的话跑出来的效果也不会太差很多框架默认就是这组参数。关于ToTensor和Normalize对灰度图的影响如果你的任务是灰度图ToTensor会把形状变成[1, H, W]Normalize的mean和std的列表长度也要改成1比如[0.5], [0.5]。很多人用RGB的3元素列表去归一化灰度图直接报元素个数不对的错。这个问题不复杂就是特别容易在深夜调试时把人搞疯。关于推理阶段的一致性训练时如果你用了RandomHorizontalFlip、RandomCrop这类随机增强推理时必须关掉。测试集的Compose里千万不能放随机操作。逻辑很简单测试时模型要对同一张图输出同一个结果来了两张内容完全相同的图片一张运气好翻转了一张没翻转预测结果就可能不一样这在业务上是完全不可接受的。关于数据流里出现list的问题如果自定义Dataset里返回的不是Tensor而是list或者Python原生数值DataLoader会自动帮忙转Tensor但有时会出错。保险的做法是在__getitem__里把所有要返回的数据都转成明确的类型label直接返回整数也可以。返回list的字段在batch时会变成list of list喂给损失函数时可能类型不匹配不如一开始就规范化。关于打印Transforms对象看配置我习惯在写新脚本时直接打印Compose内容确认model_transform transforms.Compose([...]) print(model_transform)它会输出每一步的参数有助于快速检查有没有漏掉ToTensor或者Normalize参数是不是写反了。这个小习惯能省很多排查时间。6. 写在最后的个人体会我个人做图像任务的固定流程现在回头看其实相当简单先定制好训练集和测试集的两套Transforms把Compose打印出来确认然后写一个可视化脚本看几张图。这三步做完数据这块基本就稳了后面模型怎么改都不用回头再折腾数据。很多人在项目里花大量时间调模型结构、调学习率但最后发现瓶颈恰恰是最初没做数据预处理。Transforms本身并不难难的是建立“数据也是需要设计”的意识。我见过不少从调参误区里绕出来的朋友转头仔细设计数据增强之后模型效果直接提了一个档次。这也是为什么我坚持在这篇入门教程里花大篇幅讲数据和增强——你后面跑GAN、跑检测、跑分割不管什么任务这套预处理逻辑都是通用的。希望这篇能帮你把Pytorch的数据流打通少踩几个我当年踩过的坑。
返回列表