深度学习——数据增广

图像增广对训练数据进行一系列的随机变化,生成相似但不同的训练样本,从而扩大训练集的规模,让模型见过更多样的数据变化,减少过拟合,同时提高模型对光照、角度、噪声等变化的适应性

数据增广可以处理图片和文本和语音。对于图片的处理方式包括:覆盖掉一些像素、对颜色进行变换、对亮度进行变换,以及翻转旋转,切割。

1.基础操作

1.获取图片

%matplotlib inline
import torch
import torchvision
from torch import nn
from d2l import torch as d2l
​
d2l.set_figsize() #
img = d2l.Image.open('../img/cat1.jpg') #
d2l.plt.imshow(img)
def apply(img,aug,num_row=2,num_col=4,scale=1.5):
    Y = [aug(img) for _ in range(num_rows*num_cols)]
    d2l.show_images(Y,num_rows,num_cols,scale=scale)

2.左右翻转图像

apply(img,torchvision.transforms.RandomHorizontalFlip())

image-20250918104150226

3.上下翻转图像

apply(img,torchvision.transforms.RandomVerticalFlip())

image-20250918104307368

3.随机裁剪

shape_aug = torchvision.transforms,RandomResizedCrop((200,200),scale=(0.1,1),ratio=(0.5,2))
apply(img,shape_aug)

image-20250918104957690

4.随机改变图片的亮度ColorJitter brightness

apply(img, torchvision.transforms.ColorJitter(brightness=0.5, contrast=0, saturation=0, hue=0))

5.随机改变图片的色调ColorJitter hue

apply(img, torchvision.transforms.ColorJitter(brightness=0, contrast=0, saturation=0, hue=0.5))

6.随机改变亮度,对比度,饱和度,色调 增加或者减少50%

color_aug = torchvision.transforms.ColorJitter(
    brightness=0.5,    # 亮度变化范围
    contrast=0.5,      # 对比度变化范围  
    saturation=0.5,    # 饱和度变化范围
    hue=0.5           # 色相变化范围
)
#brightness=0.5: 亮度在 [1-0.5, 1+0.5] = [0.5, 1.5] 范围内随机变化
#contrast=0.5: 对比度在 [1-0.5, 1+0.5] = [0.5, 1.5] 范围内随机变化
#saturation=0.5: 饱和度在 [1-0.5, 1+0.5] = [0.5, 1.5] 范围内随机变化
#hue=0.5: 色相在 [-0.5, 0.5] 范围内随机偏移(注意:色相是绝对偏移,不是比例)
apply(img, color_aug)

image-20250918105331962

7.结合多种图像增广方法

augs = torchvision.transforms.Compose([torchvision.transforms.RandomHorizontalFlip(),color_aug,shape_aug])
apply(img,augs)
d2l.plt.show()

image-20250918105655104

2.完整的使用示例

import torchvision.transforms as transforms
import matplotlib.pyplot as plt
from PIL import Image
import d2l
​
def apply(img, transform, show_result=True):
    """应用变换并显示结果"""
    # 应用变换
    transformed = transform(img)
    
    if show_result:
        # 显示原图和变换后的图
        fig, axes = plt.subplots(1, 2, figsize=(12, 5))
        
        axes[0].imshow(img)
        axes[0].set_title('Original Image')
        axes[0].axis('off')
        
        axes[1].imshow(transformed)
        axes[1].set_title('Color Augmented')
        axes[1].axis('off')
        
        plt.tight_layout()
        plt.show()
    
    return transformed
​
# 使用示例
d2l.set_figsize()
img = d2l.Image.open('../img/cat1.jpg')
​
# 创建颜色增广变换
color_aug = transforms.ColorJitter(
    brightness=0.5, 
    contrast=0.5, 
    saturation=0.5, 
    hue=0.5
)
​
# 应用变换
apply(img, color_aug)

不同参数效果对比:

# 1. 只调整亮度
brightness_aug = transforms.ColorJitter(brightness=0.8)
apply(img, brightness_aug)
​
# 2. 只调整对比度  
contrast_aug = transforms.ColorJitter(contrast=0.8)
apply(img, contrast_aug)
​
# 3. 只调整饱和度
saturation_aug = transforms.ColorJitter(saturation=0.8)
apply(img, saturation_aug)
​
# 4. 只调整色相
hue_aug = transforms.ColorJitter(hue=0.3)
apply(img, hue_aug)
​
# 5. 组合所有变换
all_aug = transforms.ColorJitter(
    brightness=0.5,
    contrast=0.5, 
    saturation=0.5,
    hue=0.2  # 色相变化通常设置较小
)
apply(img, all_aug)

多次应用看随机效果

# 显示同一变换的多个随机结果
fig, axes = plt.subplots(2, 4, figsize=(16, 8))
axes = axes.flatten()
​
# 原图
axes[0].imshow(img)
axes[0].set_title('Original')
axes[0].axis('off')
​
# 应用7次随机变换
for i in range(1, 8):
    augmented = color_aug(img)
    axes[i].imshow(augmented)
    axes[i].set_title(f'Augmented {i}')
    axes[i].axis('off')
​
plt.tight_layout()
plt.show()

3.参数选择建议

  1. 保守设置(适用于大多数情况)

color_aug = transforms.ColorJitter(
    brightness=0.2,    # 亮度变化±20%
    contrast=0.2,      # 对比度变化±20%  
    saturation=0.2,    # 饱和度变化±20%
    hue=0.1           # 色相偏移±0.1
)
  1. 激进设置(数据稀少时)

color_aug = transforms.ColorJitter(
    brightness=0.5,    # 亮度变化±50%
    contrast=0.5,      # 对比度变化±50%
    saturation=0.5,    # 饱和度变化±50%
    hue=0.2           # 色相偏移±0.2
)
  1. 任务特定设置

# 医学影像:避免颜色变换
medical_aug = transforms.ColorJitter(brightness=0.1, contrast=0.1)
​
# 自然场景:可以更激进
natural_aug = transforms.ColorJitter(
    brightness=0.4, contrast=0.4, saturation=0.4, hue=0.15
)

与其他变换组合

# 组合多种增广
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1),
    transforms.RandomRotation(degrees=15),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
​
# 应用完整的预处理流水线
processed = transform(img)
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