使用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)

到此这篇关于PyTorch 使用torchvision进行图片数据增广的文章就介绍到这了,更多相关PyTorch torchvision 图片增广内容请搜索Devmax以前的文章或继续浏览下面的相关文章希望大家以后多多支持Devmax!

PyTorch 使用torchvision进行图片数据增广的更多相关文章

  1. Python使用pytorch动手实现LSTM模块

    这篇文章主要介绍了Python使用pytorch动手实现LSTM模块,LSTM是RNN中一个较为流行的网络模块。主要包括输入,输入门,输出门,遗忘门,激活函数,全连接层(Cell)和输出

  2. Pytorch搭建yolo3目标检测平台实现源码

    这篇文章主要为大家介绍了Pytorch搭建yolo3目标检测平台实现源码,有需要的朋友可以借鉴参考下,希望能够有所帮助,祝大家多多进步,早日升职加薪

  3. PyTorch搭建双向LSTM实现时间序列负荷预测

    这篇文章主要为大家介绍了PyTorch搭建双向LSTM实现时间序列负荷预测,有需要的朋友可以借鉴参考下,希望能够有所帮助,祝大家多多进步,早日升职加薪

  4. pytorch使用nn.Moudle实现逻辑回归

    这篇文章主要为大家详细介绍了pytorch使用nn.Moudle实现逻辑回归,文中示例代码介绍的非常详细,具有一定的参考价值,感兴趣的小伙伴们可以参考一下

  5. pytorch加载自己的图片数据集的2种方法详解

    数据预处理在解决深度学习问题的过程中,往往需要花费大量的时间和精力,下面这篇文章主要给大家介绍了关于pytorch加载自己的图片数据集的2种方法,文中通过示例代码介绍的非常详细,需要的朋友可以参考下

  6. PyTorch实现手写数字的识别入门小白教程

    这篇文章主要介绍了python实现手写数字识别,非常适合小白入门学习,本文通过实例图文相结合给大家介绍的非常详细,对大家的学习或工作具有一定的参考借鉴价值,需要的朋友可以参考下

  7. pytorch人工智能之torch.gather算子用法示例

    这篇文章主要介绍了pytorch人工智能之torch.gather算子用法示例,有需要的朋友可以借鉴参考下,希望能够有所帮助,祝大家多多进步,早日升职加薪

  8. Pytorch深度学习addmm()和addmm_()函数用法解析

    这篇文章主要为大家介绍了Pytorch中addmm()和addmm_()函数用法解析,有需要的朋友可以借鉴参考下,希望能够有所帮助,祝大家多多进步,早日升职加薪

  9. 基于Pytorch实现逻辑回归

    这篇文章主要为大家详细介绍了基于Pytorch实现逻辑回归,文中示例代码介绍的非常详细,具有一定的参考价值,感兴趣的小伙伴们可以参考一下

  10. pytorch关于Tensor的数据类型说明

    这篇文章主要介绍了pytorch关于Tensor的数据类型说明,具有很好的参考价值,希望对大家有所帮助。如有错误或未考虑完全的地方,望不吝赐教

随机推荐

  1. 10 个Python中Pip的使用技巧分享

    众所周知,pip 可以安装、更新、卸载 Python 的第三方库,非常方便。本文小编为大家总结了Python中Pip的使用技巧,需要的可以参考一下

  2. python数学建模之三大模型与十大常用算法详情

    这篇文章主要介绍了python数学建模之三大模型与十大常用算法详情,文章围绕主题展开详细的内容介绍,具有一定的参考价值,感想取得小伙伴可以参考一下

  3. Python爬取奶茶店数据分析哪家最好喝以及性价比

    这篇文章主要介绍了用Python告诉你奶茶哪家最好喝性价比最高,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价值,需要的朋友们下面随着小编来一起学习吧

  4. 使用pyinstaller打包.exe文件的详细教程

    PyInstaller是一个跨平台的Python应用打包工具,能够把 Python 脚本及其所在的 Python 解释器打包成可执行文件,下面这篇文章主要给大家介绍了关于使用pyinstaller打包.exe文件的相关资料,需要的朋友可以参考下

  5. 基于Python实现射击小游戏的制作

    这篇文章主要介绍了如何利用Python制作一个自己专属的第一人称射击小游戏,文中的示例代码讲解详细,感兴趣的小伙伴可以跟随小编一起动手试一试

  6. Python list append方法之给列表追加元素

    这篇文章主要介绍了Python list append方法如何给列表追加元素,具有很好的参考价值,希望对大家有所帮助。如有错误或未考虑完全的地方,望不吝赐教

  7. Pytest+Request+Allure+Jenkins实现接口自动化

    这篇文章介绍了Pytest+Request+Allure+Jenkins实现接口自动化的方法,文中通过示例代码介绍的非常详细。对大家的学习或工作具有一定的参考借鉴价值,需要的朋友可以参考下

  8. 利用python实现简单的情感分析实例教程

    商品评论挖掘、电影推荐、股市预测……情感分析大有用武之地,下面这篇文章主要给大家介绍了关于利用python实现简单的情感分析的相关资料,文中通过示例代码介绍的非常详细,需要的朋友可以参考下

  9. 利用Python上传日志并监控告警的方法详解

    这篇文章将详细为大家介绍如何通过阿里云日志服务搭建一套通过Python上传日志、配置日志告警的监控服务,感兴趣的小伙伴可以了解一下

  10. Pycharm中运行程序在Python console中执行,不是直接Run问题

    这篇文章主要介绍了Pycharm中运行程序在Python console中执行,不是直接Run问题,具有很好的参考价值,希望对大家有所帮助。如有错误或未考虑完全的地方,望不吝赐教

返回
顶部