Träna en CNN-modell för text
Bra jobbat med att definiera klassen TextClassificationCNN. PyBooks behöver nu träna modellen för att optimera den för korrekt sentimentanalys av bokrecensioner.
Följande paket har importerats åt dig:
torch, torch.nn som nn, torch.nn.functional som F, torch.optim som optim.
En instans av TextClassificationCNN() med argumenten vocab_size och embed_dim har också laddats och sparats som model.
Den här övningen är en del av kursen
Djupinlärning för text med PyTorch
Övningsinstruktioner
- Definiera en förlustfunktion för binär klassificering och spara den som
criterion. - Nollställ gradienterna i början av träningsloopen.
- Uppdatera parametrarna i slutet av loopen.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
# Define the loss function
criterion = nn.____()
optimizer = optim.SGD(model.parameters(), lr=0.1)
for epoch in range(10):
for sentence, label in data:
# Clear the gradients
model.____()
sentence = torch.LongTensor([word_to_ix.get(w, 0) for w in sentence]).unsqueeze(0)
label = torch.LongTensor([int(label)])
outputs = model(sentence)
loss = criterion(outputs, label)
loss.backward()
# Update the parameters
____.____()
print('Training complete!')