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

Середній рівень Deep Learning з PyTorch

Michal Oleszak

Machine Learning Engineer

Набір даних про хмари

Зразки з набору даних про хмари: пʼять зображень різних типів хмар.

1 https://www.kaggle.com/competitions/cloud-type-classification2/data
Середній рівень Deep Learning з PyTorch

Що таке зображення?

Зображення хмари з фрагментом, збільшеним до видимих пікселів.

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

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

    • 30:

Сірий прямокутник

  • Кольорові: три цілі числа — по одному для каналів (Red, Green, Blue)
    • RGB = (52, 171, 235):

Синій прямокутник

Середній рівень Deep Learning з PyTorch

Завантаження зображень у PyTorch

Бажана структура тек:

clouds_train

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

 

  • Основні теки: clouds_train і clouds_test
  • Усередині кожної — по одній теці на категорію
  • Усередині теки класу — файли зображень
Середній рівень Deep Learning з 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, )
  • Визначте перетворення:

    • Перетворити на tensor
    • Змінити розмір до 128x128
  • Створіть набір даних, передаючи:

    • Шлях до даних
    • Заздалегідь визначені перетворення
Середній рівень Deep Learning з 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

Середній рівень Deep Learning з 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, )

Аугментація даних: створення додаткових даних шляхом випадкових перетворень оригінальних зображень

  • Збільшує розмір і різноманітність тренувального набору
  • Підвищує стійкість моделі
  • Зменшує перенавчання

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

Середній рівень Deep Learning з PyTorch

Давайте потренуємось!

Середній рівень Deep Learning з PyTorch

Preparing Video For Download...