Pytorch中transforms详解
·
初步使用
在深度学习中,考虑到模型对输入数据的格式、尺寸有着硬性要求,为了便于统一处理,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 XYXY↔XYWH↔CXCYWH |
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版本以获得更好的兼容性和功能支持。
更多推荐

所有评论(0)