Modèles à entrées multiples

Apprentissage profond intermédiaire avec PyTorch

Michal Oleszak

Machine Learning Engineer

Pourquoi des entrées multiples ?

Exploiter plus d'information

Schéma d'un modèle qui prend deux images de voiture en entrée et produit une sortie.

Modèles multimodaux

Schéma d'un modèle qui prend une image et un texte en entrée et produit un texte en sortie.

Apprentissage métrique

Schéma d'un modèle qui prend deux images de visage en entrée et prédit s'il s'agit de la même personne.

Apprentissage auto-supervisé

Schéma d'un modèle qui prend deux versions augmentées d'une même image en entrée et apprend qu'elles sont équivalentes.

Apprentissage profond intermédiaire avec PyTorch

Ensemble de données Omniglot

Échantillon d'images de l'ensemble de données Omniglot.

1 Lake, B. M., Salakhutdinov, R., and Tenenbaum, J. B. (2015). Human-level concept learning through probabilistic program induction. Science, 350(6266), 1332-1338.
Apprentissage profond intermédiaire avec PyTorch

Classification de caractères

Schéma du modèle : des images de caractères sont passées à un réseau de neurones.

Apprentissage profond intermédiaire avec PyTorch

Classification de caractères

Schéma du modèle : un vecteur « one-hot » d'alphabet est passé à un réseau de neurones.

Apprentissage profond intermédiaire avec PyTorch

Classification de caractères

Schéma du modèle : les plongements du caractère et de l'alphabet sont combinés.

Apprentissage profond intermédiaire avec PyTorch

Classification de caractères

Schéma du modèle : à partir des plongements combinés, un classificateur prédit le caractère.

Apprentissage profond intermédiaire avec PyTorch

Jeu de données à deux entrées

from PIL import Image

class OmniglotDataset(Dataset):

def __init__(self, transform, samples): self.transform = transform self.samples = samples
def __len__(self): return len(self.samples)
def __getitem__(self, idx): img_path, alphabet, label = self.samples[idx] img = Image.open(img_path).convert('L') img = self.transform(img) return img, alphabet, label
  • Affecter les échantillons et les transformations

    print(samples[0])
    
    [(
      'omniglot_train/.../0459_14.png',
       array([1., 0., 0., ..., 0., 0., 0.]),
       0
     )]
    
  • Implémenter __len__()

  • Charger et transformer l'image

  • Retourner les deux entrées et l'étiquette
Apprentissage profond intermédiaire avec PyTorch

Concaténation de tenseurs

x = torch.tensor([
  [1, 2, 3],
])

y = torch.tensor([
  [4, 5, 6],
])

Concaténation selon l'axe 0

torch.cat((x, y), dim=0)
[[1, 2, 3],
 [4, 5, 6]]

Concaténation selon l'axe 1

torch.cat((x, y), dim=1)
[[1, 2, 3, 4, 5, 6]]
Apprentissage profond intermédiaire avec PyTorch

Architecture à deux entrées

class Net(nn.Module):
    def __init__(self):
        super().__init__()

self.image_layer = nn.Sequential( nn.Conv2d(1, 16, kernel_size=3, padding=1), nn.MaxPool2d(kernel_size=2), nn.ELU(), nn.Flatten(), nn.Linear(16*32*32, 128) )
self.alphabet_layer = nn.Sequential( nn.Linear(30, 8), nn.ELU(), )
self.classifier = nn.Sequential( nn.Linear(128 + 8, 964), )
  • Définir la couche de traitement d'image
  • Définir la couche de traitement de l'alphabet
  • Définir la couche du classificateur
Apprentissage profond intermédiaire avec PyTorch

Architecture à deux entrées

def forward(self, x_image, x_alphabet):

x_image = self.image_layer(x_image)
x_alphabet = self.alphabet_layer(x_alphabet)
x = torch.cat((x_image, x_alphabet), dim=1)
return self.classifier(x)
  • Faire passer l'image par la couche d'image
  • Faire passer l'alphabet par sa couche
  • Concaténer les sorties image et alphabet
  • Faire passer le résultat au classificateur
Apprentissage profond intermédiaire avec PyTorch

Boucle d'entraînement

net = Net()
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.01)

for epoch in range(10):
    for img, alpha, labels in dataloader_train:
        optimizer.zero_grad()
        outputs = net(img, alpha)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
  • Les données d'entraînement contiennent trois éléments :
    • Image
    • Vecteur d'alphabet
    • Étiquettes
  • On transmet au modèle les images et les alphabets
Apprentissage profond intermédiaire avec PyTorch

Passons à la pratique !

Apprentissage profond intermédiaire avec PyTorch

Preparing Video For Download...