Xây dựng mô hình LSTM cho văn bản
Tại PyBooks, đội ngũ luôn nỗ lực nâng cao trải nghiệm người dùng bằng cách tận dụng các tiến bộ công nghệ mới nhất. Theo định hướng đó, họ đã giao cho bạn một nhiệm vụ quan trọng. Nhóm muốn bạn khám phá tiềm năng của một công cụ mạnh mẽ khác: LSTM, nổi tiếng với khả năng nắm bắt các mẫu dữ liệu phức tạp hơn. Bạn sẽ tiếp tục làm việc với cùng bộ dữ liệu Newsgroup, với mục tiêu không đổi: phân loại các bài báo thành ba nhóm riêng biệt:
rec.autos, sci.med, và comp.graphics.
Các gói sau đã được nạp sẵn cho bạn: torch, nn, optim.
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
- Thiết lập mô hình LSTM bằng cách hoàn thiện các lớp LSTM và tuyến tính với các tham số cần thiết.
- Khởi tạo mô hình với các tham số cần thiết.
- Huấn luyện mô hình LSTM bằng cách đặt lại gradient về 0 và truyền dữ liệu đầu vào
X_train_seqqua mô hình. - Tính loss dựa trên
outputsdự đoán và nhãn thật.
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.
# Initialize the LSTM and the output layer with parameters
class LSTMModel(nn.Module):
def __init__(self, input_size, hidden_size, num_layers, num_classes):
super(LSTMModel, self).__init__()
self.hidden_size = hidden_size
self.num_layers = num_layers
self.lstm = nn.LSTM(____, ____, ____, batch_first=True)
self.fc = nn.Linear(____, ____)
def forward(self, x):
h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size)
c0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size)
out, _ = self.lstm(x, (h0, c0))
out = out[:, -1, :]
out = self.fc(out)
return out
# Initialize model with required parameters
lstm_model = LSTMModel(____, ____, ____, ____)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(lstm_model.parameters(), lr=0.01)
# Train the model by passing the correct parameters and zeroing the gradient
for epoch in range(10):
optimizer.____
outputs = lstm_model(____)
loss = criterion(____, y_train_seq)
loss.backward()
optimizer.step()
print(f'Epoch: {epoch+1}, Loss: {loss.item()}')