ÎncepețiÎncepe gratuit

Generarea de text cu RNN – Antrenare și generare

Echipa PyBooks îți cere acum să antrenezi și să testezi modelul RNN, conceput pentru a prezice următorul caracter dintr-o secvență, pe baza datelor de intrare furnizate, în vederea completării automate a numelor de cărți. Acest proiect va ajuta echipa să dezvolte în continuare modele pentru completarea textului.

Instanța model pentru clasa RNNmodel este deja încărcată pentru tine. Variabila data a fost preprocesată și encodată ca o secvență.

Variabilele inputs și targets sunt deja încărcate pentru tine.

Acest exercițiu face parte din cursul

Deep Learning pentru text cu PyTorch

Vezi cursul

Instrucțiuni pentru exercițiu

  • Instanțiază funcția de pierdere care va fi folosită pentru a calcula eroarea modelului.
  • Instanțiază optimizatorul din modulul de optimizare al PyTorch.
  • Rulează procesul de antrenare a modelului, setând modelul în modul de antrenare și resetând gradienții la zero înainte de a efectua un pas de optimizare.
  • După procesul de antrenare, comută modelul în modul de evaluare pentru a-l testa pe un eșantion de intrare.

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# 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]}'")
Editează și rulează codul