【pytorch】transform的使用

一、transforms的用法

transforms​ 是数据预处理与增强的核心工具,主要用于将原始图像转换为模型可接受的格式,并通过随机变换丰富数据集以提高模型泛化能力。

导入方式:

python 复制代码
from torchvision import transforms

主要用法,按顺序

python 复制代码
transform_pipeline = transforms.Compose([
    transforms.Resize(256),          # 调整图像大小至256x256(保持宽高比)
    transforms.CenterCrop(224),      # 从中心裁剪224x224区域(常用预训练模型输入尺寸)
    transforms.RandomHorizontalFlip(p=0.5),  # 以50%概率水平翻转(数据增强)
    transforms.ToTensor(),           # 将PIL图像转换为Tensor(像素值缩放至[0,1])
    transforms.Normalize(            # 标准化(使用ImageNet均值/方差)
        mean=[0.485, 0.456, 0.406],  # RGB通道均值
        std=[0.229, 0.224, 0.225]    # RGB通道标准差
    )
])

二、transform的使用

将PIL图像转换成Tensor类型

python 复制代码
from PIL import Image
from torchvision import transforms

img_path = r'data/train/ants_image/0013035.jpg'
img = Image.open(img_path)
tensor_trans = transforms.ToTensor()
tensor_img = tensor_trans(img)
print(tensor_img.shape)    #CHW

通过tensor()类型的数据生成tensorboard图

python 复制代码
from PIL import Image
from torch.utils.tensorboard import SummaryWriter
from torchvision import transforms

img_path = r'data/train/ants_image/0013035.jpg'
img = Image.open(img_path)
tensor_trans = transforms.ToTensor()
tensor_img = tensor_trans(img)
# print(tensor_img.shape)    #CHW
writer = SummaryWriter('logs')
writer.add_image('tensor_img', tensor_img, 0)
writer.close()

Normalize()归一化使用

python 复制代码
from PIL import Image
from torch.utils.tensorboard import SummaryWriter
from torchvision import transforms

img_path = r'data/train/ants_image/0013035.jpg'
img = Image.open(img_path)
tensor_trans = transforms.ToTensor()
tensor_img = tensor_trans(img)
# print(tensor_img.shape)    #CHW

writer = SummaryWriter('logs')
norm_trans = transforms.Normalize([0.485, 0.456, 0.406], [0.5, 0.5, 0.5])
norm_img = norm_trans(tensor_img)

writer.add_image('tensor_img', tensor_img, 0)
writer.add_image('norm_img', norm_img, 1)
writer.close()

归一化后的图片和未归一化的图片

Resize()调整大小的使用

python 复制代码
from PIL import Image
from torch.utils.tensorboard import SummaryWriter
from torchvision import transforms

img_path = r'data/train/ants_image/0013035.jpg'
img = Image.open(img_path)
tensor_trans = transforms.ToTensor()
tensor_img = tensor_trans(img)
# print(tensor_img.shape)    #CHW

writer = SummaryWriter('logs')
norm_trans = transforms.Normalize([0.485, 0.456, 0.406], [0.5, 0.5, 0.5])
norm_img = norm_trans(tensor_img)

# print(img.size)
resize_trans = transforms.Resize((256, 256))
resize_img = resize_trans(tensor_img)
writer.add_image('resize_img', resize_img, 0)
# print(resize_img.size)
#Compose用法
trans_resize_2 = transforms.Compose([transforms.Resize((512)), transforms.ToTensor()])
img_resize_2 = trans_resize_2(img)


writer.add_image('tensor_img', tensor_img, 0)
writer.add_image('norm_img', norm_img, 1)
writer.add_image('img_resize_2', img_resize_2, 2)
writer.close()
相关推荐
风合星语2 天前
2026 具身智能技术实战(一):VLA 到底怎么控制机器人?——用 LeRobot 跑通 SmolVLA 推理
pytorch·机器人·具身智能·vla·lerobot·smolvla
Tancenter2 天前
sequeeze()和unsequeeze()
pytorch·tensorr
Thomas.Sir2 天前
第21课:PyTorch|GPU多卡训练与分布式训练基础【让多卡并行成为你的加速引擎】
人工智能·pytorch·分布式
CODER03042 天前
ubuntu22.04部署完整deepseek局域网web服务全过程(RTX5090安装黑屏+web端多人并发)
pytorch·webui·ubuntu22.04·ollama·deepseek·rtx5090·自然语言模型
宿州派大星3 天前
[NLP实战] 基于PyTorch实现N-gram词嵌入模型:输入4个词预测第5个词
人工智能·pytorch·深度学习·nlp
Liaiyang663 天前
空圈容错视角下的无人机全链路审计:从理论框架到耦合式检验
人工智能·pytorch·python·深度学习·系统架构·自动驾驶·无人机
AI模力圈3 天前
Pytorch图模式技术原理解析
pytorch·深度学习·torch.compile
Tancenter3 天前
gather和scatter API
pytorch·tensor
磁场转动100万匹3 天前
基于 dlib 与 OpenCV 的疲劳驾驶检测:眼睛纵横比(EAR)原理与代码逐段解析
pytorch·python
Dr_Fourier3 天前
AWQ量化
c++·人工智能·pytorch·ai