Ocena modeli wielowyjściowych i ważenie strat

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

Michal Oleszak

Machine Learning Engineer

Ocena modelu

acc_alpha = Accuracy(
    task="multiclass", num_classes=30
)
acc_char = Accuracy(
    task="multiclass", num_classes=964
)


net.eval() with torch.no_grad(): for images, labels_alpha, labels_char \ in dataloader_test: out_alpha, out_char = net(images)
_, pred_alpha = torch.max(out_alpha, 1) _, pred_char = torch.max(out_char, 1)
acc_alpha(pred_alpha, labels_alpha) acc_char(pred_char, labels_char)
  • Zdefiniowanie metryki dla każdego wyjścia
  • Iteracja po zbiorze testowym i uzyskanie wyjść
  • Obliczenie predykcji dla każdego wyjścia
  • Aktualizacja metryk dokładności
  • Obliczenie końcowych wyników dokładności
print(f"Alphabet: {acc_alpha.compute()}")
print(f"Character: {acc_char.compute()}")
Alphabet: 0.3166305720806122
Character: 0.24064336717128754
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Pętla treningowa dla modelu wielowyjściowego

for epoch in range(10):
    for images, labels_alpha, labels_char \
    in dataloader_train:
        optimizer.zero_grad()
        outputs_alpha, outputs_char = net(images)
        loss_alpha = criterion(
          outputs_alpha, labels_alpha
        )
        loss_char = criterion(
          outputs_char, labels_char
        )
        loss = loss_alpha + loss_char
        loss.backward()
        optimizer.step()
  • Dwie straty: dla alfabetów i znaków
  • Łączna strata jako suma strat alfabetu i znaków: loss = loss_alpha + loss_char
  • Oba zadania klasyfikacji traktowane równorzędnie
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Różne znaczenie zadań

Klasyfikacja znaków 2 razy ważniejsza niż klasyfikacja alfabetu

  • Podejście 1: Skalowanie ważniejszej straty o współczynnik 2

    loss = loss_alpha + loss_char * 2
    
  • Podejście 2: Przypisanie wag sumujących się do 1

    loss = 0.33 * loss_alpha + 0.67 * loss_char
    
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Uwaga: straty na różnych skalach

  • Straty muszą być w tej samej skali przed ważeniem i sumowaniem
  • Przykładowe zadania:

    • Predykcja ceny domu -> strata MSE
    • Predykcja jakości: niska, średnia, wysoka -> strata CrossEntropy
  • CrossEntropy zazwyczaj mieści się w zakresie jednocyfrowym

  • Strata MSE może osiągać dziesiątki tysięcy
  • Model ignorowałby zadanie oceny jakości
  • Rozwiązanie: Normalizacja obu strat przed ważeniem i sumowaniem
    loss_price = loss_price / torch.max(loss_price)
    loss_quality = loss_quality / torch.max(loss_quality)
    loss = 0.7 * loss_price + 0.3 * loss_quality
    
Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Czas na ćwiczenia!

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

Preparing Video For Download...