Začněte nyníZačněte zdarma

Definování modelu s embedding vrstvami

Definuješ Keras model, který:

  • Používá vrstvy Embedding
  • Bude trénován pomocí Teacher Forcing

Tento model bude mít dvě embedding vrstvy – jednu pro enkodér a jednu pro dekodér. Protože je model trénován pomocí Teacher Forcing, použije v Input vrstvě dekodéru délku sekvence fr_len-1.

Pro toto cvičení máš k dispozici všechny potřebné keras.layers a Model. Jsou také definovány proměnné en_len (délka anglické sekvence), fr_len (délka francouzské sekvence), en_vocab (velikost anglické slovní zásoby), fr_vocab (velikost francouzské slovní zásoby) a hsize (velikost skryté vrstvy).

Toto cvičení je součástí kurzu

Machine Translation with Keras

Zobrazit kurz

Pokyny k cvičení

  • Definuj vrstvu Input, která přijímá sekvenci ID slov.
  • Definuj vrstvu Embedding, která zakóduje en_vocab slov, má délku 96 a dokáže přijmout sekvenci ID (délka sekvence se zadává argumentem input_length).
  • Definuj vrstvu Embedding, která zakóduje fr_vocab slov, má délku 96 a dokáže přijmout sekvenci fr_len-1 ID.
  • Definuj model, který přijímá vstup z enkodéru a vstup z dekodéru (v tomto pořadí) a vrací predikce slov.

Interaktivní cvičení na vyzkoušení si v praxi

Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.

# 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'])
Upravit a spustit kód