Skip to content

Commit da503e8

Browse files
committed
make learning rate adjustable
1 parent acb1972 commit da503e8

1 file changed

Lines changed: 20 additions & 5 deletions

File tree

seq2vec/seq2seq_auto_encoder.py

Lines changed: 20 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,14 @@
1313
from .base import TrainableInterfaceMixin
1414

1515

16-
def _create_single_layer_seq2seq_model(max_length, max_index, latent_size):
16+
def _create_single_layer_seq2seq_model(
17+
max_length,
18+
max_index,
19+
latent_size,
20+
learning_rate,
21+
rho=0.9,
22+
decay=0.0,
23+
):
1724
inputs = Input(shape=(max_length, max_index))
1825
encoded = LSTM(latent_size)(inputs)
1926
decoded = Dropout(0.3)(encoded)
@@ -23,9 +30,9 @@ def _create_single_layer_seq2seq_model(max_length, max_index, latent_size):
2330
encoder = Model(inputs, encoded)
2431

2532
optimizer = RMSprop(
26-
lr=0.0001,
27-
rho=0.95,
28-
decay=0.1,
33+
lr=learning_rate,
34+
rho=rho,
35+
decay=decay,
2936
)
3037
model.compile(loss='categorical_crossentropy', optimizer=optimizer)
3138
return model, encoder
@@ -56,15 +63,23 @@ class Seq2SeqAutoEncoderUseWordHash(TrainableInterfaceMixin, BaseSeq2Vec):
5663
5764
"""
5865

59-
def __init__(self, max_index, max_length, latent_size=20):
66+
def __init__(
67+
self,
68+
max_index,
69+
max_length,
70+
learning_rate=0.0001,
71+
latent_size=20,
72+
):
6073
self.max_index = max_index
6174
self.max_length = max_length
75+
self.learning_rate = learning_rate
6276
self.latent_size = latent_size
6377

6478
model, encoder = _create_single_layer_seq2seq_model(
6579
max_length=self.max_length,
6680
max_index=self.max_index,
6781
latent_size=self.latent_size,
82+
learning_rate=self.learning_rate,
6883
)
6984
self.model = model
7085
self.encoder = encoder

0 commit comments

Comments
 (0)