Utiliser des modèles préentraînés

Deep Learning pour les images avec PyTorch

Michal Oleszak

Machine Learning Engineer

Exploiter des modèles préentraînés

  • Entraîner des modèles à partir de zéro :

    • Processus long
    • Exige beaucoup de données
  • Modèles préentraînés – modèles déjà entraînés sur une tâche

    • Réutilisables directement sur une nouvelle tâche
    • Nécessitent un ajustement à la nouvelle tâche (apprentissage par transfert)
  • Étapes pour exploiter des modèles préentraînés :

    • Enregistrer et charger des modèles localement
    • Télécharger des modèles torchvision
Deep Learning pour les images avec PyTorch

Enregistrer un modèle PyTorch complet

  • torch.save()
  • Extension du modèle : .pt ou .pth
  • Enregistrer les poids avec .state_dict()
    torch.save(model.state_dict(), "BinaryCNN.pth")
    
Deep Learning pour les images avec PyTorch

Charger des modèles PyTorch

  • Instancier un nouveau modèle

    new_model = BinaryCNN()
    
  • Charger les paramètres enregistrés

    new_model.load_state_dict(torch.load('BinaryCNN.pth'))
    
Deep Learning pour les images avec PyTorch

Télécharger des modèles torchvision

from torchvision.models import (
    resnet18, ResNet18_Weights
)


weights = ResNet18_Weights.DEFAULT
model = resnet18(weights=weights)
transforms = weights.transforms()
  • Importer l'architecture resnet et les poids
  • Extraire les poids
  • Instancier un modèle en lui passant les poids
  • Conserver les transformations de données requises
Deep Learning pour les images avec PyTorch

Préparer de nouvelles images d'entrée

from PIL import Image

image = Image.open("cat013.jpg")

image_tensor = transform(image)
image_reshaped = image_tensors.unsqueeze(0)

 

image de chat

  • Charger l'image
  • Transformer l'image
  • Remodeler l'image
Deep Learning pour les images avec PyTorch

Générer une nouvelle prédiction

model.eval()


with torch.no_grad():
pred = model(image_reshaped).squeeze(0)
pred_cls = pred.softmax(0)
cls_id = pred_cls.argmax().item()
cls_name = weights.meta["categories"][cls_id]
print(cls_name)
Egyptian cat
  • Mode évaluation pour l'inférence
  • Désactiver les gradients
  • Passer l'image au modèle et retirer la dimension lot
  • Appliquer softmax
  • Choisir la classe la plus probable et extraire son indice
  • Faire correspondre l'indice de classe à l'étiquette
  • Afficher l'étiquette de classe
Deep Learning pour les images avec PyTorch

Passons à la pratique !

Deep Learning pour les images avec PyTorch

Preparing Video For Download...