Commit 1c3044cb authored by Rudolf Chrispens's avatar Rudolf Chrispens
Browse files

notebook save for demo

parent 8f8e2c7a
Loading
Loading
Loading
Loading
+5 −2
Original line number Diff line number Diff line
@@ -58,10 +58,13 @@
   "outputs": [],
   "source": [
    "# setting path to data that gets used by our trained model\n",
    "_config.path.input_validate = '/coco.valid'\n",
    "_config.path.demo_sample_on = True\n",
    "_config.path.demo_sample_path = os.path.dirname(__file__) + \"/../../Demo/example.pickle\"\n",
    "\n",
    "\n",
    "# setting path to image file for attention visualisation\n",
    "_config.path.image_validation = '/validation/val2017'\n",
    "_config.path.image_data_path = os.path.dirname(__file__) + \"/../../cocoapi\"\n",
    "_config.path.image_validation = '/val'\n",
    "\n",
    "# enable validation and attention\n",
    "_config.validator.max_iteration = 1\n",
+3 −0
Original line number Diff line number Diff line
@@ -27,6 +27,9 @@ class default_paths(object):
    image_validation = "/validation/val2017"
    image_test = "/test/test2017"
    image_train = "/train/train2017"
    #sample validation for Demp
    demo_sample_on = False
    demo_sample_path = os.path.dirname(__file__) + "/../../Demo/example.pickle"


class default_reader(object):  # default reader config
+5 −1
Original line number Diff line number Diff line
@@ -82,8 +82,12 @@ def main(self, parameter_list):
                                                modelpath=_reader.helper.modelPath + _config.path.model_for_validation,
                                                batch_size_of_trained_model=_config.trainer.batch_size)

        #change sample
        if _config.path.demo_sample_on:
            v_data = current_reader.load_data_encoded(filename=_config.path.demo_sample_path)

        # validate trained
        validator.test(val_data, current_reader=current_reader, split='val', attention_visualization=True, save_sampled_captions=True)
        validator.test(v_data, current_reader=current_reader, split='val', attention_visualization=True, save_sampled_captions=True)

        # test trained
        #validator.test(test_data, current_reader=current_reader, split='test', attention_visualization=True, save_sampled_captions=True)