From 572b043d82732480d4068d6ed89112f84a7e4dcf Mon Sep 17 00:00:00 2001 From: Nick Date: Sat, 14 Sep 2019 20:22:28 -0600 Subject: [PATCH] constrained and unconstrained console work now --- .gitignore | 2 +- console_play.py | 54 ++++++++- .../__pycache__/web_generator.cpython-36.pyc | Bin 1553 -> 1628 bytes .../__pycache__/web_generator.cpython-37.pyc | Bin 0 -> 1632 bytes main.py | 105 +----------------- .../__pycache__/story_manager.cpython-36.pyc | Bin 3591 -> 4112 bytes .../__pycache__/story_manager.cpython-37.pyc | Bin 3582 -> 3590 bytes story/story_manager.py | 27 ++++- 8 files changed, 80 insertions(+), 108 deletions(-) create mode 100644 generator/web/__pycache__/web_generator.cpython-37.pyc 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 10f0f74b5b1db261ed4198331cc569e30e16cf74..056b43dfc0f8a107e5acd1005cf59604bc66b55e 100644 GIT binary patch delta 503 zcmZ8eJ4*vW5Wc<3C6|X!@PTLw!6LPSjUXWgiH1Z>3W?})+>LsA?jEy?SR`23*~t}_ zVrlJ<5f%#@!Op+nZepP>%(pW;-^_f&-k07=`PH$p!u!ME=`4s2_59DNz9X%mYY z`b1jiBqSSFXb)^@_w8HKx2O|Z9joHBRC&AI-fq@ub#JfPs8zd-c8k{b>s$3!w^40& z(si=5lTkJoDhKQ;Le6@{Q;-x`z(YXgXs6<+jDYy0HyI#=z+~jdf}TV%^a_`n71P9* zXU`Vq(;L#sWwk7e6kcA*fb4qOSK|R^60Rkc`~sLBB>=CY0C2)&y@(0vvjCKbAQH}W zHm?&SC>U>E!({sCd>-daMfOLXc~(vgOl3nbdX<}C#ALDqb#gv6iufp@y-Uv4$a@F@+(R zL6gaEaw4OwHd7RLdTL30YF=`FN@~$9W*|e8=@v_IYED`dS1wR2C%z!DBx5B*5!>Vw zjKQodK+_mDD>2nFGO|uyz%0wiGkG7gga{i@v9SPF_#Q;MX4Oo$n9gKlx8r{<*=C6=TrPBvhTHsJs%0U2JT zgsK>>qlg(O!3o60AR!JW4wlLLS?wA5C;w*^R{;e@kuZo50TCdjNESlOm~6$S1ON$a BL-qgw 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 0000000000000000000000000000000000000000..6ebdc86271262aff467b9214bd7f959b3127a1cb GIT binary patch literal 1632 zcmZ8h&u`l{6eg)(w5_IH(R4{M6vnU~JZx4|3@9)ZTeCC>0%T6n1n4SY&=iq&l*p1o zQBI=Dm!-Q5*m+myxPOVK0Xywq*lCZ7on@i%9-sO1=N=J+O~zQtO;y@gntbv_{0xgfs7$R6ss zf5U#%!wP>ssM_zcvKvYnM#mM^vE)foIXWr}UfDVrWtAJoUxsYLQ&|e`yCa@*5z0&i zlQo}CxG3T*eHQ6+7mYyh60G(@Km8|(YqyMtik z&E~7k!T$F8&R+GX?_EE@lS6NuF`g9O+S>E)d!`~;p&Xu0Vv(hOXT)W|(+E}wWkAp~ zlv(g(=--P3XTV1sCPgre6YjU0bAmjShkqc^0-*nV>m6oe?xk^bd>SS)OP@aVMzQqP zw_la%h-c{yAgBDm0}HO9_49K@gCI_083Y~ucpDtTZEWId`6gBU#j0!0^h(Gd)NqzV zpc$Jr{9giR%$OO{lorHBji;ux89Hu>uWHZ*X6#G+9Kj5kIhW{&!YE``ZN_+nW6Yf2 z$gFjVrWT{Mb!>}6*{&g8MHl#I3{h)U>(Cx3GENkktX0do$c}g<3-A4OLH7G;<$CkP z(c^wAP74{P5eN38ES3CB_J}gmaLko?kg;=RaUn9{cXA*;jwEDU(d($YqbwUGJj~-L zi9rv)#o{8$;Aq%!P{s*U_AnFUP%3ARLs^Yd%9-Eq(*;1iq&eJo_!*B%$(8%21d+#l zQ(xY*1qCwI(i~wRvtynD;)07w9C2}1&(g(H6cxP40Vi+L;mYX(7c1cIffWzIA-sfL zY!Zr>zzBY5Q0$WG!T%(l1G`pOtE*LR6#fZ&@DLHD+z@A>90=#f%G#(J6U{PRI{nZ{aHgikbs~;T~mYrOZoN zJ(*`J?+SiW@6J{}d*R2jCDe=YqBIie?krnBFdv#hyXW)|I7mD+p9O z2p}KIM8|FroRnd5H6wJ zPJzFAD`)%HDJ$;7UjL&!D}t~|@c@cK$U-c`xX{TWKGsfOaRW+&Oyg1X#noHB(NYyI MIEpFJe87+V3t~O8>;M1& literal 0 HcmV?d00001 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 a283cdce48f53a992d3d2b466d564a8c8c1389d1..36811e193c3099eab6ca5740b3d9de52b70b2829 100644 GIT binary patch delta 1007 zcmZXS-%k@k5XX1-esp`iYpj8m3Sy05s6YixeDOz!KOe*rCB}qEF3>G-g|v0o>@ z3bcD@;S_>$S8P|p2+sI`a!YS_kb3U30mmX3^q9dk9S}paMg@OnY|Q!qmRAP+P7MEY+x%0PR>kP za5E;C9I!2cwX^EI(FXt>s08gd8+(DOGq z+az~dH{P43PSZb~k5h03mlJ!~#wVVN$_jIcE&VC?gwN>Txe1F>Yih^+mA^)AIv-{c xl!Ud2lN~+oJ)F%#i9D1HlnD-jt^|}RTqdN~T-1l&MVEA>kdCxvL4Wg}{{uPE&~E?$ delta 523 zcmZ9IOD_Xa6vywqGo87;otYw-*0UiV)u0JG3lfPHLIfeC!IV}+a(^Sbk&xy)R2s4>@7mFwE(ZQ@?N+k71I=CBv2Gq!FA zi6xeG>OL$gA0pFrOvczWM!3Kj+?b3ug{_@xxpG+0#qju`8rJ=zN*I>64~qXA-p}AsNpFQk7>Q3bI9$>r6t}qI zCmHmWg7a8oiF_Hq-ts8W5sP@5ZHStv8JO{X5sAn+6`Xp;ms=K5hwTddN(;mbEY~*V zRGP2CojcbvgsW7-L(@IKI!*<8?+Ubfuc={59mKl(=eI$((-O;JeGNL3GyPKgG?*f! h35q~Qbw7bF)H!%d`ptR(^(00}Gg8QewBa$i{0UpLZn^*f diff --git a/story/__pycache__/story_manager.cpython-37.pyc b/story/__pycache__/story_manager.cpython-37.pyc index f35b53e9f44c712c44c10beebeda565fe12fcef7..098f93fb1943cd6d4038be1c494d8286e8b5f08a 100644 GIT binary patch delta 248 zcmew--6q59#LLUY00j9{YhxF03y7D{1U zz*fV!kTIA+lgY1$AE;-tKC?0_OIChn-sCW5Jw~a?t<0s2ER%mTPiA4PVThj`$Du#@ zAWN61Do|}P56BiqHYOfM4iIFToWQz zE)F&h<|4_-yeyKNKd~1xGCEF9<2=sjGg*(TnlX6tDz4{z;UMKqK!S~-%FcOmBD>_| m9o!y_VUrnoUNGL8{EA1NF>3y<}P8Y zVQgj$X3%8xn{39c%)(q;nmaj;S&vb2axZf!BlBctmdTUnu;>Dn)iA^}Og_odC8Yv1 zxtIrJJ|hn!2M988F;*#s<|R)yobtj9mwK5t{p3Sy{(QI-E zdljz|NHYhJ;9%n5+f`PG0b+Q4k#N=yShZ#dB n&*IKw44KTy^MWyI@+Tg3#_-8fyd8|TlNa;Gu`=<=^GO2$!u&gc 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()