PyTorch 使用torchvision进行图片数据增广
作者:峡谷的小鱼 发布时间:2023-06-19 23:09:10
标签:PyTorch,torchvision,图片增广
使用torchvision来进行图片的数据增广
数据增强就是增强一个已有数据集,使得有更多的多样性。对于图片数据来说,就是改变图片的颜色和形状等等。比如常见的:
左右翻转,对于大多数数据集都可以使用;
上下翻转:部分数据集不适合使用;
图片切割:从图片中切割出一个固定的形状,
随机高宽比(e.g. [3/4, 4/3)
随机大小(e.g. [8%, 100%])
随机位置
改变图片的颜色
改变色调,饱和度,明亮度(e.g. [0.5, 1.5])
1. 读取图片
加载相关包。
import torch
import torchvision
import matplotlib
from torch import nn
from torchvision import transforms
from PIL import Image
from IPython import display
from matplotlib import pyplot as plt
选取一个狗的图片作为示例:
def set_figsize(figsize=(3.5, 2.5)):
display.set_matplotlib_formats('svg')
plt.rcParams['figure.figsize'] = figsize
def show_images(imgs, num_rows, num_cols, titles=None, scale=1.5):
r"""
展示一列图片
img: Image对象的列表
"""
figsize = (num_cols * scale, num_rows * scale)
fig, axes = plt.subplots(num_rows, num_cols, figsize=figsize)
axes = axes.flatten()
for i, (ax, img) in enumerate(zip(axes, imgs)):
if torch.is_tensor(img):
ax.imshow(img.numpy())
else:
ax.imshow(img)
ax.get_xaxis().set_visible(False)
ax.get_yaxis().set_visible(False)
ax.get_xaxis().set_label('x')
if titles:
ax.set_title(titles[i])
return axes
set_figsize()
img = Image.open('img/dog1.jpg')
plt.imshow(img);
2. 图片增广
def apply(img, aug, num_rows=2, num_clos=4, scale=1.5):
# 对图片应用图片增广
# img: Image object
# aug: 增广操作
Y = [aug(img) for _ in range(num_clos * num_rows)]
d2l.show_images(Y, num_rows, num_clos, scale=scale)
2.1 图片水平翻转
class RandomHorizontalFlip(torch.nn.modules.module.Module):
r'''
RandomHorizontalFlip(p=0.5)
给图片一个一定概率的水平翻转操作,如果是Tensor,要求形状为[..., H, W]
Args:
p: float, 图片翻转的概率,默认值0.5
'''
def __init__(self, p=0.5):
pass
示例。可以看到,有一般的几率对图片进行了水平翻转。
aug = transforms.RandomHorizontalFlip(0.5)
apply(img, aug)
2.2 图片上下翻转
class RandomVerticalFlip(torch.nn.modules.module.Module):
r'''
RandomVerticalFlip(p=0.5)
给图片一个一定概率的上下翻转操作,如果是Tensor,要求形状为[..., H, W]
Args:
p: float, 图片翻转的概率,默认值0.5
'''
def __init__(self, p=0.5):
pass
示例。可以看到,有一般的几率对图片进行了上下翻转。
aug = transforms.RandomHorizontalFlip(0.5)
apply(img, aug)
2.3 图片旋转
class RandomRotation(torch.nn.modules.module.Module):
r'''
将图片旋转一定角度。
'''
def __init__(self,
degrees,
interpolation=<InterpolationMode.NEAREST: 'nearest'>,
expand=False,
center=None,
fill=0):
r"""
Args:
degrees: number or sequence, 可选择的角度范围(min, max),
如果是一个数字,则范围是(-degrees, +degrees)
interpolation: Default is ``InterpolationMode.NEAREST``.
expand: bool, 如果为True,则扩展输出,使其足够大来容纳整个旋转的图像
如果为False, 将输出图像与输入图像的大小相同。
center: sequence, 以左上角为原点的旋转中心,默认是图片中心。
fill: sequence or number: 旋转图像外部区域的像素填充值,默认0。
"""
pass
def forward(self, input):
r"""
Args:
img: PIL Image or Tensor, 被旋转的图片。
Return:
PIL Image or Tensor: 旋转后的图片。
"""
pass
使用实例:
aug = transforms.RandomRotation(degrees=(-90, 90), fill=128)
apply(img, aug)
2.4 中心裁切
class CenterCrop(torch.nn.modules.module.Module):
r'''
中心裁切。
'''
def __init__(self, size):
r"""
Args:
size: sequence or int, 裁切尺寸(H, W), 如果是int,尺寸为(size, size)
"""
pass
def forward(self, input):
r"""
Args:
img: PIL Image or Tensor, 被裁切的图片。
Return:
PIL Image or Tensor: 裁切后的图片。
"""
pass
实例:
aug = transforms.CenterCrop((200, 300))
apply(img, aug)
2.5 随机裁切
class RandomCrop(torch.nn.modules.module.Module):
r'''
随机裁切。
'''
def __init__(self, size):
r"""
Args:
size: sequence or int, 裁切尺寸(H, W), 如果是int,尺寸为(size, size)
padding: sequence or int, 填充大小,
如果值为 a , 四周填充a个像素
如果值为 (a, b), 左右填充a,上下填充b
如果值为 (a, b, c, d), 左上右下依次填充
pad_if_need: bool, 如果裁切尺寸大于原图片,则填充
fill: number or str or tuple: 填充像素的值
padding_mode: str, 填充类型。
`constant`: 使用 fill 填充
`edge`: 使用边缘的最后一个值填充在图像边缘。
`reflect`: 镜像填充
"""
pass
def forward(self, input):
r"""
Args:
img: PIL Image or Tensor, 被裁切的图片。
Return:
PIL Image or Tensor: 裁切后的图片。
"""
pass
示例:
aug = transforms.RandomCrop((200, 300))
apply(img, aug)
输出:
2.6 随机裁切并修改尺寸
class RandomResizedCrop(torch.nn.modules.module.Module):
r'''
随机裁切, 并重设尺寸。
'''
def __init__(self, size, scale=(0.08, 1.0), ratio=(0.75, 1.3333333333333333)):
r"""
Args:
size: sequence or int, 需要输出的尺寸(H, W), 如果是int,尺寸为(size, size)
scale: tuple of float, 原始图片中裁切大小,百分比
ratio: tuple of float, resize前的裁切的纵横比范围
"""
pass
def forward(self, input):
r"""
Args:
img: PIL Image or Tensor, 被裁切的图片。
Return:
PIL Image or Tensor: 输出的图片。
"""
pass
示例:
aug = transforms.RandomResizedCrop((200, 200), scale=(0.2, 1))
apply(img, aug)
2. 7 修改图片颜色
class ColorJitter(torch.nn.modules.module.Module):
r'''
修改颜色。
'''
def __init__(self, brightness=0, contrast=0, saturation=0, hue=0):
r"""
Args:
brightness: float or tuple of float (min, max), 亮度的偏移幅度,范围[max(0, 1 - brightness), 1 + brightness]
contrast: float or tuple of float (min, max), 对比度偏移幅度,范围[max(0, 1 - contrast), 1 + contrast]
saturation: float or tuple of float (min, max), 饱和度偏移幅度,范围[max(0, 1 - saturation), 1 + saturation]
hue: float or tuple of float (min, max), 色相偏移幅度,范围[-hue, hue]
"""
pass
def forward(self, input):
r"""
Args:
img: PIL Image or Tensor, 输入的图片。
Return:
PIL Image or Tensor: 输出的图片。
"""
pass
示例:
aug = transforms.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.5, hue=0.5)
apply(img, aug)
3. 训练数据集加载
train_augs = transforms.Compose([transforms.RandomHorizontalFlip(),
torchvision.transforms.ToTensor()])
dataset = torchvision.datasets.CIFAR10(root="../data", train=is_train,
transform=augs, download=True)
来源:https://blog.csdn.net/weixin_43276033/article/details/124581049


