Xây dựng mô hình RNN cho văn bản
Là một nhà phân tích dữ liệu tại PyBooks, bạn thường gặp các tập dữ liệu có thông tin tuần tự, như tương tác của khách hàng, dữ liệu chuỗi thời gian hoặc tài liệu văn bản. RNN có thể phân tích hiệu quả và rút ra insight từ những dữ liệu như vậy. Trong bài tập này, bạn sẽ làm việc với bộ dữ liệu Newsgroup đã được xử lý và mã hóa sẵn. Bộ dữ liệu này gồm các bài viết từ nhiều chủ đề khác nhau. Nhiệm vụ của bạn là áp dụng RNN để phân loại các bài viết vào ba nhóm:
rec.autos, sci.med, và comp.graphics.
Những mục sau đã được nạp sẵn cho bạn: torch, nn, optim.
Ngoài ra, các tham số input_size, hidden_size (32), num_layers (2), và num_classes cũng đã được nạp sẵn.
Bài này và các bài tiếp theo sử dụng bộ dữ liệu fetch_20newsgroups từ sklearn.
Bài tập này là một phần của khóa học
Deep Learning cho Văn bản với PyTorch
Hướng dẫn bài tập
- Hoàn thiện lớp RNN với một lớp RNN và một lớp fully connected (Linear).
- Khởi tạo mô hình.
- Huấn luyện mô hình RNN trong mười epoch bằng cách đưa gradient về 0 mỗi vòng.
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
# Complete the RNN class
class RNNModel(nn.Module):
def __init__(self, input_size, hidden_size, num_layers, num_classes):
super(RNNModel, self).__init__()
self.hidden_size = hidden_size
self.num_layers = num_layers
self.rnn = ____.____(input_size, hidden_size, num_layers, batch_first=True)
self.fc = ____.____(hidden_size, num_classes)
def forward(self, x):
h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size)
out, _ = self.rnn(x, h0)
out = out[:, -1, :]
out = self.fc(out)
return out
# Initialize the model
rnn_model = ____(input_size, hidden_size, num_layers, num_classes)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(rnn_model.parameters(), lr=0.01)
# Train the model for ten epochs and zero the gradients
for epoch in ____:
optimizer.____()
outputs = ____(X_train_seq)
loss = criterion(outputs, y_train_seq)
loss.backward()
optimizer.step()
print(f'Epoch: {epoch+1}, Loss: {loss.item()}')