Definiowanie dekodera modelu inferencyjnego
Model inferencyjny to model używany w praktyce do wykonywania tłumaczeń na żądanie użytkownika. W tym ćwiczeniu zaimplementujesz dekoder modelu inferencyjnego.
Dekoder modelu inferencyjnego różni się od dekodera modelu treningowego. Nie możemy podawać dekoderowi francuskich słów, bo właśnie to chcemy przewidywać. Na szczęście istnieje rozwiązanie: możemy wykorzystać przewidziane w poprzednim kroku czasowym słowo francuskie jako wejście dla dekodera modelu inferencyjnego. Dlatego podczas generowania tłumaczenia dekoder musi produkować jedno słowo naraz, przyjmując poprzednie wyjście jako wejście.
W tym ćwiczeniu zaimportowano zmienne hsize (rozmiar ukryty warstwy GRU), fr_len oraz fr_vocab. Pamiętaj, że przedrostek de odnosi się do dekodera.
To ćwiczenie jest częścią kursu
Tłumaczenie maszynowe z Keras
Instrukcje do ćwiczenia
- Zdefiniuj warstwę
Input, która przyjmuje wsad sekwencji słów francuskich zakodowanych metodą one-hot (długość sekwencji równa 1). - Zdefiniuj kolejną warstwę
Input, która przyjmuje wsad stanów o rozmiarzehsize– będzie ona służyć do przekazywania poprzedniego stanu do dekodera. - Pobierz wyjście i stan dekodera
GRU. - Zdefiniuj model, który przyjmuje warstwę
Inputze słowami francuskimi oraz warstwęInputz poprzednim stanem, a zwraca końcową predykcję i nowy stanGRU.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
import tensorflow.keras.layers as layers
from tensorflow.keras.models import Model
# Define an input layer that accepts a single onehot encoded word
de_inputs = layers.____(shape=(____, ____))
# Define an input to accept the t-1 state
de_state_in = layers.____(shape=(____,))
de_gru = layers.GRU(hsize, return_state=True)
# Get the output and state from the GRU layer
de_out, de_state_out = ____(de_inputs, initial_state=____)
de_dense = layers.Dense(fr_vocab, activation='softmax')
de_pred = de_dense(de_out)
# Define a model
decoder = Model(inputs=[____, ____], outputs=[____, ____])
print(decoder.summary())