Kom igångKom igång gratis

Textgenerering med RNN – träning och generering

Teamet på PyBooks vill nu att du tränar och testar RNN-modellen, som är utformad för att förutsäga nästa tecken i sekvensen baserat på angiven indata – för automatisk komplettering av boknamn. Projektet hjälper teamet att vidareutveckla modeller för textkomplettering.

model-instansen för klassen RNNmodel är förinstansierad åt dig. Variabeln data har förbehandlats och kodats som en sekvens.

Variablerna inputs och targets är förinstansierade åt dig.

Den här övningen är en del av kursen

Djupinlärning för text med PyTorch

Visa kurs

Övningsinstruktioner

  • Instansiera förlustfunktionen som används för att beräkna modellens fel.
  • Instansiera optimeraren från PyTorchs optimeringsmodul.
  • Kör modellträningen genom att ställa in modellen i träningsläge och nollställa gradienterna innan du utför ett optimeringssteg.
  • Växla modellen till utvärderingsläge efter träningen för att testa den på ett exempelindata.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

# 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]}'")
Redigera och kör kod