深度学习框架PyTorch笔记(四)数据转换 Data Transformation

在深度学习的实际工程中,数据预处理往往是决定模型性能的关键一步。原始数据通常以各种格式存在——图片可能是不同分辨率、不同通道顺序;文本可能是不同长度;数值特征可能存在量纲差异。PyTorch 提供了强大且灵活的 torchvision.transforms 模块,让我们能够以组合的方式对数据进行标准化、增强、裁剪、翻转等操作。本文将从实战角度出发,通过大量代码演示,带你掌握 PyTorch 中的数据转换核心技巧。## 数据转换的基本概念在 PyTorch 中,数据转换(Transformation)本质上是一个可调用对象,它接收一个输入数据(如 PIL 图像、Tensor 或 NumPy 数组),并返回转换后的输出。torchvision.transforms 模块提供了许多预定义的转换类,例如:- ToTensor:将 PIL 图像或 NumPy 数组转换为 PyTorch Tensor,并将像素值从 [0,255] 缩放到 [0.0,1.0]- Normalize:对 Tensor 进行标准化处理,即减去均值再除以标准差- Resize:调整图像尺寸- RandomCrop:随机裁剪图像- RandomHorizontalFlip:随机水平翻转这些转换可以通过 transforms.Compose 组合成一个管道,按顺序依次执行。## 实战代码示例1:图像分类中的数据增强在图像分类任务中,数据增强是防止过拟合、提升模型泛化能力的常用手段。下面我们演示如何对 CIFAR-10 数据集进行数据转换。pythonimport torchimport torchvisionimport torchvision.transforms as transformsfrom torch.utils.data import DataLoaderimport matplotlib.pyplot as pltimport numpy as np# 定义数据转换管道transform_train = transforms.Compose([ transforms.Resize(32), # 调整尺寸为32x32 transforms.RandomCrop(32, padding=4), # 随机裁剪,带4像素填充 transforms.RandomHorizontalFlip(p=0.5), # 随机水平翻转,概率0.5 transforms.ColorJitter(brightness=0.2, # 随机调整亮度、对比度、饱和度 contrast=0.2, saturation=0.2), transforms.ToTensor(), # 转换为Tensor并缩放到[0,1] transforms.Normalize((0.4914, 0.4822, 0.4465), # 标准化:减去均值 (0.2023, 0.1994, 0.2010)) # 除以标准差])# 测试集只做基本转换transform_test = transforms.Compose([ transforms.Resize(32), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))])# 加载CIFAR-10训练集和测试集trainset = torchvision.datasets.CIFAR10( root='./data', train=True, download=True, transform=transform_train)testset = torchvision.datasets.CIFAR10( root='./data', train=False, download=True, transform=transform_test)# 创建DataLoadertrainloader = DataLoader(trainset, batch_size=64, shuffle=True, num_workers=2)testloader = DataLoader(testset, batch_size=64, shuffle=False, num_workers=2)# 可视化一个batch的增强后图像def imshow(img): img = img / 2 + 0.5 # 反标准化,将图像恢复到[0,1]范围 npimg = img.numpy() plt.imshow(np.transpose(npimg, (1, 2, 0))) plt.axis('off')# 获取一个batchdataiter = iter(trainloader)images, labels = next(dataiter)# 显示图像plt.figure(figsize=(12, 8))imshow(torchvision.utils.make_grid(images[:8]))plt.title('增强后的CIFAR-10样本')plt.show()代码解读:- 训练集使用了 RandomCropRandomHorizontalFlipColorJitter 三种增强手段,有效增加了数据多样性。- Normalize 使用的均值和标准差是 CIFAR-10 数据集预先计算好的(RGB 三个通道)。- 测试集只做必要的尺寸调整和标准化,不做随机增强,以保证评估结果的稳定性。- 可视化时通过 img / 2 + 0.5 反标准化,将图像从 [-1,1] 恢复到 [0,1] 显示。## 实战代码示例2:自定义数据转换与 Lambda当内置的转换不能满足需求时,我们可以使用 transforms.Lambda 来封装自定义函数,或者直接编写继承自 torch.nn.Module 的自定义转换类。下面演示如何将图像转换为灰度图并添加椒盐噪声。pythonimport torchimport torchvision.transforms as transformsfrom PIL import Imageimport numpy as npimport random# 自定义转换:添加椒盐噪声class AddSaltPepperNoise(object): """以概率prob添加椒盐噪声""" def __init__(self, prob=0.02): self.prob = prob def __call__(self, img): # img为PIL图像,先转换为numpy数组 img_np = np.array(img) if len(img_np.shape) == 2: # 灰度图 height, width = img_np.shape else: # RGB图像 height, width, channels = img_np.shape # 生成噪声掩码 salt_mask = np.random.random((height, width)) < self.prob / 2 pepper_mask = np.random.random((height, width)) < self.prob / 2 # 应用椒盐噪声 if len(img_np.shape) == 2: img_np[salt_mask] = 255 img_np[pepper_mask] = 0 else: img_np[salt_mask] = [255, 255, 255] img_np[pepper_mask] = [0, 0, 0] return Image.fromarray(img_np)# 使用Lambda实现自定义转换:将图像转换为灰度图to_grayscale = transforms.Lambda(lambda img: img.convert('L'))# 构建转换管道custom_transform = transforms.Compose([ to_grayscale, # 先转成灰度图 AddSaltPepperNoise(prob=0.03), # 添加3%的椒盐噪声 transforms.Resize((128, 128)), # 调整尺寸 transforms.ToTensor(), # 转为Tensor transforms.Normalize(mean=[0.5], std=[0.5]) # 灰度图只有一个通道])# 测试自定义转换def test_custom_transform(): # 创建一个随机的彩色图像(模拟真实数据) fake_image = Image.fromarray( np.random.randint(0, 256, (256, 256, 3), dtype=np.uint8) ) print(f"原始图像尺寸: {fake_image.size}, 模式: {fake_image.mode}") # 应用转换 transformed = custom_transform(fake_image) print(f"转换后Tensor形状: {transformed.shape}") # 应为 [1, 128, 128] print(f"像素值范围: [{transformed.min().item():.4f}, {transformed.max().item():.4f}]") # 验证是否确实是灰度图 assert transformed.shape[0] == 1, "应该是单通道灰度图" assert transformed.shape[1] == 128 and transformed.shape[2] == 128, "尺寸应为128x128" print("自定义转换验证通过!")test_custom_transform()代码解读:- AddSaltPepperNoise 类实现了可调用接口 __call__,内部使用 numpy 生成随机噪声掩码。- transforms.Lambda 允许我们传入任意函数,这里用匿名函数将彩色图转为灰度图。- 自定义转换必须谨慎处理数据类型转换:PIL 图像 ↔ NumPy 数组 ↔ Tensor 之间的兼容性。- 在 Normalize 时,灰度图只需提供单通道的均值和标准差。## 进阶技巧:混合使用 torchvision 与 albumentations在实际工程中,很多团队会结合使用 torchvision.transformsalbumentations 库,后者在图像增强方面性能更优、功能更丰富。下面是一个混合使用的示例:pythonimport albumentations as Afrom albumentations.pytorch import ToTensorV2# 使用albumentations定义增强albumentations_transform = A.Compose([ A.RandomResizedCrop(height=224, width=224, scale=(0.8, 1.0)), A.HorizontalFlip(p=0.5), A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5), A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), ToTensorV2()])# 包装成PyTorch兼容的转换类class AlbumentationsWrapper: def __init__(self, transform): self.transform = transform def __call__(self, img): # 确保输入是PIL图像或numpy数组 img_np = np.array(img) augmented = self.transform(image=img_np) return augmented['image']# 使用包装后的转换final_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), AlbumentationsWrapper(albumentations_transform)])## 性能优化与注意事项1. 批处理转换 vs 逐样本转换transforms.Compose 是逐样本执行的,对于大数据集会有性能瓶颈。在 DataLoader 中设置 num_workers>0 可以利用多进程并行处理。2. 避免在 GPU 上做转换:数据转换应在 CPU 上完成,GPU 只负责模型计算。使用 pin_memory=True 可以加速 CPU→GPU 的数据传输。3. 缓存预处理结果:如果数据量巨大且增强方式固定(如只做标准化),可以预先处理并保存为 .pt 文件,训练时直接加载。4. 保持随机性的一致性:对于验证集或测试集,不要使用随机增强。如果需要可视化,可以单独创建一个无增强的 Dataset。## 总结本文从实战角度深入讲解了 PyTorch 中的数据转换机制。我们首先介绍了 torchvision.transforms 的基本使用方式,然后通过两个完整的代码示例演示了图像分类中的常见增强策略和自定义转换的实现方法。此外,还介绍了如何混合使用 albumentations 库来扩展功能,并给出了性能优化的实用建议。数据转换是连接原始数据和深度学习模型之间的桥梁,它直接影响模型的训练效率和最终性能。掌握好 transforms.Composetransforms.Lambda 以及自定义转换类的编写,将让你在处理各种复杂数据时游刃有余。在后续的工程实践中,建议根据具体任务需求灵活组合不同的变换,并通过可视化验证转换后的数据是否符合预期。记住:好的数据转换策略,往往比复杂的模型架构更能带来性能提升。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