mirror of
https://github.com/wassname/Clover-Edition.git
synced 2026-09-09 11:13:26 +08:00
update
This commit is contained in:
+5
-5
@@ -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()
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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
@@ -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", ):
|
||||
|
||||
Reference in New Issue
Block a user