Użycie funkcji straty do oceny prognoz modelu

Wprowadzenie do uczenia głębokiego z PyTorch

Jasmin Ludolf

Senior Data Science Content Developer, DataCamp

Po co jest funkcja straty?

  • Mierzy jakość modelu podczas trenowania
  • Przyjmuje prognozę $\hat{y}$ i prawdziwą etykietę $y$
  • Zwraca wartość float

$$

Diagram funkcji straty

Wprowadzenie do uczenia głębokiego z PyTorch

Po co jest funkcja straty?

  • Klasa 0 – ssak, klasa 1 – ptak, klasa 2 – gad
Hair Feathers Eggs Milk Fins Legs Tail Domestic Catsize Class
1 0 0 1 0 4 0 0 1 0

$$

  • Prognozowana klasa = 0 -> poprawna = niska strata
  • Prognozowana klasa = 1 -> błędna = wysoka strata
  • Prognozowana klasa = 2 -> błędna = wysoka strata

$$

  • Celem jest minimalizacja straty
Wprowadzenie do uczenia głębokiego z PyTorch

Koncepcja kodowania one-hot

  • $loss = F(y, \hat{y})$
  • $y$ to pojedyncza liczba całkowita (etykieta klasy)
    • np. $y=0$, gdy $y$ to ssak
  • $\hat{y}$ to tensor (prognoza przed softmax)
    • Jeśli N to liczba klas, np. N = 3
    • $\hat{y}$ jest tensorem o N wymiarach,
      • np. $\hat{y}$ = [-5.2, 4.6, 0.8]
Wprowadzenie do uczenia głębokiego z PyTorch

Koncepcja kodowania one-hot

  • Konwertuje liczbę całkowitą y na tensor zer i jedynek

Kodowanie one-hot

Wprowadzenie do uczenia głębokiego z PyTorch

Przekształcanie etykiet za pomocą kodowania one-hot

import torch.nn.functional as F

print(F.one_hot(torch.tensor(0), num_classes = 3))
tensor([1, 0, 0])
print(F.one_hot(torch.tensor(1), num_classes = 3))
tensor([0, 1, 0])
print(F.one_hot(torch.tensor(2), num_classes = 3))
tensor([0, 0, 1])
Wprowadzenie do uczenia głębokiego z PyTorch

Entropia krzyżowa w PyTorch

from torch.nn import CrossEntropyLoss

scores = torch.tensor([-5.2, 4.6, 0.8])
one_hot_target = torch.tensor([1, 0, 0])

criterion = CrossEntropyLoss()
print(criterion(scores.double(), one_hot_target.double()))

$$

tensor(9.8222, dtype=torch.float64)
Wprowadzenie do uczenia głębokiego z PyTorch

Podsumowanie

Funkcja straty przyjmuje:

  • scores – prognozy modelu przed końcową funkcją softmax
  • one_hot_target – etykieta docelowa zakodowana metodą one-hot

Funkcja straty zwraca:

  • loss – pojedynczą wartość float

Diagram funkcji straty z wartościami

Wprowadzenie do uczenia głębokiego z PyTorch

Czas na ćwiczenia!

Wprowadzenie do uczenia głębokiego z PyTorch

Preparing Video For Download...