Obsługa obrazów w PyTorch

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Michal Oleszak

Machine Learning Engineer

Zbiór danych o chmurach

Próbki ze zbioru danych o chmurach: pięć obrazów przedstawiających różne typy chmur.

1 https://www.kaggle.com/competitions/cloud-type-classification2/data
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Czym jest obraz?

Obraz chmury z powiększonym fragmentem ukazującym widoczne piksele.

  • Obraz składa się z pikseli (ang. "picture elements")
  • Każdy piksel zawiera informację o kolorze

  • Obrazy w skali szarości: liczba całkowita z zakresu 0–255

    • 30:

Szare pole koloru

  • Obrazy kolorowe: trzy liczby całkowite, po jednej na kanał (czerwony, zielony, niebieski)
    • RGB = (52, 171, 235):

Niebieskie pole koloru

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Wczytywanie obrazów do PyTorch

Wymagana struktura katalogów:

clouds_train

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

 

  • Główne foldery: clouds_train i clouds_test
  • W każdym głównym folderze: jeden folder na kategorię
  • W każdym folderze klasy: pliki obrazów
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Wczytywanie obrazów do 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, )
  • Zdefiniuj transformacje:

    • Konwersja do tensora
    • Zmiana rozmiaru do 128x128
  • Utwórz zbiór danych, podając:

    • Ścieżkę do danych
    • Zdefiniowane transformacje
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Wyświetlanie obrazów

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()

Obraz chmury będący wynikiem plt.show

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Augmentacja danych

train_transforms = transforms.Compose([

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

Augmentacja danych: generowanie dodatkowych danych przez losowe transformacje oryginalnych obrazów

  • Zwiększenie rozmiaru i różnorodności zbioru treningowego
  • Poprawa odporności modelu
  • Redukcja przeuczenia

Trzy obrazy chmur przedstawiające transformację obrotu

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Czas na ćwiczenia!

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Preparing Video For Download...