From f743428c15117ded8c7f9a3e6c56e86d10ad7d47 Mon Sep 17 00:00:00 2001 From: Janis Date: Sun, 12 Aug 2018 03:37:49 +1000 Subject: [PATCH 1/2] Attempt to pickle model --- persephone_api/api_endpoints/model.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/persephone_api/api_endpoints/model.py b/persephone_api/api_endpoints/model.py index 3e9e108..f83851d 100644 --- a/persephone_api/api_endpoints/model.py +++ b/persephone_api/api_endpoints/model.py @@ -41,6 +41,9 @@ def create_RNN_CTC_model(model: TranscriptionModel, corpus_storage_path: Path, beam_width=model.beam_width, decoding_merge_repeated=model.decoding_merge_repeated ) + model_pickle_path = model_path / "model.p" + with model_pickle_path.open('wb') as pickle_file: + pickle.dump(model, pickle_file, protocol=4) # TODO: pickle model at this point? def search(): @@ -74,7 +77,6 @@ def post(modelInfo): model_uuid = uuid.uuid1() - import pdb; pdb.set_trace() current_model = TranscriptionModel( name=modelInfo['name'], corpus=current_corpus, From 2cdbf98b75392b73f172aebe286297ede41d9858 Mon Sep 17 00:00:00 2001 From: Janis Date: Sun, 12 Aug 2018 21:03:13 +1000 Subject: [PATCH 2/2] Attempt to debug source of pickle error --- persephone_api/api_endpoints/model.py | 20 ++++++++++++++++++-- 1 file changed, 18 insertions(+), 2 deletions(-) diff --git a/persephone_api/api_endpoints/model.py b/persephone_api/api_endpoints/model.py index f83851d..4dae2ae 100644 --- a/persephone_api/api_endpoints/model.py +++ b/persephone_api/api_endpoints/model.py @@ -16,6 +16,21 @@ from ..db_models import DBcorpus, TranscriptionModel from ..serialization import TranscriptionModelSchema +class MyPickler(pickle._Pickler): + def save(self, obj): + print('pickling object {0} of type {1}'.format(obj, type(obj))) + try: + pickle._Pickler.save(self, obj) + except: + import pdb; pdb.set_trace() + raise + +# attempting workaround to get TF to pickle properly: +import tensorflow as tf +setattr(tf.contrib.rnn.GRUCell, '__deepcopy__', lambda self, _: self) +setattr(tf.contrib.rnn.BasicLSTMCell, '__deepcopy__', lambda self, _: self) +setattr(tf.contrib.rnn.MultiRNNCell, '__deepcopy__', lambda self, _: self) + def create_RNN_CTC_model(model: TranscriptionModel, corpus_storage_path: Path, models_storage_path: Path): """Create a persephone RNN CTC model @@ -33,7 +48,7 @@ def create_RNN_CTC_model(model: TranscriptionModel, corpus_storage_path: Path, corpus = pickle.load(pickle_file) corpus_reader = CorpusReader(corpus) - model = rnn_ctc.Model( + persephone_model = rnn_ctc.Model( exp_dir, corpus_reader, num_layers=model.num_layers, @@ -43,7 +58,8 @@ def create_RNN_CTC_model(model: TranscriptionModel, corpus_storage_path: Path, ) model_pickle_path = model_path / "model.p" with model_pickle_path.open('wb') as pickle_file: - pickle.dump(model, pickle_file, protocol=4) + p = MyPickler(pickle_file, protocol=4) + p.dump(persephone_model) # TODO: pickle model at this point? def search():