Tworzenie modelu RNN z mechanizmem uwagi
Zespół PyBooks bada różne architektury głębokiego uczenia. Po przeprowadzeniu analizy postanawiasz zaimplementować sieć RNN z mechanizmem uwagi (ang. Attention), która będzie przewidywać następne słowo w zdaniu. Dysponujesz zbiorem danych zawierającym zdania oraz słownik utworzony na ich podstawie.
Następujące pakiety zostały już zaimportowane: torch, nn.
Następujące elementy są wstępnie załadowane:
vocabivocab_size: zbiór słownictwa i jego rozmiarword_to_ixiix_to_word: słowniki mapujące słowa na indeksy i indeksy na słowainput_dataitarget_data: zbiór danych przekształcony na pary wejście–wyjścieembedding_dimihidden_dim: wymiary osadzeń i ukrytego stanu RNN
Możesz sprawdzić zmienną data w konsoli, aby zobaczyć przykładowe zdania.
To ćwiczenie jest częścią kursu
Uczenie głębokie dla tekstu z PyTorch
Instrukcje do ćwiczenia
- Utwórz warstwę osadzeń dla słownictwa z podanym wymiarem
embedding_dim. - Zastosuj transformację liniową do sekwencji wyjściowej RNN, aby uzyskać wyniki uwagi.
- Wyznacz wagi uwagi na podstawie tych wyników.
- Oblicz wektor kontekstu jako ważoną sumę wyjść RNN i wag uwagi.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
class RNNWithAttentionModel(nn.Module):
def __init__(self):
super(RNNWithAttentionModel, self).__init__()
# Create an embedding layer for the vocabulary
self.embeddings = nn.____(vocab_size, embedding_dim)
self.rnn = nn.RNN(embedding_dim, hidden_dim, batch_first=True)
# Apply a linear transformation to get the attention scores
self.attention = nn.____(____, 1)
self.fc = nn.____(hidden_dim, vocab_size)
def forward(self, x):
x = self.embeddings(x)
out, _ = self.rnn(x)
# Get the attention weights
attn_weights = torch.nn.functional.____(self.____(out).____(2), dim=1)
# Compute the context vector
context = torch.sum(____.____(2) * out, dim=1)
out = self.fc(context)
return out
attention_model = RNNWithAttentionModel()
optimizer = torch.optim.Adam(attention_model.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss()
print("Model Instantiated")