mirror of
https://github.com/wassname/keras-language-modeling.git
synced 2026-09-09 11:25:29 +08:00
merged changes
This commit is contained in:
+9
-5
@@ -26,9 +26,12 @@ class AttentionLSTM(LSTM):
|
||||
name='{}_U_m'.format(self.name))
|
||||
self.b_m = K.zeros((self.output_dim,), name='{}_b_m'.format(self.name))
|
||||
|
||||
self.U_s = self.inner_init((self.output_dim, self.output_dim),
|
||||
# self.U_s = self.inner_init((self.output_dim, self.output_dim),
|
||||
# name='{}_U_s'.format(self.name))
|
||||
# self.b_s = K.zeros((self.output_dim,), name='{}_b_s'.format(self.name))
|
||||
self.U_s = self.inner_init((self.output_dim, 1),
|
||||
name='{}_U_s'.format(self.name))
|
||||
self.b_s = K.zeros((self.output_dim,), name='{}_b_s'.format(self.name))
|
||||
self.b_s = K.zeros((1,), name='{}_b_s'.format(self.name))
|
||||
|
||||
self.trainable_weights += [self.U_a, self.U_m, self.U_s, self.b_a, self.b_m, self.b_s]
|
||||
|
||||
@@ -43,9 +46,10 @@ class AttentionLSTM(LSTM):
|
||||
m = K.tanh(K.dot(h, self.U_a) * attention + self.b_a)
|
||||
# Intuitively it makes more sense to use a sigmoid (was getting some NaN problems
|
||||
# which I think might have been caused by the exponential function -> gradients blow up)
|
||||
s = K.exp(K.dot(m, self.U_s) + self.b_s)
|
||||
# s = K.sigmoid(K.dot(m, self.U_s) + self.b_s)
|
||||
h = h * s
|
||||
# s = K.exp(K.dot(m, self.U_s) + self.b_s)
|
||||
s = K.tanh(K.dot(m, self.U_s) + self.b_s)
|
||||
h = h * K.repeat_elements(s, self.output_dim, axis=1)
|
||||
# h = h * s
|
||||
|
||||
return h, [h, c]
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from time import strftime, gmtime
|
||||
|
||||
import pickle
|
||||
|
||||
from keras.optimizers import Adam
|
||||
from keras.optimizers import Adam, RMSprop
|
||||
from scipy.stats import rankdata
|
||||
|
||||
from keras_models import *
|
||||
@@ -98,6 +98,7 @@ class Evaluator:
|
||||
|
||||
questions = self.padq(questions)
|
||||
good_answers = self.pada(good_answers)
|
||||
# bad_answers = self.pada(random.sample(self.answers.values(), len(good_answers)))
|
||||
|
||||
for i in range(nb_epoch):
|
||||
# bad_answers = np.roll(good_answers, random.randint(10, len(questions) - 10))
|
||||
@@ -204,7 +205,7 @@ if __name__ == '__main__':
|
||||
'question_len': 20,
|
||||
'answer_len': 100,
|
||||
'n_words': 22353, # len(vocabulary) + 1
|
||||
'margin': 0.009,
|
||||
'margin': 0.02,
|
||||
|
||||
'training_params': {
|
||||
'save_every': 1,
|
||||
@@ -212,12 +213,12 @@ if __name__ == '__main__':
|
||||
'batch_size': 128,
|
||||
'nb_epoch': 1000,
|
||||
'validation_split': 0.2,
|
||||
'optimizer': 'adam',
|
||||
'optimizer': RMSprop(clip_norm=0.1), # Adam(clip_norm=0.1),
|
||||
'n_eval': 20,
|
||||
|
||||
'evaluate_all_threshold': {
|
||||
'mode': 'all',
|
||||
'top1': 0.55,
|
||||
'top1': 0.5,
|
||||
},
|
||||
},
|
||||
|
||||
@@ -264,8 +265,8 @@ if __name__ == '__main__':
|
||||
|
||||
# train the model
|
||||
# evaluator.load_epoch(model, 25)
|
||||
# evaluator.train(model)
|
||||
evaluator.train(model)
|
||||
|
||||
# evaluate mrr for a particular epoch
|
||||
evaluator.load_epoch(model, 115)
|
||||
evaluator.get_mrr(model, evaluate_all=True)
|
||||
# evaluator.load_epoch(model, 53)
|
||||
# evaluator.get_mrr(model, evaluate_all=True)
|
||||
|
||||
@@ -259,6 +259,10 @@ class AttentionModel(LanguageModel):
|
||||
question_rnn = merge([f_rnn(question_dropout), b_rnn(question_dropout)], mode='concat', concat_axis=-1)
|
||||
question_dropout = dropout(question_rnn)
|
||||
|
||||
# regularize
|
||||
regularize = ActivityRegularization(l2=0.0001)
|
||||
question_dropout = regularize(question_dropout)
|
||||
|
||||
# could add convolution layer here (as in paper)
|
||||
|
||||
# maxpooling
|
||||
@@ -272,6 +276,7 @@ class AttentionModel(LanguageModel):
|
||||
# b_rnn = LSTM(self.model_params.get('n_lstm_dims', 141), return_sequences=True, go_backwards=True)
|
||||
answer_rnn = merge([f_rnn(answer_dropout), b_rnn(answer_dropout)], mode='concat', concat_axis=-1)
|
||||
answer_dropout = dropout(answer_rnn)
|
||||
answer_dropout = regularize(answer_dropout)
|
||||
answer_pool = maxpool(answer_dropout)
|
||||
|
||||
# activation
|
||||
|
||||
Reference in New Issue
Block a user