Работа с изображениями в PyTorch

Глубокое обучение на PyTorch: средний уровень

Michal Oleszak

Machine Learning Engineer

Набор данных облаков

Примеры из набора данных облаков: пять изображений с разными типами облаков.

1 https://www.kaggle.com/competitions/cloud-type-classification2/data
Глубокое обучение на PyTorch: средний уровень

Что такое изображение?

Изображение облака с увеличенным фрагментом, на котором видны пиксели.

  • Изображение состоит из пикселей («picture elements»)
  • Каждый пиксель содержит информацию о цвете

  • Оттенки серого: целое число от 0 до 255

    • 30:

Серый прямоугольник

  • Цветные изображения: три целых числа — по одному на каждый цветовой канал (Red, Green, Blue)
    • RGB = (52, 171, 235):

Синий прямоугольник

Глубокое обучение на PyTorch: средний уровень

Загрузка изображений в PyTorch

Требуемая структура директорий:

clouds_train

- cumulus
- 75cbf18.jpg - ...
- cumulonimbus - ...
clouds_test
- cumulus - cumulonimbus - ...

 

  • Основные папки: clouds_train и clouds_test
  • Внутри каждой основной папки — по одной папке на категорию
  • Внутри каждой папки класса — файлы изображений
Глубокое обучение на PyTorch: средний уровень

Загрузка изображений в PyTorch

from torchvision.datasets import ImageFolder
from torchvision import transforms


train_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Resize((128, 128)), ])
dataset_train = ImageFolder( "data/clouds_train", transform=train_transforms, )
  • Задайте преобразования:

    • Преобразование в тензор
    • Изменение размера до 128×128
  • Создайте набор данных, передав:

    • Путь к данным
    • Заданные преобразования
Глубокое обучение на PyTorch: средний уровень

Отображение изображений

dataloader_train = DataLoader(
    dataset_train, 
    shuffle=True, 
    batch_size=1,
)

image, label = next(iter(dataloader_train))
print(image.shape)
torch.Size([1, 3, 128, 128])
image = image.squeeze().permute(1, 2, 0)
print(image.shape)
torch.Size([128, 128, 3])
import matplotlib.pyplot as plt
plt.imshow(image)
plt.show()

Вывод изображения облака после вызова plt.show

Глубокое обучение на PyTorch: средний уровень

Аугментация данных

train_transforms = transforms.Compose([

transforms.RandomHorizontalFlip(), transforms.RandomRotation(45),
transforms.ToTensor(), transforms.Resize((128, 128)), ])
dataset_train = ImageFolder( "data/clouds/train", transform=train_transforms, )

Аугментация данных: генерация дополнительных данных путём случайных преобразований исходных изображений

  • Увеличение объёма и разнообразия обучающей выборки
  • Повышение устойчивости модели
  • Снижение переобучения

Три изображения облаков, демонстрирующих преобразование поворота

Глубокое обучение на PyTorch: средний уровень

Давайте потренируемся!

Глубокое обучение на PyTorch: средний уровень

Preparing Video For Download...