From bd6653113cd59b5bdbab624ddef0eca92266e15a Mon Sep 17 00:00:00 2001 From: Nick Date: Tue, 24 Sep 2019 14:23:30 -0600 Subject: [PATCH] fixed mental issue? --- generator/ctrl/ctrl_generator.py | 16 ++++++++-------- story/story_manager.py | 2 +- 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index ebbbc2a..9d40914 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -22,7 +22,7 @@ def loss(labels, logits): class CTRLGenerator(): - def __init__(self, control_code="Horror Text: ", generate_num=100, temperature=0.75): + def __init__(self, control_code="Horror Text: ", generate_num=100, temperature=0.3): self.generate_num=generate_num model_dir = "generator/ctrl/model/seqlen256_v1.ckpt/" @@ -185,6 +185,11 @@ class CTRLGenerator(): def generate(self, prompt, options=None): prompt = self.prompt_replace(prompt) + if options is None: + options = dict() + if "used_verbs" not in options: + options["used_verbs"] = set() + if prompt[-1] != " ": prompt = prompt + " " first_token = True @@ -192,11 +197,6 @@ class CTRLGenerator(): prompt = second_to_first_person(prompt) prompt = self.control_code + prompt - - print("******************************") - print(" DEBUG:: Prompt to generate by is \n", prompt) - print("******************************") - prompt_length = len(prompt) # tokenize provided prompt @@ -264,7 +264,8 @@ class CTRLGenerator(): # Make sure only a possible verb is chosen. if first_token: for word in get_possible_verbs(): - prompt_logits[_token][self.word2idx[word]] += 5 + if word not in options["used_verbs"]: + prompt_logits[_token][self.word2idx[word]] += 5 # compute probabilities from logits prompt_probs = np.exp(prompt_logits[_token]) @@ -324,7 +325,6 @@ class CTRLGenerator(): tokens_generated_so_far = ' '.join([self.idx2word[c] for c in tokens_generated[0].squeeze()[:token + 2]]) tokens_generated_so_far = re.sub('(@@ )', '', string=tokens_generated_so_far) tokens_generated_so_far = re.sub('(@@ ?$)', '', string=tokens_generated_so_far) - print(tokens_generated_so_far) result = tokens_generated_so_far[prompt_length:] first_token = False diff --git a/story/story_manager.py b/story/story_manager.py index b3b8b8d..5d85332 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -164,7 +164,7 @@ class CTRLStoryManager(ConstrainedStoryManager): results = [] for phrase in self.action_phrases: options = dict() - options["used_verbs"] = used_verbs + options["used_verbs"] = set(used_verbs) result = self.generate_action_result(self.story_context(), phrase, options=options) used_verb = result[0].split()[1]