diff --git a/console_play.py b/console_play.py index 2ba2c89..729f646 100644 --- a/console_play.py +++ b/console_play.py @@ -40,7 +40,7 @@ def play_unconstrained(): def play_constrained(): #generator = WebGenerator(CRED_FILE) generator = CTRLGenerator() - story_start = "classic" + story_start = "haunted" verbs_key = "anything" prompt = get_story_start(story_start) story_manager = ConstrainedStoryManager(generator, prompt, action_verbs_key=verbs_key) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index 53c0b86..f16d7b9 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -137,11 +137,13 @@ class CTRLGenerator(): self.topk = 0 - def generate(self, prompt): + def generate(self, prompt, first_verb_whitelist=True): if prompt[-1] != " ": prompt = prompt + " " + first_token = True + prompt = second_to_first_person(prompt) prompt = self.control_code + prompt @@ -198,13 +200,15 @@ class CTRLGenerator(): prompt_logits[_token][generated_token] /= self.penalty # disallow some tokens - prompt_logits[_token][self.word2idx['']] = -1e8 - prompt_logits[_token][self.word2idx['\n']] = -1e8 + forbidden_tokens = ['', 'Sco@@'] + for token in forbidden_tokens: + prompt_logits[_token][self.word2idx[token]] = -1e8 - # sometimes, when generating from reddit, - # it tries to generate the Score (reddit Karma) immediately after generating the Title: - # to disallow this, we can just prevent it from generating Score - prompt_logits[_token][self.word2idx['Sco@@']] = -1e8 + # Make sure only a possible verb is chosen. + if first_token: + for word in get_possible_verbs(): + prompt_logits[_token][self.word2idx[word]] += 0.1 + first_token = False # compute probabilities from logits prompt_probs = np.exp(prompt_logits[_token]) @@ -228,7 +232,7 @@ class CTRLGenerator(): # for instance, if you want to disallow anything with the phrase `http`, # you can delete theme from the pruned_list # you can comment this out, I'm keeping it in for demonstration purpose - tokens_to_disallow = [] + tokens_to_disallow = ["\n\n"] for _ in range(len(pruned_list)): if 'http' in self.idx2word[pruned_list[_]]: tokens_to_disallow.append(_) diff --git a/story/story_manager.py b/story/story_manager.py index 7fa3373..4d8d6a7 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -109,7 +109,7 @@ class ConstrainedStoryManager(StoryManager): return [self.generate_action_result(self.story_context(), phrase) for phrase in self.action_phrases] def generate_action_result(self, prompt, phrase): - action = phrase + self.generator.generate(prompt + phrase) + action = phrase + " " + self.generator.generate(prompt + phrase) action_result = cut_trailing_sentence(action) action, result = split_first_sentence(action_result) diff --git a/story/utils.py b/story/utils.py index 5070223..7e1509d 100644 --- a/story/utils.py +++ b/story/utils.py @@ -133,6 +133,12 @@ def second_to_first_person(text): return capitalize_first_letters(text) + +possible_verbs = ["go", "run", "open", "look", "walk", "make", "try", "say", "tell", "attack", "use", "turn", "fight", "scream", "yell"] + +def get_possible_verbs(): + return possible_verbs + if __name__ == '__main__': f = open("test.txt", "r") test_text = f.read()