This commit is contained in:
Nick Walton
2019-09-21 16:57:36 -06:00
parent 0fe73ab7f6
commit a5bdaa24e4
5 changed files with 51 additions and 13 deletions
+5 -5
View File
@@ -4,7 +4,7 @@ from google.cloud import storage
import json
from story.story_manager import *
from generator.web.web_generator import *
from generator.ctrl.ctrl_generator import *
#from generator.ctrl.ctrl_generator import *
import tensorflow as tf
import textwrap
@@ -40,12 +40,12 @@ def play_unconstrained():
def play_constrained():
print("\n")
#generator = WebGenerator(CRED_FILE)
generator = CTRLGenerator()
generator = WebGenerator(CRED_FILE)
#generator = CTRLGenerator()
story_start = "haunted"
verbs_key = "anything"
prompt = get_story_start(story_start)
story_manager = ConstrainedStoryManager(generator, prompt, action_verbs_key=verbs_key)
story_manager = CTRLStoryManager(generator, prompt, action_verbs_key=verbs_key)
console_print(str(story_manager.story))
possible_actions = story_manager.get_possible_actions()
@@ -87,7 +87,7 @@ def play_cached():
if __name__ == '__main__':
play_unconstrained()
play_constrained()
+18 -4
View File
@@ -139,7 +139,23 @@ class CTRLGenerator():
self.topk = 0
def generate(self, prompt, first_verb_whitelist=True):
def configure_verb_probs(self, probabilities, options):
# Make sure only a possible verb is chosen.
for word in get_possible_verbs():
probabilities[self.word2idx[word]] += 100
# Disallow used verbs
if "used_verbs" in options:
for verb in options["used_verbs"]:
probabilities[self.word2idx[verb]] = -1e8
return probabilities
def generate(self, prompt, options=None):
if options is None:
options = {}
if prompt[-1] != " ":
prompt = prompt + " "
@@ -217,10 +233,8 @@ 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():
prompt_logits[_token][self.word2idx[word]] += 5
prompt_logits[_token] = self.configure_verb_probs(prompt_logits[_token], options)
# compute probabilities from logits
prompt_probs = np.exp(prompt_logits[_token])
+1 -1
View File
@@ -41,7 +41,7 @@ class TFGenerator():
ckpt = tf.train.latest_checkpoint(model_path)
saver.restore(self.sess, ckpt)
def generate(self, prompt):
def generate(self, prompt, options={}):
context_tokens = self.enc.encode(prompt)
out = self.sess.run(self.output, feed_dict={
self.context: [context_tokens for _ in range(1)]
+1 -1
View File
@@ -31,7 +31,7 @@ class WebGenerator():
return response['predictions']
def generate(self, prompt):
def generate(self, prompt, options={}):
while (True):
context_tokens = self.enc.encode(prompt)
try:
+26 -2
View File
@@ -107,8 +107,11 @@ class ConstrainedStoryManager(StoryManager):
def get_action_results(self):
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)
def generate_action_result(self, prompt, phrase, options=None):
if options is None:
options = {}
action = phrase + " " + self.generator.generate(prompt + phrase, options)
action_result = cut_trailing_sentence(action)
action, result = split_first_sentence(action_result)
@@ -118,6 +121,27 @@ class ConstrainedStoryManager(StoryManager):
return action, result
class CTRLStoryManager(ConstrainedStoryManager):
def __init__(self, generator, story_prompt, action_verbs_key="classic"):
super().__init__(generator, story_prompt)
self.action_phrases = get_action_verbs("anything")
def get_action_results(self):
used_verbs = []
results = []
for phrase in self.action_phrases:
options = {}
options["used_verbs"] = used_verbs
result = self.generate_action_result(self.story_context(), phrase, options=options)
used_verb = result[0].split()[1]
used_verbs.append(used_verb)
results.append(result)
return results
class CachedStoryManager(ConstrainedStoryManager):
def __init__(self, generator, prompt_num, seed, credentials_file, action_verbs_key="classic", ):