Antrenarea modelului bazat pe încorporări de cuvinte
Aici vei învăța cum să implementezi procesul de antrenare pentru un model de traducere automată care folosește încorporări de cuvinte (word embeddings). Un cuvânt este reprezentat printr-un singur număr, în loc de un vector one-hot encodat, așa cum ai făcut în exercițiile anterioare. Vei antrena modelul pe mai multe epoci, parcurgând întregul set de date în loturi (batch-uri).
Pentru acest exercițiu ai la dispoziție date de antrenament (tr_en și tr_fr) sub forma unei liste de propoziții. Vei folosi doar un eșantion foarte mic (1.000 de propoziții) din datele reale, deoarece antrenarea pe întregul set ar dura foarte mult. Ai, de asemenea, funcția sents2seqs() și modelul nmt_emb, pe care l-ai implementat în exercițiul anterior. Reține că folosim en_x pentru intrările encoderului și de_x pentru intrările decoderului.
Acest exercițiu face parte din cursul
Traducere automată cu Keras
Instrucțiuni pentru exercițiu
- Obține un singur lot de propoziții în franceză fără one-hot encoding, folosind funcția
sents2seqs(). - Obține toate cuvintele din
de_xy, mai puțin ultimul. - Obține toate cuvintele din
de_xy_oh(cuvinte în franceză cu one-hot encoding), mai puțin primul. - Antrenează modelul folosind un singur lot de date.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
for ei in range(3):
for i in range(0, train_size, bsize):
en_x = sents2seqs('source', tr_en[i:i+bsize], onehot=False, reverse=True)
# Get a single batch of French sentences with no onehot encoding
de_xy = ____('target', ____[i:i+bsize], ____=____)
# Get all words except the last word in that batch
de_x = de_xy[:,____]
de_xy_oh = sents2seqs('target', tr_fr[i:i+bsize], onehot=True)
# Get all words except the first from de_xy_oh
de_y = de_xy_oh[____,____,____]
# Training the model on a single batch of data
nmt_emb.train_on_batch([____,____], ____)
res = nmt_emb.evaluate([en_x, de_x], de_y, batch_size=bsize, verbose=0)
print("{} => Loss:{}, Train Acc: {}".format(ei+1,res[0], res[1]*100.0))