PyTorch中图像随机裁剪场景下如何正确对张量执行归一化
CIFAR-10带随机增强的张量归一化实现方案
首选方案:使用原始数据集的官方统计量
常规图像分类训练流程中,归一化所用的通道均值和标准差不需要匹配随机增强后的结果,因为随机裁剪、水平翻转属于无偏变换,不会改变数据集整体的通道分布,直接使用CIFAR-10官方提供的预计算统计量即可:
from torchvision import transforms, datasets trafo = transforms.Compose([ transforms.Pad(padding = 4, fill = 0, padding_mode = "constant"), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomCrop(size = (32, 32)), transforms.ToTensor(), # 使用CIFAR-10预计算的三通道均值和标准差 transforms.Normalize(mean = (0.4914, 0.4822, 0.4465), std = (0.2023, 0.1994, 0.2010)) ]) cifar10_full = datasets.CIFAR10(root = "CIFAR-10", train = True, transform = trafo, download = True)
该方案是业界通用做法,不需要额外计算统计量,性能稳定。
自定义方案:计算当前增强策略下的统计量再归一化
如果确实需要适配你当前增强策略的专属统计量,可以按以下两步实现,不需要固定随机种子:
第一步:遍历增强后的数据集计算统计量
随机增强带来的统计量波动极小,训练过程对该级别的波动完全鲁棒,无需固定种子对齐结果。
import torch mean = torch.zeros(3) std = torch.zeros(3) sample_count = len(cifar10_full) for img, _ in cifar10_full: # 单张图像维度为 (3, 32, 32),对高、宽维度求均值和标准差 mean += img.mean(dim=(1, 2)) std += img.std(dim=(1, 2)) # 得到全数据集的平均通道均值与标准差 mean /= sample_count std /= sample_count
第二步:应用归一化
两种实现方式可选:
方式1:重新构建数据集拼接归一化变换
trafo_final = transforms.Compose([ transforms.Pad(padding = 4, fill = 0, padding_mode = "constant"), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomCrop(size = (32, 32)), transforms.ToTensor(), transforms.Normalize(mean=mean, std=std) ]) cifar10_full_final = datasets.CIFAR10(root = "CIFAR-10", train = True, transform = trafo_final, download = True)
方式2:包装现有数据集无需重复加载
自定义数据集包装类直接对已加载的cifar10_full做归一化:
from torch.utils.data import Dataset class NormalizedCIFAR10(Dataset): def __init__(self, origin_dataset, mean, std): self.origin = origin_dataset # 调整均值、标准差维度适配张量广播规则 self.mean = torch.as_tensor(mean).view(3, 1, 1) self.std = torch.as_tensor(std).view(3, 1, 1) def __len__(self): return len(self.origin) def __getitem__(self, idx): img, label = self.origin[idx] return (img - self.mean) / self.std, label normalized_cifar = NormalizedCIFAR10(cifar10_full, mean, std)
内容的提问来源于stack exchange,提问作者Imahn
相关产品推荐
相关产品推荐

