Entraîner et tester le modèle Transformer
Avec le modèle TransformerEncoder en place, la prochaine étape chez PyBooks est d'entraîner le modèle sur des critiques d'exemple et d'évaluer sa performance. L'entraînement sur ces critiques aidera PyBooks à comprendre les tendances de sentiment dans leur vaste répertoire. En obtenant un modèle performant, PyBooks pourra ensuite automatiser l'analyse de sentiment, afin que les lecteurs reçoivent des recommandations et des commentaires éclairants.
Les modules suivants ont été importés pour vous : torch, nn, optim.
L'instance model de la classe TransformerEncoder, token_embeddings, ainsi que train_sentences, train_labels, test_sentences, test_labels sont déjà chargés pour vous.
Cette activité fait partie du cours
Apprentissage profond pour le texte avec PyTorch
Instructions de l’exercice
- Dans la boucle d'entraînement, divisez les phrases en jetons et empilez les embeddings.
- Remettez les gradients à zéro et effectuez une passe arrière.
- Dans la fonction
predict, désactivez les calculs de gradient, puis obtenez la prédiction de sentiment.
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
for epoch in range(5):
for sentence, label in zip(train_sentences, train_labels):
# Split the sentences into tokens and stack the embeddings
tokens = ____
data = torch.____([token_embeddings[token] for token in ____], dim=1)
output = model(data)
loss = criterion(output, torch.tensor([label]))
# Zero the gradients and perform a backward pass
optimizer.____()
loss.____()
optimizer.step()
print(f"Epoch {epoch}, Loss: {loss.item()}")
def predict(sentence):
model.eval()
# Deactivate the gradient computations and get the sentiment prediction.
with torch.____():
tokens = sentence.split()
data = torch.stack([token_embeddings.get(token, torch.rand((1, 512))) for token in tokens], dim=1)
output = model(data)
predicted = torch.____(output, dim=1)
return "Positive" if predicted.item() == 1 else "Negative"
sample_sentence = "This product can be better"
print(f"'{sample_sentence}' is {predict(sample_sentence)}")