Trénování a vyhodnocování RNN

Intermediate Deep Learning with PyTorch

Michal Oleszak

Machine Learning Engineer

Ztrátová funkce MSE

  • Chyba:

    $$prediction - target$$

  • Kvadratická chyba:

    $$(prediction - target)^2$$

  • Střední kvadratická chyba:

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

Umocnění chyby:

  • Zabraňuje vzájemnému rušení kladných a záporných chyb
  • Více penalizuje velké chyby
  • V PyTorchi:
      criterion = nn.MSELoss()
    
Intermediate Deep Learning with PyTorch

Rozšiřování tenzorů

  • Rekurentní vrstvy očekávají vstup tvaru (batch_size, seq_length, num_features)
  • Máme (batch_size, seq_length)
  • Je nutné přidat jednu dimenzi na konec
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])
Intermediate Deep Learning with PyTorch

Zmenšování tenzorů

  • Ve vyhodnocovací smyčce je nutné vrátit změnu tvaru provedenou v trénovací smyčce
  • Štítky mají tvar (batch_size)

    for seqs, labels in test_loader:
      print(labels.shape)
    
    torch.Size([32])
    
  • Výstupy modelu mají tvar (batch_size, 1)

    out = net(seqs)
    
    torch.Size([32, 1])
    
  • Tvary výstupů modelu a štítků musí pro ztrátovou funkci odpovídat
  • Poslední dimenzi výstupů modelu lze odstranit

    out = net(seqs).squeeze()
    
    torch.Size([32])
    
Intermediate Deep Learning with PyTorch

Trénovací smyčka

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()
  • Vytvoření modelu, definice ztrátové funkce a optimizátoru
  • Iterace přes epochy a dávky dat
  • Změna tvaru vstupní sekvence
  • Zbytek: jako obvykle
Intermediate Deep Learning with PyTorch

Vyhodnocovací smyčka

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
  • Nastavení metriky MSE
  • Iterace přes testovací data bez gradientů
  • Změna tvaru vstupů modelu
  • Zmenšení výstupů modelu
  • Aktualizace metriky
  • Výpočet výsledné hodnoty metriky
Intermediate Deep Learning with PyTorch

LSTM vs. GRU

  • LSTM:
Test MSE: 0.13292162120342255
  • GRU:
Test MSE: 0.12187089771032333
  • GRU je výhodnější: stejné nebo lepší výsledky s nižšími nároky na výpočetní výkon
Intermediate Deep Learning with PyTorch

Pojďme si procvičit!

Intermediate Deep Learning with PyTorch

Preparing Video For Download...