Обучение модели CNN для классификации текста
Отличная работа по определению класса TextClassificationCNN! Теперь PyBooks нужно обучить модель, чтобы оптимизировать её для точного анализа тональности отзывов о книгах.
Следующие пакеты уже импортированы:
torch, torch.nn как nn, torch.nn.functional как F, torch.optim как optim.
Экземпляр класса TextClassificationCNN() с аргументами vocab_size и embed_dim также уже создан и сохранён как model.
Это упражнение является частью курса
Глубокое обучение для работы с текстом на PyTorch
Инструкции к упражнению
- Определите функцию потерь для бинарной классификации и сохраните её как
criterion. - Обнулите градиенты в начале цикла обучения.
- Обновите параметры модели в конце цикла.
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# 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!')