Definiowanie modelu z osadzeniami
Zdefiniujesz model Keras, który:
- korzysta z warstw
Embedding, - będzie trenowany z użyciem Teacher Forcing.
Model będzie miał dwie warstwy osadzeń (embeddings) – jedną dla enkodera i jedną dla dekodera. Ponieważ model jest trenowany z użyciem Teacher Forcing, warstwa Input dekodera przyjmuje sekwencje o długości fr_len-1.
W tym ćwiczeniu wszystkie potrzebne elementy keras.layers oraz Model są już zaimportowane. Dostępne są też zmienne: en_len (długość sekwencji angielskiej), fr_len (długość sekwencji francuskiej), en_vocab (rozmiar słownika angielskiego), fr_vocab (rozmiar słownika francuskiego) oraz hsize (rozmiar warstwy ukrytej).
To ćwiczenie jest częścią kursu
Tłumaczenie maszynowe z Keras
Instrukcje do ćwiczenia
- Zdefiniuj warstwę
Input, która przyjmuje sekwencję identyfikatorów słów. - Zdefiniuj warstwę
Embeddingosadzającą słowa zen_vocab, o długości 96, która przyjmuje sekwencję identyfikatorów (długość sekwencji określa argumentinput_length). - Zdefiniuj warstwę
Embeddingosadzającą słowa zfr_vocab, o długości 96, która przyjmuje sekwencjęfr_len-1identyfikatorów. - Zdefiniuj model, który przyjmuje dane wejściowe z enkodera i dekodera (w tej kolejności) i zwraca predykcje słów.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# Define an input layer which accepts a sequence of word IDs
en_inputs = Input(____=(____,))
# Define an Embedding layer which accepts en_inputs
en_emb = ____(____, ____, input_length=____)(en_inputs)
en_out, en_state = GRU(hsize, return_state=True)(en_emb)
de_inputs = Input(shape=(fr_len-1,))
# Define an Embedding layer which accepts de_inputs
de_emb = Embedding(____, 96, input_length=____)(____)
de_out, _ = GRU(hsize, return_sequences=True, return_state=True)(de_emb, initial_state=en_state)
de_pred = TimeDistributed(Dense(fr_vocab, activation='softmax'))(de_out)
# Define the Model which accepts encoder/decoder inputs and outputs predictions
nmt_emb = Model([____, ____], ____)
nmt_emb.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['acc'])