Skip to content

Commit 52b50c4

Browse files
authored
Merge pull request #11 from Yoctol/seq2seq-variable
Seq2seq variable
2 parents acb1972 + fe04b81 commit 52b50c4

1 file changed

Lines changed: 26 additions & 7 deletions

File tree

seq2vec/seq2seq_auto_encoder.py

Lines changed: 26 additions & 7 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
@@ -88,10 +103,14 @@ def _generate_padding_array(self, seqs):
88103
array.append(np_seq)
89104
return np.array(array)
90105

91-
def fit(self, train_seqs, verbose=2, nb_epoch=10, validation_split=0.0):
106+
def fit(self, train_seqs, predict_seqs=None, verbose=2, nb_epoch=10, validation_split=0.0):
92107
train_x = self._generate_padding_array(train_seqs)
108+
if predict_seqs is None:
109+
train_y = train_x
110+
else:
111+
train_y = self._generate_padding_array(predict_seqs)
93112
self.model.fit(
94-
train_x, train_x,
113+
train_x, train_y,
95114
verbose=verbose,
96115
nb_epoch=nb_epoch,
97116
validation_split=validation_split,

0 commit comments

Comments
 (0)