Trening i ocena RNN

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Michal Oleszak

Machine Learning Engineer

Funkcja straty MSE

  • Błąd:

    $$prediction - target$$

  • Błąd kwadratowy:

    $$(prediction - target)^2$$

  • Średni błąd kwadratowy:

    $$avg[(prediction - target)^2]$$

Potęgowanie błędu:

  • Zapobiega wzajemnemu znoszeniu się błędów dodatnich i ujemnych
  • Mocniej penalizuje duże błędy
  • W PyTorch:
      criterion = nn.MSELoss()
    
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Rozszerzanie tensorów

  • Warstwy rekurencyjne oczekują kształtu wejścia (batch_size, seq_length, num_features)
  • Posiadamy (batch_size, seq_length)
  • Należy dodać jeden wymiar na końcu
for seqs, labels in dataloader_train:
    print(seqs.shape)
torch.Size([32, 96])
seqs = seqs.view(32, 96, 1)
print(seqs.shape)
torch.Size([32, 96, 1])
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Ściskanie tensorów

  • W pętli ewaluacyjnej należy cofnąć zmianę kształtu z pętli treningowej
  • Etykiety mają kształt (batch_size)

    for seqs, labels in test_loader:
      print(labels.shape)
    
    torch.Size([32])
    
  • Wyjścia modelu mają kształt (batch_size, 1)

    out = net(seqs)
    
    torch.Size([32, 1])
    
  • Kształty wyjść modelu i etykiet muszą być zgodne dla funkcji straty
  • Można usunąć ostatni wymiar z wyjść modelu

    out = net(seqs).squeeze()
    
    torch.Size([32])
    
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Pętla treningowa

net = Net()
criterion = nn.MSELoss()
optimizer = optim.Adam(
  net.parameters(), lr=0.001
)


for epoch in range(num_epochs): for seqs, labels in dataloader_train:
seqs = seqs.view(32, 96, 1)
outputs = net(seqs) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step()
  • Inicjalizacja modelu, definicja straty i optymalizatora
  • Iteracja po epokach i batchach danych
  • Zmiana kształtu sekwencji wejściowej
  • Reszta: jak zwykle
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Pętla ewaluacyjna

mse = torchmetrics.MeanSquaredError()


net.eval() with torch.no_grad(): for seqs, labels in test_loader:
seqs = seqs.view(32, 96, 1)
outputs = net(seqs).squeeze()
mse(outputs, labels)
print(f"Test MSE: {mse.compute()}")
Test MSE: 0.13292162120342255
  • Konfiguracja metryki MSE
  • Iteracja po danych testowych bez gradientów
  • Zmiana kształtu wejść modelu
  • Ściśnięcie wyjść modelu
  • Aktualizacja metryki
  • Obliczenie końcowej wartości metryki
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

LSTM vs. GRU

  • LSTM:
Test MSE: 0.13292162120342255
  • GRU:
Test MSE: 0.12187089771032333
  • GRU preferowane: takie same lub lepsze wyniki przy mniejszym zużyciu zasobów
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Czas na ćwiczenia!

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Preparing Video For Download...