normal web version works now. Now I need toupdate console versions and get unconstrained version working

This commit is contained in:
Max
2019-09-24 16:13:21 -06:00
parent 04dd210360
commit f008a336b9
6 changed files with 128 additions and 105 deletions
+5 -2
View File
@@ -24,7 +24,8 @@ def play_unconstrained():
generator = CTRLGenerator()
#generator = WebGenerator(CRED_FILE)
prompt = get_story_start("haunted")
story_manager = UnconstrainedStoryManager(generator, prompt)
story_manager = UnconstrainedStoryManager(generator)
story_manager.start_new_story(prompt)
print("\n")
console_print(str(story_manager.story))
@@ -45,7 +46,8 @@ def play_constrained():
story_start = "haunted"
verbs_key = "anything"
prompt = get_story_start(story_start)
story_manager = CTRLStoryManager(generator, prompt, action_verbs_key=verbs_key)
story_manager = CTRLStoryManager(generator)
story_manager.start_new_story(prompt, action_verbs_key=verbs_key)
console_print(str(story_manager.story))
possible_actions = story_manager.get_possible_actions()
@@ -69,6 +71,7 @@ def play_constrained():
def play_cached():
generator = WebGenerator(CRED_FILE)
story_manager = CachedStoryManager(generator, 0, 0, CRED_FILE)
story_manager.start_new_story()
console_print(str(story_manager.story))
possible_actions = story_manager.get_possible_actions()
+3
View File
@@ -32,6 +32,9 @@ class WebGenerator():
return response['predictions']
def generate(self, prompt, options={}):
print("Prompt to generate from is ", prompt)
while (True):
context_tokens = self.enc.encode(prompt)
try:
+29 -26
View File
@@ -12,13 +12,17 @@ import numpy as np
app = Flask(__name__)
app.secret_key = '#d\xe0\xd1\xfb\xee\xa4\xbb\xd0\xf0/e)\xb5g\xdd<`\xc7\xa5\xb0-\xb8d0S'
CRED_FILE = "./AI-Adventure-2bb65e3a4e2f.json"
generator = WebGenerator(CRED_FILE)
story_manager = CachedStoryManager(generator, CRED_FILE)
def get_response_string(story_text, possible_actions):
string_list = ["\n\n", story_text, "\n\nOptions:" + "\n"]
for i, action in enumerate(possible_actions):
string_list.append(str(i) + ") " + action + "\n")
string_list.append("\nWhich action do you choose? ")
# Initializes everything for a session
def story_init(session, seed):
session["generator"] = WebGenerator(CRED_FILE)
session["seed"] = seed
session["story_manager"] = CachedStoryManager(generator, 0, session["seed"], CRED_FILE)
response = "".join(string_list)
return response
# Shows about. (Should also link to paper when published)
@app.route('/about.html')
@@ -28,38 +32,37 @@ def about():
# Bread and butter of app, updates story and returns based on choice
@app.route('/generate', methods=['POST'])
def generate():
print("Entered generate")
action = request.form["action"]
if "story_manager" not in session:
print("not initialized")
# If there is no story in session, make a new one
if "story" not in session or session["story"] is None:
print("Starting new story")
seed = np.random.randint(100)
story_init(session, seed)
story_manager = session["story_manager"]
prompt = get_story_start("classic")
story_manager.start_new_story(prompt, seed)
possible_actions = story_manager.get_possible_actions()
string_list = [str(story_manager.story), "\n\nOptions:" + "\n"]
for i, action in enumerate(possible_actions):
string_list.append(str(i) + ") " + action)
response = "".join(string_list)
response = get_response_string(str(story_manager.story), possible_actions)
# If there is a story in session continue from it.
else:
print("initialized")
story_manager = session["story_manager"]
action = request.form["action"]
result, possible_actions = story_manager.act(action_choice)
if result is None:
response = "Invalid choice. Must be a number from 0 to 3. \n"
else:
string_list = [response]
for i, action in enumerate(possible_actions):
string_list.append(str(i) + ") " + action)
response = "".join(string_list)
print("Using existing story")
story = session["story"]
story_manager.load_story(story, from_json=True)
result, possible_actions = story_manager.act(action)
if result is None:
response = "\nInvalid choice. Must be a number from 0 to 3. \n" + "\nWhich action do you choose? "
else:
response = get_response_string(result, possible_actions)
session["story"] = story_manager.json_story()
print("Returning response")
return response
# Routes to index
@app.route('/')
def root():
session["story"] = None
return render_template('index.html')
if __name__ == '__main__':
+4 -3
View File
@@ -11,8 +11,8 @@ class cacher():
self.bucket = self.storage_client.get_bucket("dungeon-cache")
pass
def cache_file(self, seed, prompt_num, choices, response, tag, print_result=False):
def cache_file(self, seed, choices, response, tag, print_result=False):
prompt_num=0
blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag
for action in choices:
blob_file_name = blob_file_name + str(action)
@@ -22,7 +22,8 @@ class cacher():
if print_result: print("File ", blob_file_name, " cached")
def retrieve_from_cache(self, seed, prompt_num, choices, tag, print_result=False):
def retrieve_from_cache(self, seed, choices, tag, print_result=False):
prompt_num = 0
blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag
for action in choices:
+10 -9
View File
@@ -1,4 +1,8 @@
start_text = "<span id='a'>Adventurer@AIDungeon</span>:<span id='b'>~</span><span id='c'>$</span> ./EnterDungeon \n <br/><!-- laglaglaglaglaglaglaglaglaglaglag-->"
start_text = "<span id='a'>Adventurer@AIDungeon</span>:<span id='b'>~</span><span id='c'>$</span> ./EnterDungeon <br/><!-- laglaglaglaglaglaglaglaglaglaglag-->"
function isMobileDevice() {
return /Android|webOS|iPhone|iPad|iPod|BlackBerry|IEMobile|Opera Mini/i.test(navigator.userAgent)
};
// Used to control the terminal like screen typing
var Typer={
@@ -76,12 +80,12 @@ var Typer={
onKeyPressFunc:function(evt) {
if(acceptInput && !isMobileDevice()){
if(Typer.acceptInput && !isMobileDevice()){
evt = evt || window.event
var charCode = evt.keyCode || evt.which
if(charCode == 13){
acceptInput = false
Typer.acceptInput = false
Typer.sendInput()
}
else{
@@ -94,13 +98,12 @@ var Typer={
},
}
function onButtonClick(num){
document.getElementById('buttons').style.visibility='hidden';
if (acceptInput == true){
acceptInput = false
if (Typer.acceptInput == true){
Typer.acceptInput = false
num = String(num)
Typer.appendToText(num)
Typer.inputStr = num
@@ -121,9 +124,7 @@ function start(){
Typer.appendToText(start_text)
Typer.startTyping()
request_str = ""
$.post("/generate", {text: request_str}, receiveResponse)
console.log("Not mobile device");
$.post("/generate", {action: request_str}, receiveResponse)
document.getElementById('buttons').style.visibility='hidden';
}
+77 -65
View File
@@ -5,8 +5,7 @@ import json
class Story():
def __init__(self, story_start):
def __init__(self, story_start, seed=None):
self.story_start = story_start
# list of actions. First action is the prompt length should always equal that of story blocks
@@ -15,6 +14,20 @@ class Story():
# list of story blocks first story block follows prompt and is intro story
self.results = []
# Only needed in constrained/cached version
self.seed = seed
self.choices = []
self.possible_action_results = []
def initialize_from_json(self, json_string):
story_dict = json.loads(json_string)
self.story_start = story_dict["story_start"]
self.seed = story_dict["seed"]
self.actions = story_dict["actions"]
self.results = story_dict["results"]
self.choices = story_dict["choices"]
self.possible_action_results = story_dict["possible_action_results"]
def add_to_story(self, action, story_block):
self.actions.append(action)
self.results.append(story_block)
@@ -35,20 +48,41 @@ class Story():
return "".join(story_list)
def to_json(self):
story_dict = {}
story_dict["story_start"] = self.story_start
story_dict["seed"] = self.seed
story_dict["actions"] = self.actions
story_dict["results"] = self.results
story_dict["choices"] = self.choices
story_dict["possible_action_results"] = self.possible_action_results
return json.dumps(story_dict)
class StoryManager():
def __init__(self, generator, story_prompt):
def __init__(self, generator):
self.generator = generator
self.story_prompt = story_prompt
def init_story(self):
block = self.generator.generate(self.story_prompt)
def start_new_story(self, story_prompt):
block = self.generator.generate(story_prompt)
block = cut_trailing_sentence(block)
block = story_replace(block)
story_start = self.story_prompt + block
story_start = story_prompt + block
self.story = Story(story_start)
return story_start
def load_story(self, story, from_json=False):
if from_json:
self.story = Story("")
self.story.initialize_from_json(story)
else:
self.story = story
return str(story)
def json_story(self):
return self.story.to_json()
def story_context(self):
return self.story.latest_result()
@@ -56,10 +90,6 @@ class StoryManager():
class UnconstrainedStoryManager(StoryManager):
def __init__(self, generator, story_prompt):
super().__init__(generator, story_prompt)
self.init_story()
def act(self, action_choice):
result = self.generate_result(action_choice)
self.story.add_to_story(action_choice, result)
@@ -71,21 +101,27 @@ class UnconstrainedStoryManager(StoryManager):
block = story_replace(block)
return block
class ConstrainedStoryManager(StoryManager):
def __init__(self, generator, story_prompt, action_verbs_key="classic"):
super().__init__(generator, story_prompt)
self.init_story()
self.possible_action_results = None
def __init__(self, generator, action_verbs_key="classic"):
self.generator = generator
self.action_phrases = get_action_verbs(action_verbs_key)
def get_possible_actions(self):
if self.possible_action_results is None:
self.possible_action_results = self.get_action_results()
def start_new_story(self, story_prompt):
super().start_new_story(story_prompt)
self.story.possible_action_results = self.get_action_results()
return story.story_start
return [action_result[0] for action_result in self.possible_action_results]
def load_story(self, story, from_json=False):
story_string = super().load_story(story, from_json=from_json)
return story_string
def get_possible_actions(self):
if self.story.possible_action_results is None:
self.story.possible_action_results = self.get_action_results()
return [action_result[0] for action_result in self.story.possible_action_results]
def act(self, action_choice_str):
@@ -99,9 +135,10 @@ class ConstrainedStoryManager(StoryManager):
print("Error invalid choice.")
return None, None
action, result = self.possible_action_results[action_choice]
self.story.choices.append(action_choice)
action, result = self.story.possible_action_results[action_choice]
self.story.add_to_story(action, result)
self.possible_action_results = self.get_action_results()
self.story.possible_action_results = self.get_action_results()
return result, self.get_possible_actions()
def get_action_results(self):
@@ -122,9 +159,8 @@ class ConstrainedStoryManager(StoryManager):
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 __init__(self, generator, action_verbs_key="anything"):
super().__init__(generator, action_verbs_key)
def get_action_results(self):
@@ -144,52 +180,28 @@ class CTRLStoryManager(ConstrainedStoryManager):
class CachedStoryManager(ConstrainedStoryManager):
def __init__(self, generator, prompt_num, seed, credentials_file, action_verbs_key="classic", ):
def __init__(self, generator, credentials_file, action_verbs_key="classic"):
super().__init__(generator, action_verbs_key=action_verbs_key)
self.cacher = cacher(credentials_file)
prompt = get_story_start("classic")
super().__init__(generator, prompt, action_verbs_key)
self.seed = seed
self.prompt_num = prompt_num
self.choices = []
result = self.cacher.retrieve_from_cache(seed, prompt_num, [], "story")
def start_new_story(self, prompt, seed):
result = self.cacher.retrieve_from_cache(seed, [], "story")
if result is not None:
story_start = result
self.story = Story(story_start)
story_start = prompt + result
self.story = Story(story_start, seed=seed)
else:
story_start = super().start_new_story(prompt)
self.story.seed = seed
self.cacher.cache_file(seed, [], story_start, "story")
story_start = self.init_story()
self.cacher.cache_file(seed, prompt_num, [], story_start, "story")
self.story.possible_action_results = None
self.possible_action_results = None
def get_possible_actions(self):
if self.possible_action_results is None:
self.possible_action_results = self.get_action_results()
return [action_result[0] for action_result in self.possible_action_results]
def act(self, action_choice_str):
try:
action_choice = int(action_choice_str)
except:
print("Error invalid choice.")
return None, None
if action_choice < 0 or action_choice >= len(self.action_phrases):
print("Error invalid choice.")
return None, None
self.choices.append(action_choice)
action, result = self.possible_action_results[action_choice]
self.story.add_to_story(action, result)
self.possible_action_results = self.get_action_results()
return result, self.get_possible_actions()
return story_start
def get_action_results(self):
response = self.cacher.retrieve_from_cache(self.seed, self.prompt_num, self.choices, "choices")
response = self.cacher.retrieve_from_cache(self.story.seed, self.story.choices, "choices")
if response is not None:
action_results = json.loads(response)
@@ -197,7 +209,7 @@ class CachedStoryManager(ConstrainedStoryManager):
print("Not found in cache. Generating...")
action_results = super().get_action_results()
response = json.dumps(action_results)
self.cacher.cache_file(self.seed, self.prompt_num, self.choices, response, "choices")
self.cacher.cache_file(self.story.seed, self.story.choices, response, "choices")
return action_results