Skip to content

Commit a949f04

Browse files
author
plliao
committed
update readme
1 parent 7d27baf commit a949f04

1 file changed

Lines changed: 8 additions & 26 deletions

File tree

README.md

Lines changed: 8 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -193,23 +193,16 @@ transformer = Seq2VecR2RWord(
193193
learning_rate=0.05
194194
)
195195

196-
input_transformer = WordEmbeddingTransformer(
197-
word2vec, max_length
198-
)
199-
output_transformer = WordEmbeddingTransformer(
200-
word2vec, max_length
201-
)
202-
203196
train_data = DataGenterator(
204197
corpus_for_training_path,
205-
input_transformer,
206-
output_transformer,
198+
transformer.input_transformer,
199+
transformer.output_transformer,
207200
batch_size=128
208201
)
209202
test_data = DataGenterator(
210203
corpus_for_validation_path,
211-
input_transformer,
212-
output_transformer,
204+
transformer.input_transformer,
205+
transformer.output_transformer,
213206
batch_size=128
214207
)
215208

@@ -256,6 +249,10 @@ class YourSeq2Vec(TrainableSeq2VecBase):
256249
self.input_transformer = YourInputTransformer()
257250
self.output_transformer = YourOutputTransformer()
258251

252+
# add your customized layer
253+
self.custom_objects = {}
254+
self.custom_objects[customized_class_name] = customized_class
255+
259256
super(YourSeq2Vec, self).__init__(
260257
max_length,
261258
latent_size,
@@ -270,21 +267,6 @@ class YourSeq2Vec(TrainableSeq2VecBase):
270267
model.compile(loss)
271268
return model, encoder
272269

273-
def transform(self, seqs):
274-
# define how your encoder transform input sequences
275-
# into fixed length vectors
276-
return fixed_length_vectors
277-
278-
def load_customed_model(self, file_path):
279-
# if you use customized layer in yklz or with your
280-
# own layers, you have to sepcify them here.
281-
return keras.models.load_model(
282-
file_path,
283-
custom_objects={
284-
'CustomizedLayer':CustomizedLayer
285-
}
286-
)
287-
288270
def load_model(self, file_path):
289271
# load your seq2vec model here and set its attribute values
290272
self.model = self.load_customed_model(file_path)

0 commit comments

Comments
 (0)