欢迎光临
我们一直在努力

transforms

一.transforms 是什么?

1.名词释义

torchvision.transforms :PyTorch 专门为图像数据设计的预处理&数据增强工具包。

2.通俗理解

神经网络不能直接看懂手机/电脑里的图片,就像人吃食物要先洗菜、切菜、调味。
transforms 就是给图片“加工处理”:改大小、转格式、调颜色、随机翻转等,把原始图片变成模型能识别、能学得更好的数据。

3.必备导入代码

#核心库
import torch
#图像变换工具
from torchvision import transforms
#Python原生图片读取工具
from PIL import Image

二.两个核心图片格式

所有变换都是围绕这两种格式展开

1.PIL 图像(原始图片格式)

  • 来源:用 Image.open() 读取本地图片得到的格式
  • 维度规则:(高度H, 宽度W, 通道数C)
    例:一张 500高、300宽 的彩色图 → (500, 300, 3)
  • 像素值域:每个像素值范围** 0 ~ 255**(整数)
  • 特点:人类能直观打开查看,绝大多数基础变换只支持该格式

2.Tensor 张量(模型专属格式)

  • 来源:经过 ToTensor() 转换得到
  • 维度规则:(通道数C, 高度H, 宽度W)
    例:上面同一张图 → (3, 500, 300)
  • 像素值域: **0.0 ~ 1.0 (**浮点数)
  • 特点:PyTorch 神经网络只认这种格式

一句话总结顺序:原始图片(PIL) → 各种裁剪/翻转 → 转Tensor → 标准化 → 送入模型

三.总控制器:Compose 组合变换

1.名词释义

transforms.Compose(transform_list):变换流水线,把多个处理步骤按书写顺序依次执行。

2.通俗理解

相当于“流水线工序单”,把改大小、翻转、转格式等多个操作排好队,一张图片进来,自动走完所有流程,不用手动一步步处理。

3.语法 & 示例

#定义一套处理流程,列表里按顺序写操作
train_transform = transforms.Compose([
transforms.Resize(224), # 第1步:缩放图片
transforms.RandomHorizontalFlip(), # 第2步:随机左右翻转
transforms.ToTensor(), # 第3步:转为张量
transforms.Normalize(...) # 第4步:标准化
])

⚠️ 关键规则:列表顺序 = 代码执行顺序,顺序错了直接报错/效果异常。

四.分类详解:所有常用变换

第一类:尺寸与裁剪(修改图片大小)

