Trenowanie modelu CNN do klasyfikacji tekstu
Świetna robota z definiowaniem klasy TextClassificationCNN. PyBooks musi teraz wytrenować model, aby zoptymalizować go pod kątem trafnej analizy sentymentu recenzji książek.
Następujące pakiety zostały już zaimportowane:
torch, torch.nn jako nn, torch.nn.functional jako F, torch.optim jako optim.
Instancja klasy TextClassificationCNN() z argumentami vocab_size i embed_dim została również załadowana i zapisana jako model.
To ćwiczenie jest częścią kursu
Uczenie głębokie dla tekstu z PyTorch
Instrukcje do ćwiczenia
- Zdefiniuj funkcję straty używaną do klasyfikacji binarnej i zapisz ją jako
criterion. - Wyzeruj gradienty na początku pętli treningowej.
- Zaktualizuj parametry na końcu pętli.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# 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!')