Generování textu pomocí RNN – trénování a generování
Tým PyBooks teď potřebuje natrénovat a otestovat model RNN, který je navržený tak, aby předpovídal další znak v sekvenci na základě zadaného vstupu – a umožnil tak automatické doplňování názvů knih. Tento projekt pomůže týmu dále rozvíjet modely pro doplňování textu.
Instance model třídy RNNmodel je pro tebe předem načtena. Proměnná data byla předzpracována a zakódována jako sekvence.
Proměnné inputs a targets jsou také předem načteny.
Toto cvičení je součástí kurzu
Deep Learning for Text with PyTorch
Pokyny k cvičení
- Vytvoř instanci funkce ztráty, která se použije k výpočtu chyby modelu.
- Vytvoř instanci optimalizátoru z optimalizačního modulu PyTorche.
- Spusť trénování modelu: přepni ho do trénovacího režimu a před krokem optimalizace vynuluj gradienty.
- Po dokončení trénování přepni model do režimu vyhodnocení a otestuj ho na vzorovém vstupu.
Interaktivní cvičení na vyzkoušení si v praxi
Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.
# 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]}'")