初步使用

在深度学习中,考虑到模型对输入数据的格式、尺寸有着硬性要求,为了便于统一处理,Pytorch提供了transforms用于数据预处理。transforms包括v1和v2两个版本,出于兼容性考虑,直接用from torchvision import transforms默认为v1,但推荐推荐使用v2,常见用法如下

import matplotlib.pyplot as plt
from PIL import Image
from torchvision.transforms import v2
import torch

transform = v2.Compose([
    v2.Resize((448, 448)),
    v2.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.05),
    v2.ToImage(),
    v2.ToDtype(torch.float32, scale=True),
    v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
path = f"test.jpg"
img = Image.open(path).convert("RGB")
imgT = transform(img)

代码中,

  • 【Compose】用于将多个变换打包,按列表顺序依次调用。
  • 【Resize】用于统一图像尺寸,将图片强制缩放为 448 × 448 448\times448 448×448
  • 【ColorJitter】用于随机扰动图像颜色,防止模型过拟合。上面的示例中,输入参数分别为亮度变化 ± 0.1 \pm0.1 ±0.1,对比度变化 ± 0.1 \pm0.1 ±0.1,饱和度变化 ± 0.1 \pm0.1 ±0.1,色调变化 ± 0.05 \pm0.05 ±0.05
  • 【ToImage】将数据类型转换为【tv_tensor.Image】类型
  • 【ToDtype】将通道顺序从 h × w × c \mathrm {h\times w\times c} h×w×c调整为 c × h × w \mathrm {c\times h\times w} c×h×w,且将像素值从 [ 0 , 255 ] [0,255] [0,255]线性缩放到 [ 0.0 , 1.0 ] [0.0, 1.0] [0.0,1.0]
  • 【Nomarlize】归一化,其输出值output=(input-mean)/std

在v1版本中,【ToImage, ToDtype】用【ToTensor】一步实现,所以看到【ToTensor】不要慌。

处理结果

img经过transform处理前后对比为

在这里插入图片描述

可视化代码如下,需要注意,由于在变换过程中,【Tensor】将其通道顺序变成了 c × h × w \mathrm {c\times h\times w} c×h×w,在可视化之前应该重新排布其维度。

ax = plt.subplot(121)
ax.imshow(img)
ax.axis('off')
ax = plt.subplot(122)
ax.imshow(imgT.permute(1,2,0))
ax.axis('off')
plt.show()

常用变换

transforms中提供了许多变换方案,其v2版本中提供的变换类型如下列诸表所示。


组合与控制流 功能 v2 特性
Compose(transforms) 顺序执行变换列表 支持传入 tv_tensors,自动分发到各子变换
RandomApply(transforms, p=0.5) 以概率 p 应用一组变换 内部状态统一管理,支持多模态同步
RandomChoice(transforms) 随机挑选 1 个变换执行 可设置各变换权重(p 列表)
RandomOrder(transforms) 随机打乱顺序后执行 适合需要增强随机性的场景
Lambda(func) 包装任意 Python 函数 仍保留,但官方更推荐继承 nn.Module 自定义
类型转换 功能 替代 v1 对应项
ToImage() 将 PIL/numpy/Tensor 统一转为 tv_tensors.Image 替代隐式类型推断,保留空间元数据
ToDtype(dtype, scale=False) 转换数据类型。scale=True 时自动 [0,255]→[0,1] 替代 ToTensor()(官方已标记废弃)
PILToTensor() PIL → Tensor(整数,不缩放) 适合需保留原始整型像素的任务
ToPILImage(mode=None) Tensor → PIL(调试/保存用) 行为与 v1 一致
Normalize(mean, std, inplace=False) 通道标准化(仅支持 Tensor) 必须在 ToDtype 之后使用
几何与空间变换 功能 关键参数/备注
Resize(size, interpolation, antialias=True) 调整尺寸 antialias 默认开启,防锯齿
CenterCrop(size) 中心裁剪 支持 (h,w) 或单值
RandomCrop(size, padding, pad_if_needed, fill, padding_mode) 随机裁剪 fill 支持标量/元组/字符串
RandomResizedCrop(size, scale, ratio, interpolation) 随机缩放+裁剪 ResNet/ViT 训练标配
RandomHorizontalFlip(p), RandomVerticalFlip(p) 随机翻转 默认 p=0.5
RandomRotation(degrees, interpolation, expand, center, fill) 随机旋转 expand=True 防裁剪黑边
Affine(degrees, translate, scale, shear, interpolation, fill, center) 仿射变换 统一替代 RandomAffine
Pad(padding, fill, padding_mode) 填充 padding_mode 支持 constant/edge/reflect/symmetric
RandomErasing(p, scale, ratio, value, inplace) 随机遮挡 防过拟合经典策略
ElasticTransform(alpha, sigma, interpolation, fill) 弹性形变 医学图像/OCR 常用
Perspective(distortion_scale, interpolation, fill) 透视变换 模拟相机视角变化
FiveCrop(size), TenCrop(size, vertical_flip) 多区域裁剪 传统验证集增强(现多被 Resize+CenterCrop 替代)
颜色与像素级变换 功能 参数说明
ColorJitter(brightness, contrast, saturation, hue) 综合颜色扰动 参数 <1 为比例因子,>1 为绝对值
Grayscale(num_output_channels=1), RandomGrayscale(p) 灰度化 模拟单通道传感器
GaussianBlur(kernel_size, sigma), RandomGaussianBlur(..., p) 高斯模糊 sigma 支持元组范围或标量
RandomAdjustSharpness(sharpness_factor, p) 随机锐化 factor>1 锐化,<1 模糊
RandomPosterize(bits, p) 降低色彩深度 bits 范围 0~8
RandomSolarize(threshold, p) 色调反转 threshold 以上像素反转
RandomEqualize(p), RandomInvert(p), RandomAutocontrast(p) 直方图均衡/反色/自动对比度 模拟不同拍摄条件
RandomPhotometricDistort(p, contrast, brightness, saturation, hue) 组合光照扰动 高效替代多次 ColorJitter
现代增强策略 功能 使用注意
RandAugment(num_ops, magnitude, interpolation) 自动增强策略 num_ops 通常 2~3,magnitude 0~1
TrivialAugmentWide(num_magnitude_bins, interpolation) 单变换强扰动 计算量低,ImageNet 竞赛常用
AugMix(magnitude, alpha, width, depth, interpolation) 多分支混合增强 抗分布偏移(OOD)效果显著
MixUp(alpha, labels_getter) Batch 级样本混合 需配合 DataLoader,修改 Loss 计算方式
CutMix(alpha, labels_getter) Batch 级区域裁剪混合 同上,labels_getter 处理标签字典/元组
检测框/分割掩码专用 功能 适用场景
SanitizeBoundingBoxes(min_size, labels_getter) 移除越界/过小 BBox 目标检测数据清洗
ClampBoundingBoxes() 将 BBox 限制在图像边界内 防止裁剪/缩放后坐标越界
ConvertBoundingBoxFormat(format) 坐标格式转换 X Y X Y ↔ X Y W H ↔ C X C Y W H XYXY ↔ XYWH ↔ CXCYWH XYXYXYWHCXCYWH
RandomIoUCrop(scale, ratio, min_ious, sampler) SSD 专用随机裁剪 同步变换 Image + BBoxes + Labels
ScaleJitter(target_size, scale_range) 多尺度抖动 YOLO/DETR 等检测模型训练

本文介绍了PyTorch中torchvision.transforms模块的数据预处理方法,重点对比了v1和v2版本的区别。主要内容包括:1)基本使用方法,通过Compose组合Resize、ColorJitter等变换;2)处理结果可视化注意事项;3)详细分类整理了v2版本提供的各类变换,包括组合控制、类型转换、几何变换、颜色调整等,并标注了与v1版本的差异。特别强调了v2版本新增的现代增强策略(如RandAugment、MixUp)和目标检测专用变换,推荐使用v2版本以获得更好的兼容性和功能支持。

Logo

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

更多推荐