猜你喜欢
- 生成静态页的方法有很多种,我比较喜欢用xmlhttp的方法生成,因为我不用考虑很多东西,我只要把动态的asp页面编写好就行了。<% s
- 安装golang使用homebrew安装golang。homebrew是MacOS 平台下的软件包管理工具,拥有安装、卸载、更新、查看、搜索
- 在本篇文章中,我们将探讨如何使用YOLOv5车牌识别系统实现实时监控与分析。我们将介绍如何将模型应用于实时视频流,以及如何分析车牌识别结果以
- 程序介绍本程序利用1.密码必须由数字、字母及特殊字符三种组合2.密码只能由字母开头3.密码长度不能低于16位来判断密码程度。首先,把可输入的
- Python关于删除list中的某个元素,一般有两种方法,pop()和remove()。remove() 函数用于移除列表中某个值的第一个匹
- WEB开发者不光要解决程序的效率问题,对数据库的快速访问和相应也是一个大问题。希望本文能对大家掌握MySQL优化技巧有所帮助。1. 优化你的
- 网站流量上来后,日志按天甚至小时存储更方便查看和管理,而Python的logging模块也提供了TimedRotatingFileHandl
- 1. 用qt designer编写主窗体,窗体类型是MainWindow,空白窗口上一个按钮。并转换成mainWindow.py# -*-
- 背景故事:我需要对一张图片做一些处理,是在图像像素级别上的数值处理,以此来反映图片 * 定区域的图像特征,网上查了很多,大多关于opencv的
- 本文介绍了python十进制和二进制的转换方法(含浮点数),分享给大家,也给自己留个笔记,具体如下:我终于写完了 , 十进制转二进制的小数部
- 小编想把用python将列表[1,1,1,1,1,1,1,1,1,1] 和 列表 [2,2,2,2,2,2,2,2,2,2]对应相加成[3,
- 一,什么是mycat一个彻底开源的,面向企业应用开发的大数据库集群支持事务、ACID、可以替代MySQL的加强版数据库一个可以视为MySQL
- 这段时间服务器崩溃2次,一直没有找到原因,今天看到论坛发出的错误信息邮件,想起可能是MySQL的默认连接数引起的问题,一查果然,老天,默认
- 1. 查找图像中出现的人脸代码示例:#导入face_recognition模块import face_recognition#将j
- 在开始本文之前,首先要保证你的mysql的密码是对的不然就要想起他的办法了。下面话不多说了,下面来一起看看吧。一、首先进入cmd 切入MyS
- 官网下载就好, https://www.python.org/downloads/release/python-352/用installer
- 给zblog添加上“运行代码”的功能,这是“密陀僧”修改z-blog源码,给z-bog增添的新功能。这个方法出来很久了,我现在才加上还不晚吧
- 1.什么是SQL语句sql语言:结构化的查询语言。(Structured Query Language),是关系数据库管理系统的标准语言。它
- 一、MySQL5.6安装后,不能正常启用压缩版MySQL,解压完后在:我的电脑->属性->高级->环境变量选择PATH,在
- 前言前段时间想实现一个短信验证码的功能,但是卡了很长时间。首先我用的是阿里云的短信服务业务,其首次接入流程如下:在阿里云上开通短信服务后需要