1.Resize 缩放
  • 作用:把图片缩放到指定尺寸,统一图片大小(模型要求输入尺寸固定)
  • 参数说明
  • 传入单个数字: Resize(n) → 保持图片原有比例,把短边缩为 n
    例:原图(500,300) → Resize(224) → 短边300→224,长边等比例缩小
  • 传入元组(H, W): Resize((h,w)) → 强制改成 高h、宽w,可能拉伸图片
  • :::info
    torchvision 的 Resize 还有个参数 ** max_size **,搭配单数字使用:

    transforms.Resize(224,max_size=400)

    含义:

    短边缩到224,如果算出来的长边超过 400,就把长边限制在400,防止图片变得过长/过宽。

    :::

    • 代码示例

    transforms.Resize(256) # 等比例缩放,短边=256
    transforms.Resize((224,224)) # 强制改为 224×224 正方形

    • 使用场景:训练、测试集通用,几乎所有图像任务必用
    2.CenterCrop 中心裁剪
    • 作用:从图片正中间抠出一块指定大小的区域
    • 示例: CenterCrop(224) → 在图片中心裁剪出 224×224 的图
    • 使用场景:测试集/验证集专用
    • **原因:**测试时要保证结果稳定,固定从中心裁剪,不随机改动图片
    3.RandomCrop 随机裁剪
    • 作用:在图片随机位置抠图,属于数据增强
    • 使用场景:训练集专用
    • 通俗理解:同一张图,每次训练随机裁不同位置,相当于凭空多了很多训练数据,提升模型泛化能力
    4.RandomResizedCrop随机缩放+随机裁剪
    • 作用:先随机把图片放大/缩小,再随机裁剪,ImageNet 分类任务标准增强操作
    • 示例: RandomResizedCrop(224)
    • 使用场景:大型图像分类训练集首选

    第二类:翻转与旋转(数据增强主力)

    这类都是随机操作,只用于训练集,测试集禁用!

    1.RandomHorizontalFlip 随机水平翻转
    • 参数:p=0.5 (概率,默认50%概率翻转)
    • 作用:图片左右翻转(类似照镜子)
    • 示例:人脸、猫狗图片,左右翻转不改变物体含义
    • 使用场景:90% 图像训练任务必加,安全、效果好
    2.RandomVerticalFlip 随机垂直翻转
    • 作用:图片上下翻转(倒立)
    • 慎用场景:文字、人脸、建筑物、车牌等,倒立后语义完全改变
    3.RandomRotation 随机旋转
    • 参数:

    角度范围

  • RandomRotation(10):**随机旋转 ±10°
    **2. RandomRotation((0, 180)) :在 0~180° 之间随机旋转
    • 使用场景:航拍图、手写数字、花卉等允许旋转的任务

    第三类:格式转换

    1.ToTensor()
    • 专业功能:完成两大转换
  • 维度变换:PIL格式 (H,W,C) → Tensor格式 (C,H,W)
  • 值域归一化:像素值 [0,255](整数)→ [0.0, 1.0](浮点数)
    • 通俗理解:把“人类看懂的图片”翻译成“模型能看懂的张量”
    • 硬性规则:
  • 只能放在所有裁剪、翻转操作之后
  • 必须放在 Normalize 之前
  • 输入只能是 PIL 图片,不能是张量
  • 2.ToPILImage()
    • 作用:反向转换,Tensor → PIL 图片
    • 使用场景:训练后可视化图片、查看处理效果

    第四类:Normalize 标准化(模型收敛关键)

    1.公式(专业)

    输出像素=(输入像素-mean)/std

    • mean:均值,std:标准差
    2.通俗解释

    把 ToTensor() 后 [0,1] 的像素值,再次调整分布,让数据均值趋近于0、方差趋近于1。
    作用:加速神经网络训练,让模型更快收敛、精度更高。

    3.语法与参数

    transforms.Normalize(mean, std)

    • mean :每个通道的均值列表
    • std :每个通道的标准差列表
    4.通用标准参数
  • RGB彩色图(3通道)(ResNet/VGG/预训练模型通用)
  • mean = [0.485, 0.456, 0.406]
    std = [0.229, 0.224, 0.225]

  • 灰度图(单通道)
  • mean = [0.5]
    std = [0.5]

    5.铁律

    Normalize 必须紧跟在 ToTensor() 后面,顺序颠倒会计算错误。

    第五类:色彩增强(应对光照变化)

    ColorJitter 颜色抖动

    1.作用:随机修改图片的亮度、对比度、饱和度、色相,模拟不同光照环境
    2.参数详解(通俗版)

    – `brightness `:亮度,数值越大图片忽明忽暗越明显
    – ` contrast `:对比度,明暗差距变化
    – `saturation` :色彩鲜艳程度
    – `hue `:色调(颜色偏向,建议数值小于0.5)

    3.代码示例

    #亮度/对比度随机浮动20%
    transforms.ColorJitter(brightness=0.2, contrast=0.2)

    4.场景:户外图片、监控图像、光照不稳定的数据集

    第六类:遮挡增强(提升鲁棒性)

    RandomErasing 随机擦除

    1.作用:在图片上随机擦掉一块区域(打马赛克)
    2.通俗理解:训练时故意遮挡部分物体,强迫模型学习全局特征,不会因为物体被遮挡就识别错误
    3.位置要求:放在 Normalize 之后
    4.场景:分类、目标检测训练集

    第七类:高斯模糊

    GaussianBlur 高斯模糊

    1.作用:随机模糊图片,模拟拍照失焦、画质模糊的情况
    2.场景:低画质图像、监控图像增强

    五、工业级标准模板

    核心原则:
    ✅ 训练集:加所有随机增强(翻转、裁剪、色彩抖动),提升模型能力
    ✅ 测试/验证集:禁用所有随机操作,保证结果稳定

    1.训练集变换(带数据增强)

    train_transform = transforms.Compose([
    transforms.RandomResizedCrop(224), # 随机缩放裁剪
    transforms.RandomHorizontalFlip(0.5), # 随机左右翻转
    transforms.ColorJitter(0.2, 0.2), # 色彩增强
    transforms.ToTensor(), # 转张量
    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225]), # 标准化
    transforms.RandomErasing(p=0.2) # 随机遮挡
    ])

    2.测试/验证集变换(无增强)

    val_transform = transforms.Compose([
    transforms.Resize(256), # 等比例缩放
    transforms.CenterCrop(224), # 中心裁剪(无随机)
    transforms.ToTensor(),
    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])
    ])

    六.自定义 Transform(拓展用法)

    当内置操作不够用时,可以自己写处理逻辑,两种常用方式:

    1.Lambda 简单自定义(单行逻辑)

    把一个简易函数嵌入流水线

    # 示例:把彩色图转为灰度图
    transform = transforms.Compose([
    transforms.Lambda(lambda img: img.convert("L")),
    transforms.Resize(224),
    transforms.ToTensor()
    ])

    2.类继承(复杂逻辑,标准写法)

    继承 torch.nn.Module ,适合多行复杂处理,和官方变换用法完全一致

    class MyCustomTransform(torch.nn.Module):
    # 固定写法:__call__ 方法,输入图片,返回处理后图片
    def __call__(self, img):
    # 自定义你的处理逻辑
    return img

    # 放入流水线即可使用
    transform = transforms.Compose([MyCustomTransform()])

    七.高频易错点

    1.顺序错误(最常见)
    错误: Normalize → ToTensor()
    正确:裁剪/翻转 → ToTensor() → Normalize

  • 维度混淆
    PIL: (H, W, C) | Tensor:(C, H, W),打印 shape 即可检查

  • 测试集使用随机操作
    测试集加 RandomFlip/RandomCrop → 每次运行结果不一样,指标波动大

  • 通道不匹配
    灰度图用了RGB的 mean/std → 维度报错,单通道必须用 [0.5], [0.5]

  • 值域混淆
    PIL:0~255 | ToTensor后:0~1 | Normalize后:正负浮点数

  • 八.完整可运行代码

    from PIL import Image
    from torchvision import transforms

    # 1. 定义变换流水线
    trans = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])
    ])

    # 2. 读取本地图片(替换成你的图片路径)
    img = Image.open("test.jpg").convert("RGB")

    # 3. 执行所有变换
    img_tensor = trans(img)

    # 4. 查看结果:形状 + 像素范围
    print("张量形状:", img_tensor.shape) # 输出 torch.Size([3, 224, 224])
    print("像素最大值:", img_tensor.max().item())
    print("像素最小值:", img_tensor.min().item())

    赞(0)
    未经允许不得转载:171主机测评 » transforms
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址