Skapa en RNN-modell med attention
På PyBooks har teamet utforskat olika djupinlärningsarkitekturer. Efter lite efterforskningar bestämmer du dig för att implementera ett RNN med en attentionsmekanism för att förutsäga nästa ord i en mening. Du får en datamängd med meningar och ett vokabulär skapat utifrån dem.
Följande paket har importerats åt dig: torch, nn.
Följande har förinslästs åt dig:
vocabochvocab_size: vokabulärmängden och dess storlekword_to_ixochix_to_word: ordlistor för mappning från ord till index och från index till ordinput_dataochtarget_data: datamängden konverterad till indata-utdata-parembedding_dimochhidden_dim: dimensioner för inbäddning och RNN:ens dolda tillstånd
Du kan inspektera variabeln data i konsolen för att se exempelmeningarna.
Den här övningen är en del av kursen
Djupinlärning för text med PyTorch
Övningsinstruktioner
- Skapa ett inbäddningslager för vokabuläret med det angivna
embedding_dim. - Tillämpa en linjär transformation på RNN:ens sekvensutdata för att beräkna attentionspoängen.
- Hämta attentionsvikterna från poängen.
- Beräkna kontextvektorn som den viktade summan av RNN-utdata och attentionsvikter.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
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")