ARTICLE DETAIL

资讯详情

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

PyTorch图像归一化:深入理解transforms.Normalize的mean、std与实战

PyTorch图像归一化:深入理解transforms.Normalize的mean、std与实战 我在整理 pytorch 个人项目时几乎每个图像类的 DataLoader 里都会看到transforms.Normalize(...)这一行。刚接触的时候很容易把它当成一个“把像素缩小”的固定步骤抄完就继续写网络结构直到有一次在自定义数据集上换了自己的 mean/std又对比了不归一化的效果才意识到这个函数不只是预处理工具它直接关系到预训练模型能不能正常发挥、模型收敛速度差多少。这篇笔记把 Normalize() 从头到尾拆一遍包括 mean、std、inplace 这几个参数的真实含义出镜率最高的[0.485, 0.456, 0.406]是从哪冒出来的以及在自己数据集上怎么科学地算出合适的统计值。想真正搞懂 PyTorch 数据流、而不是停留在“照着抄 transforms”层面的读者这篇应该能帮上忙。1. Normalize() 在做的事远不止“把数字变小”1.1 从卷积核的期望说起图像输入网络之前通常会被表示成一个三维张量(C, H, W)C 是通道数H 和 W 是高宽。未经处理的像素值如果是从 PIL 或 OpenCV 读出来的原始值范围是 0 到 255如果经过transforms.ToTensor()范围会变成 0 到 1。不管是哪种数值分布都偏向正数而且不同通道之间亮度有差异比如 RGB 三个通道的均值不会完全一样。神经网络第一层的卷积核在做的事情本质上是把输入像素和权重做加权求和。这里的风险在于如果输入非常大即使权重初始化得比较小加权求和的结果也可能让激活值饱和梯度变得非常小。更麻烦的是不同通道数值基线不一样网络第一层就得花额外的力气去“适应”这种偏移。很多人以为神经网络天生能自适应数据分布理论上它能但实际训练中这种自适应能力需要大量迭代去补偿。Normalize()做的事情就是一个按通道的线性变换output (input - mean) / std也就是说先把每个通道减去均值让数据分布中心回到 0 附近再除以标准差把尺度压到差不多统一的量级。这一步不改变图像语义只改变数值分布但正是这个分布让第一层卷积的工作条件变得友好得多。1.2 不归一化时损失曲线为什么容易飘我自己做过一个不严谨但很直观的对照实验同一个 ResNet 结构同一个数据集一个走完整归一化流程一个只做ToTensor()也就是把像素限制在 0 到 1 但不做减均值除方差。结果是归一化版本在第五个 epoch 左右验证准确率就明显拉开损失下降更平滑不归一化的版本前几个 epoch 损失下降很慢有时候还会在某个小范围内反复震荡。原因不难理解。虽然 0 到 1 的范围看起来已经很小但每个通道的基线不同。比如一张偏暗的夜景图三个通道均值都很低一张阳光下拍的照片各通道均值都很高。同一个 batch 内如果混入了亮度差异极大的图像梯度方向会被少数极端样本带偏。减去均值后这些跨样本的基线差异被拉平模型注意力可以放在真正重要的纹理和结构信息上。当然现在的网络里普遍有 BatchNorm很多人会反问有 BatchNorm 还要手动归一化干嘛我的理解是BatchNorm 是在网络中间层做归一化属于“事后补偿”而输入层的 Normalize 相当于给所有后续层提供一个稳定的起点。两者解决的问题有重叠但早年训练 VGG 这类不带 BN 的模型时输入归一化几乎是必需品。即使带 BN输入归一化也能让训练早期更稳定尤其是迁移学习场景预训练权重本来就假设输入数据服从特定分布。2. mean、std、inplace 参数逐个拆开2.1 mean 和 std 的匹配方式torchvision.transforms.Normalize的签名很简洁torchvision.transforms.Normalize(mean, std, inplaceFalse)最常被忽略的是“mean 和 std 怎么匹配通道”。官方推荐的是一个和通道数等长的序列比如三通道 RGB 图像就写成transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])这三个均值分别对应 R、G、B 三个通道三个标准差同理。实现内部会把均值和标准差转成形状为(C, 1, 1)的张量然后和图像张量做广播运算也就是每个通道用自己的减去自己的均值、除以自己的标准差。如果你只有一个数值比如transforms.Normalize(0.5, 0.5)PyTorch 在较新版本里也能处理会把这个标量广播到所有通道。这种写法在灰度图、或者你想把 RGB 三个通道统一归一化到某个范围时比较省事。但我的建议是除非清楚自己在做什么否则不要用标量去糊弄三通道图像。每个通道的亮度分布天然不同用一个共同均值和标准差强行拉齐等于抹掉了通道间的色彩差异信息。通道数越多的数据比如多光谱图像越需要逐通道给出统计值。另外注意mean和std的顺序必须和一进一出的张量通道顺序完全一致。后面会专门说 OpenCV 的 BGR/RGB 坑这里先埋个伏笔。2.2 inplace 到底是省了什么inplaceFalse是默认值代表每次归一化会返回一个新的张量原来的输入数据不会被修改。把它改成 True就会直接在原张量上做减法除法不产生新的内存对象。在训练大模型时inplace 听起来很诱人每张图都少一份拷贝显存不就省下来了吗实际使用要分场景。DataLoader 里每个 batch 的张量在collate时已经新生成过经过 Normalize 再复制一份确实有点浪费打开 inplace 是安全的因为该 batch 不会在别处继续复用。但如果在某个数据增强组合里同一个张量后面还要经过其他自定义 transform而那个 transform 又依赖原始像素分布打开 inplace 就会在无意识中破坏后续步骤的输入。从我自己的习惯来说训练阶段默认不开 inplace 更稳妥。归一化本身是极轻量的线性运算真正吃显存的是中间 feature map而不是预处理阶段一次拷贝。省内存优先考虑梯度累积、混合精度、减少 batch size在 Normalize 上扣 inplace 属于性价比很低的操作。如果是做大规模数据预处理离线缓存倒是可以开能减少不少内存分配压力。3. 那张出镜率最高的 mean/std 表ImageNet 预训练统计值3.1 0.485、0.456、0.406 是怎么算出来的几乎每个 PyTorch 图像项目的 transform 里都有这套数字mean[0.485, 0.456, 0.406] std[0.229, 0.224, 0.225]它们不是拍脑袋定的也不是 PyTorch 官方发明的魔法值而是 ImageNet 数据集全量图像在 RGB 三个通道上的统计结果。ImageNet 有上百万张自然图片每一张经过尺寸处理后在 R 通道的像素均值约 0.485G 通道约 0.456B 通道约 0.406。标准差对应 0.229、0.224、0.225。观察这三个均值可以发现整体都低于 0.5也就是说自然图片里暗部占比其实偏高纯白背景的图并没有想象中那么多。而三个通道均值之间也有细微差别R 通道最高B 通道最低这其实符合日常经验天空和植被让蓝绿通道在某些图片里有偏高贡献但大量室内、肤色、土壤元素又把整体统计往偏暖方向带。我们在用单张图去感受颜色时很容易只盯着局部色彩只有站在百万张图的角度统计才能看出这种全局规律。标准差在 0.22 到 0.23 之间说明像素值围绕均值附近集中并不像很多人想象的那样均匀铺满 0 到 1。正因为像素集中在中间区域减去均值再除以标准差之后绝大多数值会落在 -1 到 1 之间这正是大多数激活函数和权重初始化最喜欢的输入范围。3.2 用了预训练模型就必须跟着用同一套统计值很多同学问我用自己的数据集能直接照抄 ImageNet 的均值方差吗答案分两种场景。如果是从 torchvision 仓库加载在 ImageNet 上预训练好的 ResNet、VGG、EfficientNet 等权重那么强烈建议照抄这套统计值。原因很朴素预训练权重是在这个统计口径下训练出来的卷积核的分布已经适应了“减 ImageNet 均值、除 ImageNet 标准差”之后的输入。你用别的 mean/std输入分布和权重期望的分布不一致再好的预训练权重也发挥不出来。迁移学习本来就指望前面的层提取通用特征分布被改掉之后通用性直接打折。如果是从零开始训练不加载任何预训练权重那 ImageNet 这套数字也可以先用因为它是从海量自然图像里统计出来的对大多数普通图像数据集有很强的代表性。但如果你的数据比较特殊比如医学影像、红外图像、灰度图、卫星多光谱图建议还是老老实实算自己的统计值。后面第 4 节会给完整计算方法。4. 自定义数据集的 mean/std 该怎么算4.1 用求和方式算全局均值和方差计算整个数据集的 mean/std最直观的想法是遍历每一张图求出每张图的通道均值和标准差然后对图像数量求平均。这个做法代码简单但严格说是有偏差的每张图内部像素数量不同时简单平均会低估或高估某些图的贡献而先按图平均标准差再求平均得到的并不是全局标准差。更严谨的做法是把所有图像的像素值全部汇总累计每个通道的像素和、像素平方和最后一次性算均值与方差。我这里给一个可以在 DataLoader 上运行的 PyTorch 实现def compute_mean_std(loader): channel_sum torch.zeros(3) channel_sum_sq torch.zeros(3) pixel_count 0 for images, _ in loader: # images 形状: (B, C, H, W)来自 ToTensor范围 [0, 1] b, c, h, w images.shape im images.permute(0, 2, 3, 1).reshape(-1, c).float() channel_sum im.sum(dim0) channel_sum_sq (im ** 2).sum(dim0) pixel_count im.shape[0] mean channel_sum / pixel_count variance channel_sum_sq / pixel_count - mean ** 2 std torch.sqrt(variance.clamp(min1e-8)) return mean, std注意代码里用的是pixel_count im.shape[0]这里im已经 reshape 成了(H*W*B, C)所以im.shape[0]是单个通道的像素总数。算方差的公式来自Var(X) E[X^2] - E[X]^2最后用clamp(min1e-8)防止浮点误差导致根号里出现极小负数。这个计算要在图像尺寸已经确定、并且ToTensor()之后做因为不同尺寸的图像像素数不一样会影响累计权重的合理性。实际项目里通常先把所有图 Resize 到统一尺寸再统计这样每个样本贡献的像素数一致结果更符合“全局”语义。4.2 统计阶段要不要带数据增强这是另一个容易纠结的问题。数据增强里的随机裁剪、翻转、旋转会改变图像分布比如 RandomCrop 可能裁掉天空部分让整体均值变高色彩抖动更是直接改变像素数值。如果带着这些随机增强去统计 mean/std得到的结果会带进增强带来的随机扰动反而失真。我的做法是统计 mean/std 时只用基础预处理也就是 Resize、CenterCrop、ToTensor关闭所有随机增强。统计集可以选训练集的一部分通常几千张图已经足够稳定不需要遍历全部数据。验证集的分布和训练集整体接近用它来算也没有原则性错误但为了不造成任何数据泄漏的疑虑我一般只用训练集。算出结果后把这个 mean/std 写死在 transform 里mean [0.4921, 0.4672, 0.4153] # 举例 std [0.2245, 0.2198, 0.2217] transform_train transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(meanmean, stdstd), ])如果你的数据集会持续更新统计值也要定期重新算不要一套统计值用半年。数据分布漂移了mean/std 应该跟着变。5. Compose 管线与 ToTensor 的配合顺序以及 BGR/RGB 的坑5.1 标准管线里 Normalize 的位置不能乱torchvision 的 transform 是按写进 Compose 的顺序从左到右执行的。一条典型的训练管线是transform transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])Normalize 被放在 ToTensor 之后是必须的。ToTensor干两件事把 PIL Image 或 numpy 数组转成 torch.Tensor同时把 0 到 255 的整型值缩放到 0 到 1 的浮点。Normalize 实现里对输入张量有强制要求必须是浮点型张量否则直接抛异常。如果你把顺序写反先 Normalize 后 ToTensor大概率会看到类似TypeError的报错更隐蔽的是有些代码能运行但归一化的数值完全对不上。RandomResizedCrop、HorizontalFlip 这类几何增强必须放在 ToTensor 之前因为 torchvision 的几何增强接口对 PIL 图像支持最完备对张量的处理函数反而是另一个命名空间。ColorJitter 也一样对 PIL 图做颜色增强更自然转成张量之后再去改颜色就得小心数值范围。有个细节容易被忽略normalize 之后的张量已经不在 0 到 1 范围了负值完全正常。如果后面还有自定义 transform 假设输入是 0 到 1会出问题。自定义 transform 最好统一约定一个契约进入 Normalize 之前是 0 到 1之后是近似 -2 到 2 的标准化分布后面所有模块按这个标准去写。5.2 用 OpenCV 读图、albumentations 增强时的等价写法用 OpenCV 读图得到的是 HWC 排列的 numpy 数组而且通道顺序是 BGR。如果你直接把这张图转成 Tensor 再 Normalizemean 和 std 的通道对应关系就全错了蓝色通道用了 R 的均值红色通道用了 B 的均值。训练时模型能跑但收敛速度和最终精度都会受影响而且这种错误极其隐蔽不仔细查代码根本发现不了。正确做法是先换通道顺序import cv2 import torch img cv2.imread(sample.jpg) # HWC, BGR img_rgb cv2.cvtColor(img, cv2.COLOR_BGR2RGB) tensor torch.from_numpy(img_rgb).permute(2, 0, 1).float() / 255.0 tensor transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])(tensor)如果你在用 albumentations 做增强它的Normalize参数和 torchvision 类似但默认输入是 numpy 数组而且可以选max_pixel_value255内部自己完成缩放。等价写法是import albumentations as A transform A.Compose([ A.Resize(224, 224), A.HorizontalFlip(), A.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225], max_pixel_value255.0), ])albumentations 的 Normalize 输出的已经是 float32 numpy之后你还需要手动torch.from_numpy(transformed[image])转成 Tensor再处理通道顺序。正是因为不同增强库对输入输出格式的约定不一样预处理这一层的通道顺序、数值范围、dtype 三件事必须有一条清晰的流水线不能靠感觉。6. 高频报错、逆归一化与两个进阶用法6.1 我遇到过的几个高频故障这里把常见的坑集中整理成一张表方便排查。现象可能原因修复方式报错TypeError: t is not a torch Tensor把 Normalize 放在 ToTensor 之前调整 Compose 顺序报错RuntimeError: The size of tensor a (...) must match the size of tensor b (...)mean/std 数量与通道数不匹配比如灰度图用了三通道统计值检查通道数灰度图用长度为 1 的统计值报错TypeError: Expected tensor to be floating point输入张量还是 int64 或 uint8手动.float()或先经过 ToTensor训练 Loss 不下降验证集波动极大通道顺序搞错OpenCV 的 BGR 没转 RGB读图后先cvtColor再进管线迁移学习效果和随机初始化差不多用了自定义 mean/std 但加载了 ImageNet 预训练权重统一数据统计口径换回 ImageNet 标准值还有一个不算报错但影响结果的问题推理阶段忘记写model.eval()。这跟 Normalize 本身无关但 BatchNorm 依赖的 running stats 在训练和推理模式下行为不同。很多人在验证集上指标奇怪费半天劲查预处理最后发现是model.eval()没调用。建议把预处理、模型模式、反向传播开关这几件事放在同一个检查清单里排错效率高很多。6.2 逆归一化把标准化后的张量还原成可看图片训练过程中如果想可视化中间结果比如把网络输入或生成图片保存下来就需要逆归一化。正常归一化是(x - mean) / std逆回去就是x * std mean。写成一个通用函数def denormalize(tensor, mean, std): mean torch.as_tensor(mean).view(-1, 1, 1) std torch.as_tensor(std).view(-1, 1, 1) tensor tensor * std mean tensor tensor.clamp(0, 1) return tensor这里clamp(0, 1)很重要。由于网络输出或中间张量可能存在超出归一化范围的值逆运算后的像素可能小于 0 或大于 1保存图片前需要裁剪回合法区间。最后用torchvision.utils.save_image或transforms.ToPILImage()时记得先把张量放回 CPU 并把通道维调整好。6.3 把 Normalize 当作“归一到任意区间”的工具很多 GAN 相关项目里会遇到另一种写法transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))代入公式算一下(x - 0.5) / 0.5 2x - 1原本 0 到 1 的输入会被映射到 -1 到 1。这不是什么特殊魔法只是把 mean 设为区间的中点、std 设为区间的一半。GAN 的生成器最后一层常用 Tanh 激活输出范围是 -1 到 1输入数据也归一到同一个范围训练会更稳定。因此 Normalize 不只是给预训练模型用的它本身就是一个灵活的线性映射工具你想把数据落到哪个区间反推对应 mean/std 就行。这类用法再次印证了一个核心观点Normalize 的参数不是随便填的每一个数字背后都应该有明确的分布假设。无论是照搬 ImageNet 统计值让自己的数据和预训练权重对齐还是自己算数据集统计值还是为了配合激活函数输出范围而构造特定映射你要做的是理解这行 transform 在数值上到底做了什么而不是把它当作“每次都要写的模板代码”。把这一点想清楚以后在各类图像任务里遇到预处理问题排查思路会清晰很多。
返回列表