Trénování a testování RNN modelu s attention
Tým v PyBooks dříve sestavil RNN model pro predikci slov bez mechanismu attention. Tento původní model, označovaný jako rnn_model, je již natrénovaný a jeho instance je předem načtena. Tvým úkolem je nyní natrénovat nový RNNWithAttentionModel a porovnat jeho predikce s výsledky staršího rnn_model.
Následující objekty jsou pro tebe předem načteny:
inputs: seznam vstupních sekvencí jako tensorytargets: tensor obsahující cílová slova pro každou vstupní sekvencioptimizer: optimalizátor Adamcriterion: funkce CrossEntropyLosspad_sequences: funkce pro zarovnání vstupních sekvencí pro dávkové zpracováníattention_model: třída modelu definovaná v předchozím cvičenírnn_model: natrénovaný RNN model od týmu PyBooks
Toto cvičení je součástí kurzu
Deep Learning for Text with PyTorch
Pokyny k cvičení
- Před testováním s testovacími daty přepni RNN model do režimu vyhodnocení.
- Získej výstup RNN předáním příslušného vstupu do RNN modelu.
- Z výstupu RNN vyber slovo s nejvyšším skóre predikce.
- Stejným způsobem vyber z výstupu attention modelu slovo s nejvyšším skóre predikce.
Interaktivní cvičení na vyzkoušení si v praxi
Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.
for epoch in range(epochs):
attention_model.train()
optimizer.zero_grad()
padded_inputs = pad_sequences(inputs)
outputs = attention_model(padded_inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
for input_seq, target in zip(input_data, target_data):
input_test = torch.tensor(input_seq, dtype=torch.long).unsqueeze(0)
# Set the RNN model to evaluation mode
rnn_model.____()
# Get the RNN output by passing the appropriate input
rnn_output = ____(____)
# Extract the word with the highest prediction score
rnn_prediction = ix_to_word[torch.____(____).item()]
attention_model.eval()
attention_output = attention_model(input_test)
# Extract the word with the highest prediction score
attention_prediction = ix_to_word[torch.____(____).item()]
print(f"\nInput: {' '.join([ix_to_word[ix] for ix in input_seq])}")
print(f"Target: {ix_to_word[target]}")
print(f"RNN prediction: {rnn_prediction}")
print(f"RNN with Attention prediction: {attention_prediction}")