dice-group / dice-group/EuroPython-2018

Embedding is not used in model

Open
#1 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
8
Forks
4
PR merge metrics
No merged PRs in 30d

Description

You're creating embedding layer in following cell, but then it's not used in building a model:

```python
from keras.layers import Embedding

ques_embedding_layer = Embedding(len(word_index) + 1, #input_dim: vocab_size
EMBEDDING_DIM, # the size of the output vectors from this layer
weights=[embedding_matrix],
input_length=ques_maxlen, # length of input sequences
trainable=False)
context_embedding_layer = Embedding(len(word_index) + 1,
EMBEDDING_DIM,
weights=[embedding_matrix],
input_length=context_maxlen,
trainable=False)
```
```python
import keras
from keras import backend as K
from keras.models import Sequential, Model
from keras.layers.embeddings import Embedding
from keras.layers import Input, Activation, Dense, Permute, Dropout, concatenate, RepeatVector, multiply
from keras.layers import LSTM, Input, Bidirectional, Masking, Lambda, TimeDistributed, Flatten

P = Input(shape=(context_maxlen, EMBEDDING_DIM), name='P')
Q = Input(shape=(ques_maxlen, EMBEDDING_DIM), name='Q')
W = 28
passage_input = P
question_input = Q
encoder = Bidirectional(LSTM(units=W,return_sequences=True))

passage_encoding = P
passage_encoding = encoder(passage_encoding)
passage_encoding = TimeDistributed(Dense(W, use_bias=False, trainable=True, weights=np.concatenate((np.eye(W), np.eye(W)), axis=1)))(passage_encoding)

question_encoding = Q
question_encoding = encoder(question_encoding)
question_encoding = TimeDistributed(Dense(W, use_bias=False, trainable=True, weights=np.concatenate((np.eye(W), np.eye(W)), axis=1)))(question_encoding)

question_attention_vector = TimeDistributed(Dense(1))(question_encoding)
question_attention_vector = Lambda(lambda q: keras.activations.softmax(q, axis=1))(question_attention_vector)

question_attention_vector = Lambda(lambda q: q[0] * q[1])([question_encoding, question_attention_vector])
question_attention_vector = Lambda(lambda q: K.sum(q, axis=1))(question_attention_vector)
question_attention_vector = RepeatVector(context_maxlen)(question_attention_vector)

answer_start = Lambda(lambda arg: concatenate([arg[0], arg[1], arg[2]]))([
passage_encoding,
question_attention_vector,
multiply([passage_encoding, question_attention_vector])])

answer_start = TimeDistributed(Dense(W, activation='relu'))(answer_start)
answer_start = TimeDistributed(Dense(1))(answer_start)
answer_start = Flatten()(answer_start)
answer_start = Activation('softmax')(answer_start)

# Answer end prediction depends on the start prediction
def s_answer_feature(x):
maxind = K.argmax( x,axis=1,)
return maxind

x = Lambda(lambda x: K.tf.cast(s_answer_feature(x), dtype=K.tf.int32))(answer_start)
start_feature = Lambda(lambda arg: K.tf.gather_nd(arg[0], K.tf.stack(
[K.tf.range(K.tf.shape(arg[1])[0]), K.tf.cast(arg[1], K.tf.int32)], axis=1)))([passage_encoding, x])
start_feature = RepeatVector(context_maxlen)(start_feature)

# Answer end prediction
answer_end = Lambda(lambda arg: concatenate([
arg[0],
arg[1],
arg[2],
multiply([arg[0], arg[1]]),
multiply([arg[0], arg[2]])]))([passage_encoding, question_attention_vector, start_feature])

answer_end = TimeDistributed(Dense(W, activation='relu'))(answer_end)
answer_end = TimeDistributed(Dense(1))(answer_end)
answer_end = Flatten()(answer_end)
answer_end = Activation('softmax')(answer_end)

input_placeholders = [P, Q]
inputs = input_placeholders
outputs = [answer_start, answer_end]
```

also there is no actual construction of a model so far and it's fitting to data. To compile model I would propose adding:
```
model = Model(inputs=inputs, outputs=outputs)
model.compile(optimizer='adam', loss='binary_crossentropy')
model.summary()
```
But I'm not sure since there is no embedding layer in there.

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.