From da7f2b776cabf8ae93ae1cd93b7ae01ee2beab90 Mon Sep 17 00:00:00 2001 From: Nick Date: Wed, 25 Sep 2019 05:41:35 -0600 Subject: [PATCH] update --- generator/ctrl/ctrl_generator.py | 23 +++++++++++++---------- story/story_manager.py | 30 ++++++------------------------ 2 files changed, 19 insertions(+), 34 deletions(-) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index 3a63b55..b446a18 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -13,8 +13,6 @@ from story.utils import * import warnings warnings.filterwarnings("ignore") -pos_action_starts = ["You attack", "You tell", "You use", "You go"] - # the loss function is a simple categorical crossentropy between the logits and the labels def loss(labels, logits): return tf.keras.losses.sparse_categorical_crossentropy(labels, logits, from_logits=True) @@ -187,7 +185,7 @@ class CTRLGenerator(): return result - def generate_next_token(self, token, tokens_generated, options, num_new_lines, first_token=False): + def generate_next_token(self, token, tokens_generated, options, num_new_lines, token_num, first_token=False): # get the logits from the prediction function # the logic here is a bit convoluted because we are allowing generation past 512 tokens @@ -231,11 +229,15 @@ class CTRLGenerator(): for forbidden_token in forbidden_tokens: prompt_logits[_token][self.word2idx[forbidden_token]] = -1e8 - # Make sure only a possible verb is chosen. - if first_token: - for word in get_possible_verbs(): - if word not in options["used_verbs"]: - prompt_logits[_token][self.word2idx[word]] += 100 + # Set whitelist + if "word_whitelist" in options and token_num in options["word_whitelist"].keys(): + for word in options["word_whitelist"][token_num]: + prompt_logits[_token][self.word2idx[word]] += 100 + + # Set blacklist, overwrites whitelist + if "word_blacklist" in options and token_num in options["word_blacklist"].keys(): + for word in options["word_blacklist"][token_num]: + prompt_logits[_token][self.word2idx[word]] = -1e8 # compute probabilities from logits prompt_probs = np.exp(prompt_logits[_token]) @@ -316,9 +318,10 @@ class CTRLGenerator(): tokens_generated = np.tile(padded_text, (1, 1)) result = "" + token_num = 0 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, first_token=first_token) + idx = self.generate_next_token(token, tokens_generated, options, num_new_lines, token_num, first_token=first_token) if self.idx2word[idx] is "\n": num_new_lines += 1 @@ -329,7 +332,7 @@ class CTRLGenerator(): tokens_generated_so_far = re.sub('(@@ ?$)', '', string=tokens_generated_so_far) result = tokens_generated_so_far - first_token = False + token_num += 1 print("PROMPT: \n", prompt) print("RESULT: \n", result) diff --git a/story/story_manager.py b/story/story_manager.py index f53cfec..bf70a75 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -21,7 +21,6 @@ class Story(): if game_state is None: game_state = dict() - game_state["current_room"] = possible_rooms[0] self.game_state = game_state @@ -73,10 +72,10 @@ class StoryManager(): def __init__(self, generator): self.generator = generator - def start_new_story(self, story_prompt): + def start_new_story(self, story_prompt, game_state=None): block = self.generator.generate(story_prompt) block = cut_trailing_sentence(block) - self.story = Story(story_prompt + block) + self.story = Story(story_prompt + block, game_state=None) return self.story def load_story(self, story, from_json=False): @@ -113,7 +112,7 @@ class ConstrainedStoryManager(StoryManager): super().__init__(generator) self.action_phrases = get_action_verbs(action_verbs_key) - def start_new_story(self, story_prompt): + def start_new_story(self, story_prompt, game_state=None): super().start_new_story(story_prompt) self.story.possible_action_results = self.get_action_results() @@ -166,9 +165,8 @@ class CTRLStoryManager(ConstrainedStoryManager): def __init__(self, generator, action_verbs_key="anything"): super().__init__(generator, action_verbs_key) - def start_new_story(self, story_prompt): + def start_new_story(self, story_prompt, game_state=None): super().start_new_story(story_prompt) - self.story.game_state["current_room"] = possible_rooms[0] return self.story.story_start @@ -178,7 +176,8 @@ class CTRLStoryManager(ConstrainedStoryManager): results = [] for phrase in self.action_phrases: options = dict() - options["used_verbs"] = set(used_verbs) + options["word_blacklist"] = {0: used_verbs} + options["word_whitelist"] = {0: get_possible_verbs()} result = self.generate_action_result(self.story_context(), phrase, options=options) used_verb = result[0].split()[1] @@ -187,23 +186,6 @@ class CTRLStoryManager(ConstrainedStoryManager): results.append(result) return results - def game_state_text(self): - current_room = self.story.game_state["current_room"] - text_list = ["You are currently in the ", current_room, ". You could go to the "] - for i in range(len(possible_rooms)): - if possible_rooms[i] is current_room: - continue - if i is len(possible_rooms) -1: - text_list.append(", or the ") - elif i is not 0: - text_list.append(", the ") - text_list.append(possible_rooms[i]) - text_list.append(".") - return "".join(text_list) - - def story_context(self): - return self.story.latest_result() + " " + self.game_state_text() - class CachedStoryManager(ConstrainedStoryManager):