From 1f9b356f97297ef5750c12fafd48444d9dfc2f69 Mon Sep 17 00:00:00 2001 From: Nick Date: Wed, 30 Oct 2019 11:15:26 -0600 Subject: [PATCH] update --- generator/ctrl/ctrl_generator.py | 2 +- story/story_manager.py | 1 + story/utils.py | 6 ++++++ 3 files changed, 8 insertions(+), 1 deletion(-) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index 60273a6..d0cc8e3 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -297,7 +297,7 @@ class CTRLGenerator(): num_new_lines = 0 for token in range(len(text) - 1, total_text_len - 1): idx = self.generate_next_token(token, tokens_generated, options, num_new_lines, token_num, first_token=first_token, forbid_newline=False) - if self.idx2word[idx] == '\n' and token_num < 10: + if self.idx2word[idx] == '\n' and token_num > 7: return self.result_replace(result) elif self.idx2word[idx] == '\n': idx = self.generate_next_token(token, tokens_generated, options, num_new_lines, token_num, diff --git a/story/story_manager.py b/story/story_manager.py index c2ee15d..2bb8ead 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -209,6 +209,7 @@ class CTRLStoryManager(ConstrainedStoryManager): def get_action_results_generate(self): results = [] options = {"word_blacklist": {0:[]}} + options["word_whitelist"] = {0: get_allowed_ctrl_verbs()} for phrase in self.action_phrases: result = self.generate_action_result(self.story_context(), phrase, options=options) action_verb = result[0].split()[1] diff --git a/story/utils.py b/story/utils.py index 70d656c..586550a 100644 --- a/story/utils.py +++ b/story/utils.py @@ -13,6 +13,12 @@ def get_context(key): return data_loaded["contexts"][key] +def get_allowed_ctrl_verbs(): + with open(YAML_FILE, 'r') as stream: + data_loaded = yaml.safe_load(stream) + + return data_loaded["ctrl_verbs"]["movement"] + data_loaded["ctrl_verbs"]["non_movement"] + def get_story_start(key): with open(YAML_FILE, 'r') as stream: data_loaded = yaml.safe_load(stream)