You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 16:18:02