transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)的计算过程

cifar10数据集的众多demo中,在数据加载环节,transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)这条指令是经常看到的。这是一个 PyTorch 中用于图像数据标准化的函数调用,它将图像的每个通道的值进行标准化处理,使得数据的均值变为 (0.4914, 0.4822, 0.4465),标准差变为 (0.2023, 0.1994, 0.2010)。

关于均值、均方差以及标准化函数transforms.Normalize()的文章太多了,这里记录一下计算过程。

对于 CIFAR-10 数据集,均值和标准差的计算方法如下:

1、收集数据集: 首先,你需要加载整个 CIFAR-10 数据集。CIFAR-10 数据集包含 60,000 张 32x32 的彩色图像,分为 10 个类别。

2、计算每个通道的均值: 对于每个图像,将 RGB 三个通道的值提取出来。然后对所有图像的每个通道的像素值求和,然后除以总像素数(图像数量乘以每个图像的像素数)。

**3、计算每个通道的标准差:**对于每个图像,计算每个通道的像素值与该通道均值的差的平方。再对所有图像的每个通道的平方差求和,然后除以总像素数,最后取平方根。

python 复制代码
import torch
from torchvision import datasets, transforms

# 定义数据预处理
transform = transforms.Compose([
    transforms.ToTensor()
])

# 加载CIFAR-10数据集
train_dataset = datasets.CIFAR10(root='./data', train=True, download=False, transform=transform)

# 将数据集转换为Tensor
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=1, shuffle=False)

# 初始化均值和标准差
mean = torch.zeros(3)
std = torch.zeros(3)

# 计算均值和标准差
for images, _ in train_loader:
    for i in range(3):  # 遍历RGB三个通道
        mean[i] += images[:, i, :, :].mean()   # 计算每个通道的均值
        std[i] += images[:, i, :, :].std()     # 计算每个通道的标准差

# 对三个通道的均值和标准差求平均
mean /= 3
std /= 3

# 计算平均值
mean /= len(train_loader)
std /= len(train_loader)

print(f'均值: {mean}')   # 均值: tensor([0.4914, 0.4822, 0.4465])
print(f'标准差: {std}')  # 标准差: tensor([0.2023, 0.1994, 0.2010])

上述代码稍加改造,就可用于自定义数据集的计算:

python 复制代码
import torch
from torchvision import transforms
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import os


# 自定义数据集类
class CustomDataset(Dataset):
    def __init__(self, img_dir, transform=None):
        self.img_dir = img_dir   # 图片文件夹的路径
        self.transform = transform   # 数据预处理
        self.img_files = os.listdir(img_dir)  # 图片文件列表

    def __len__(self):   # 获取数据集大小
        return len(self.img_files)

    def __getitem__(self, idx):  # 获取图片数据
        img_path = os.path.join(self.img_dir, self.img_files[idx])
        image = Image.open(img_path).convert('RGB')
        if self.transform:
            image = self.transform(image)
        return image


# 定义数据预处理
transform = transforms.Compose([
    transforms.ToTensor()
])

# 创建自定义数据集实例
custom_dataset = CustomDataset(img_dir='自定义数据集的文件夹路径', transform=transform)

# 创建数据加载器
custom_loader = DataLoader(custom_dataset, batch_size=1, shuffle=False)

# 初始化均值和标准差
mean = torch.zeros(3)
std = torch.zeros(3)

# 计算均值和标准差
for images in custom_loader:
    for i in range(3):  # 遍历RGB三个通道
        mean[i] += images[:, i, :, :].mean()  # 计算每个通道的均值
        std[i] += images[:, i, :, :].std()  # 计算每个填充的标准差

# 计算平均值
mean /= len(custom_loader)
std /= len(custom_loader)

print(f'均值: {mean}')
print(f'标准差: {std}')
相关推荐
AI人工智能+1 分钟前
基于深度学习的泰国文字识别系统,通过图像预处理、文字定位、整行序列预测、CNN特征提取、LSTM+CTC解码及语言模型校正六步,实现高精度泰文识别
深度学习·ocr·泰国文字识别
在线考试系统推荐3 分钟前
拆成多道小题,优考试“理解题“怎么用?
服务器·人工智能·学习
Elastic 中国社区官方博客4 分钟前
机构如何统一智慧城市数据以改善公共服务?
大数据·人工智能·物联网·elasticsearch·搜索引擎·全文检索·智慧城市
小白学大数据5 分钟前
API 调用错误处理实战:分层定位、有纪律地重试与可观测性设计
开发语言·网络·人工智能
Raas1006 分钟前
MAI Gateway(魔芋企业级AI网关)详解:企业为什么需要AI网关,一文读懂企业AI流量治理
大数据·人工智能·gateway·api·ai网关·mai gateway
yangdaxiageo7 分钟前
白帽GEO的结构化信任工程:杨大侠GEO商业方法论研讨
人工智能·科技·aigc·agi
Swift社区8 分钟前
Wi-Fi 6 的核心技术特性
人工智能·python
人才瘾大9 分钟前
从Agent Loop到PlanMode:一套能跑生产的AI Agent工程骨架
人工智能·ai编程
大任视点10 分钟前
《怪兽引擎:能源转型的驯化与共生》开幕 以艺术思辨锚定后化石时代的文明坐标
大数据·人工智能
硅谷秋水12 分钟前
Zero-WAM:基于人类视频的上下文世界-动作建模,用于开放式任务泛化
人工智能·深度学习·计算机视觉·语言模型·机器人·音视频