diff --git a/.gitignore b/.gitignore index 0efd241..8d6f8df 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,3 @@ **/__pychache__ .idea -RL +*.json diff --git a/console_play.py b/console_play.py index cf46c81..911eba1 100644 --- a/console_play.py +++ b/console_play.py @@ -14,17 +14,67 @@ def console_print(str): print((textwrap.fill(str, 80))) -if __name__ == '__main__': +def play_unconstrained(): generator = WebGenerator("./AI-Adventure-2bb65e3a4e2f.json") prompt = "You enter a dungeon with your trusty sword and shield. You are searching for the evil necromancer who killed your family. You've heard that he resides at the bottom of the dungeon, guarded by legions of the undead. You enter the first door and see" story_manager = UnconstrainedStoryManager(generator, prompt) console_print(str(story_manager.story)) - while(True): + while (True): action = input("> ") action = "You " + action result = story_manager.act(action) console_print(action + result) + # + # + # + # def act(self, 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.possible_action_results + # + # def story_context(self): + # return self.story.latest_result() + # + # 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) + # action_result = cut_trailing_sentence(action) + # + # action, result = split_first_sentence(action_result) + # result = story_replace(action_result) + # action = action_replace(action) + # + # return action, result + +def play_constrained(): + generator = WebGenerator("./AI-Adventure-2bb65e3a4e2f.json") + prompt = "You enter a dungeon with your trusty sword and shield. You are searching for the evil necromancer who killed your family. You've heard that he resides at the bottom of the dungeon, guarded by legions of the undead. You enter the first door and see" + story_manager = ConstrainedStoryManager(generator, prompt) + + console_print(str(story_manager.story)) + possible_actions = story_manager.get_possible_actions() + while (True): + console_print("\nOptions:") + for i, action in enumerate(possible_actions): + console_print(str(i) + ") " + action) + + result = None + while(result == None): + action_choice = input("Which action do you choose? ") + print("\n") + result, possible_actions = story_manager.act(action_choice) + + console_print(result) + + +if __name__ == '__main__': + play_constrained() + diff --git a/generator/web/__pycache__/web_generator.cpython-36.pyc b/generator/web/__pycache__/web_generator.cpython-36.pyc index 10f0f74..056b43d 100644 Binary files a/generator/web/__pycache__/web_generator.cpython-36.pyc and b/generator/web/__pycache__/web_generator.cpython-36.pyc differ diff --git a/generator/web/__pycache__/web_generator.cpython-37.pyc b/generator/web/__pycache__/web_generator.cpython-37.pyc new file mode 100644 index 0000000..6ebdc86 Binary files /dev/null and b/generator/web/__pycache__/web_generator.cpython-37.pyc differ diff --git a/main.py b/main.py index c2cb365..bda8215 100644 --- a/main.py +++ b/main.py @@ -1,119 +1,24 @@ from flask import g from flask import session import os -import googleapiclient.discovery from story.utils import * -from google.cloud import storage import json from flask import Flask, render_template, request, abort -from generator import StoryGenerator -import gpt2.src.encoder as encoder +from story.story_manager import * +from generator.web.web_generator import * +from other.caching import * - -# App Info -phrases = [" You attack", " You use", " You tell", " You go"] -prompts = ["You enter a dungeon with your trusty sword and shield. You are searching for the evil necromancer who killed your family. You've heard that he resides at the bottom of the dungeon, guarded by legions of the undead. You enter the first door and see"] -continuing_prompts = ["You are in a dungeon with your sword and shield. You are on a quest to defeat the necromancer. This dungeon is full of zombie and skeletons."] app = Flask(__name__) app.secret_key = '#d\xe0\xd1\xfb\xee\xa4\xbb\xd0\xf0/e)\xb5g\xdd<`\xc7\xa5\xb0-\xb8d0S' -# Encoder Info -encoder_path='gpt2/models/117M' -enc = encoder.get_encoder(encoder_path) -# Model/Cache Info -project = "ai-adventure" -model = "generator_v1" -version = "version2" -os.environ['GOOGLE_APPLICATION_CREDENTIALS']="./AI-Adventure-2bb65e3a4e2f.json" -storage_client = storage.Client() -bucket = storage_client.get_bucket("dungeon-cache") - -# Local generator functionality -RUN_LOCAL = False -local_generator = None -def get_local_generator(): - if "gen" not in g: - if "sess" not in g: - g.sess = tf.Session() - g.gen = StoryGenerator(g.sess) - - return g.gen - - -@app.teardown_appcontext -def teardown_sess(_): - sess = g.pop("sess", None) - - if sess is not None: - sess.close() - -def predict(context_tokens): - service = googleapiclient.discovery.build('ml', 'v1') - name = 'projects/{}/models/{}'.format(project, model) - instance = context_tokens - - if version is not None: - name += '/versions/{}'.format(version) - - response = service.projects(). predict( - name=name, - body={'instances': [{'context': instance}]} - ).execute() - - if 'error' in response: - raise RuntimeError(response['error']) - - return response['predictions'] - - -def generate(prompt): - - while(True): - context_tokens = enc.encode(prompt) - try: - pred = predict(context_tokens) - pred = pred[0]["output"][len(context_tokens):] - output = enc.decode(pred) - return output - except: - print("generate request failed, trying again") - continue - - -def generate_story_block(prompt, local=False): - - if local: - generator = get_local_generator() - block = generator.generate(prompt) - else: - block = generate(prompt) - - block = cut_trailing_sentence(block) - block = story_replace(block) - return block - - -def generate_action_result(prompt, phrase, local=False): - - if local: - generator = get_local_generator() - action = phrase + generator.generate(prompt + phrase) - else: - action = phrase + generate(prompt + phrase) - - action_result = cut_trailing_sentence(action) - action_result = story_replace(action_result) - action = first_sentence(action) - - return action, action_result - @app.route('/') def root(): seed = -1 data = {'seed': seed} return render_template('index.html', data=data) + @app.route('/') def rootseed(seed): if seed == "": @@ -124,11 +29,13 @@ def rootseed(seed): session["seed"] = seed return render_template('index.html', data=data) + @app.route('/index.html') def index(): data = {'seed': -1} return render_template('index.html', data=data) + @app.route('/about.html') def about(): return render_template('about.html') diff --git a/story/__pycache__/story_manager.cpython-36.pyc b/story/__pycache__/story_manager.cpython-36.pyc index a283cdc..36811e1 100644 Binary files a/story/__pycache__/story_manager.cpython-36.pyc and b/story/__pycache__/story_manager.cpython-36.pyc differ diff --git a/story/__pycache__/story_manager.cpython-37.pyc b/story/__pycache__/story_manager.cpython-37.pyc index f35b53e..098f93f 100644 Binary files a/story/__pycache__/story_manager.cpython-37.pyc and b/story/__pycache__/story_manager.cpython-37.pyc differ diff --git a/story/story_manager.py b/story/story_manager.py index f59d5ba..80acb9e 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -63,22 +63,37 @@ class UnconstrainedStoryManager(): class ConstrainedStoryManager(): def __init__(self, generator, story_prompt): + self.generator = generator + self.action_phrases = ["You attack", "You tell", "You use", "You go"] block = self.generator.generate(story_prompt) block = cut_trailing_sentence(block) block = story_replace(block) story_start = story_prompt + block - self.story = Story(story_start) - self.generator = generator - self.possible_action_results = self.get_action_results() - self.action_phrases = ["You attack", "You tell", "You use", "You go"] + self.possible_action_results = None - def act(self, action_choice): + 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 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.possible_action_results + return result, self.get_possible_actions() def story_context(self): return self.story.latest_result()