Gestion des images avec PyTorch

Apprentissage profond intermédiaire avec PyTorch

Michal Oleszak

Machine Learning Engineer

Ensemble de données « Clouds »

Échantillons de l'ensemble de données sur les nuages : cinq images montrant différents types de nuages.

1 https://www.kaggle.com/competitions/cloud-type-classification2/data
Apprentissage profond intermédiaire avec PyTorch

Qu'est-ce qu'une image ?

Image de nuage avec une partie agrandie laissant voir les pixels.

  • Une image est faite de pixels (« éléments d'image »)
  • Chaque pixel contient des infos de couleur

  • Niveaux de gris : entier de 0 à 255

    • 30 :

Carré gris

  • Images en couleur : trois entiers, un par canal (rouge, vert, bleu)
    • RGB = (52, 171, 235) :

Carré bleu

Apprentissage profond intermédiaire avec PyTorch

Charger des images dans PyTorch

Structure de répertoires souhaitée :

clouds_train

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

 

  • Dossiers principaux : clouds_train et clouds_test
  • Dans chaque dossier principal : un dossier par catégorie
  • Dans chaque dossier de classe : des fichiers image
Apprentissage profond intermédiaire avec PyTorch

Charger des images dans 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, )
  • Définir les transformations :

    • Convertir en tenseur
    • Redimensionner à 128x128
  • Créer l'ensemble de données en passant :

    • Le chemin des données
    • Les transformations définies
Apprentissage profond intermédiaire avec PyTorch

Afficher des images

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

Image de nuage affichée par plt.show

Apprentissage profond intermédiaire avec PyTorch

Augmentation des données

train_transforms = transforms.Compose([

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

Augmentation des données : créer plus de données en appliquant des transformations aléatoires aux images originales

  • Accroître la taille et la diversité de l'ensemble d'entraînement
  • Améliorer la robustesse du modèle
  • Réduire le surapprentissage

Trois images de nuages illustrant la rotation

Apprentissage profond intermédiaire avec PyTorch

Passons à la pratique !

Apprentissage profond intermédiaire avec PyTorch

Preparing Video For Download...