Zacznij terazZacznij za darmo

Trenowanie modelu

W tym ćwiczeniu wytrenujesz wcześniej zaimplementowany model. Czy wiesz, że model tłumaczenia maszynowego oparty na architekturze enkoder-dekoder firmy Google wymagał 2–4 dni trenowania?

W tym ćwiczeniu będziesz korzystać z małego zbioru danych zawierającego 1500 zdań (tj. en_text i fr_text). Taka ilość danych nie wystarczy do osiągnięcia wysokiej jakości tłumaczeń, ale metoda pozostaje taka sama – chodzi o trenowanie na większej ilości danych przez dłuższy czas. Masz też dostęp do modelu nmt oraz funkcji sents2seqs(), którą zaimplementowałeś wcześniej. W tym ćwiczeniu odwrócimy tekst wejściowy enkodera, aby poprawić jakość modelu. Tutaj en_x oznacza dane wejściowe enkodera, natomiast de_x oznacza dane wejściowe dekodera.

To ćwiczenie jest częścią kursu

Tłumaczenie maszynowe z Keras

Zobacz kurs

Instrukcje do ćwiczenia

  • Pobierz pojedynczą partię danych wejściowych enkodera (zdania angielskie od indeksu i do i+bsize) przy użyciu funkcji sents2seqs(). Dane wejściowe muszą być odwrócone i zakodowane metodą one-hot.
  • Pobierz pojedynczą partię danych wyjściowych dekodera (zdania francuskie od indeksu i do i+bsize) przy użyciu funkcji sents2seqs(). Dane wejściowe muszą być zakodowane metodą one-hot.
  • Wytrenuj model na pojedynczej partii danych zawierającej en_x i de_y.
  • Wyznacz metryki ewaluacji dla en_x i de_y, oceniając model z parametrem batch_size ustawionym na bsize.

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

n_epochs, bsize = 3, 250

for ei in range(n_epochs):
  for i in range(0,data_size,bsize):
    # Get a single batch of encoder inputs
    en_x = ____('source', ____, onehot=____, reverse=____)
    # Get a single batch of decoder outputs
    de_y = sents2seqs('target', fr_text[____], onehot=____)
    
    # Train the model on a single batch of data
    nmt.____(____, ____)    
    # Obtain the eval metrics for the training data
    res = nmt.____(____, de_y, batch_size=____, verbose=0)
    print("{} => Train Loss:{}, Train Acc: {}".format(ei+1,res[0], res[1]*100.0))  
Edytuj i uruchom kod