Bắt đầu ngayBắt đầu miễn phí

Sinh văn bản bằng RNN - Huấn luyện và Sinh

Nhóm PyBooks muốn bạn huấn luyện và kiểm thử mô hình RNN, mô hình này được thiết kế để dự đoán ký tự tiếp theo trong chuỗi dựa trên đầu vào nhằm tự động hoàn thành tên sách. Dự án này sẽ giúp nhóm phát triển thêm các mô hình để hoàn thành văn bản.

Instance model của lớp RNNmodel đã được nạp sẵn cho bạn. Biến data đã được tiền xử lý và mã hóa thành một chuỗi.

Các biến inputstargets cũng đã được nạp sẵn.

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

Xem khóa học

Hướng dẫn bài tập

  • Khởi tạo hàm mất mát sẽ dùng để tính sai số của mô hình.
  • Khởi tạo bộ tối ưu từ mô-đun tối ưu hóa của PyTorch.
  • Chạy quá trình huấn luyện bằng cách đặt mô hình ở chế độ train và đưa gradient về 0 trước khi thực hiện một bước tối ưu.
  • Sau quá trình huấn luyện, chuyển mô hình sang chế độ đánh giá để kiểm thử trên một đầu vào mẫu.

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.

# Instantiate the loss function
criterion = nn.____()
# Instantiate the optimizer
optimizer = torch.optim.____(model.parameters(), lr=0.01)

# Train the model
for epoch in range(100):
    model.____()
    outputs = model(inputs)
    loss = criterion(outputs, targets)
    optimizer.____()
    loss.backward()
    optimizer.step()
    if (epoch+1) % 10 == 0:
        print(f'Epoch {epoch+1}/100, Loss: {loss.item()}')

# Test the model
model.____()
test_input = char_to_ix['r']
test_input = nn.functional.one_hot(torch.tensor(test_input).view(-1, 1), num_classes=len(chars)).float()
predicted_output = model(test_input)
predicted_char_ix = torch.argmax(predicted_output, 1).item()
print(f"Test Input: 'r', Predicted Output: '{ix_to_char[predicted_char_ix]}'")
Chỉnh sửa và Chạy Mã