From ef3a8ca864b28583d477c1fbf7c02327cb77121d Mon Sep 17 00:00:00 2001 From: Nick Date: Mon, 16 Sep 2019 19:40:54 -0600 Subject: [PATCH] got cached version working --- console_play.py | 32 ++- generator/__pycache__/__init__.cpython-36.pyc | Bin 137 -> 0 bytes generator/__pycache__/__init__.cpython-37.pyc | Bin 141 -> 0 bytes .../generator_local.cpython-37.pyc | Bin 3102 -> 0 bytes .../src/__pycache__/__init__.cpython-36.pyc | Bin 144 -> 0 bytes .../src/__pycache__/__init__.cpython-37.pyc | Bin 154 -> 0 bytes .../tf/src/__pycache__/encoder.cpython-36.pyc | Bin 4982 -> 0 bytes .../tf/src/__pycache__/encoder.cpython-37.pyc | Bin 4977 -> 0 bytes .../tf/src/__pycache__/model.cpython-37.pyc | Bin 6724 -> 0 bytes .../tf/src/__pycache__/sample.cpython-37.pyc | Bin 2606 -> 0 bytes .../__pycache__/web_generator.cpython-36.pyc | Bin 1628 -> 0 bytes .../__pycache__/web_generator.cpython-37.pyc | Bin 1632 -> 0 bytes main.py | 101 +------ old_main.py | 264 ++++++++++++++++++ other/cacher.py | 59 ++-- .../__pycache__/story_manager.cpython-36.pyc | Bin 4112 -> 0 bytes .../__pycache__/story_manager.cpython-37.pyc | Bin 3590 -> 0 bytes story/__pycache__/utils.cpython-36.pyc | Bin 2095 -> 0 bytes story/__pycache__/utils.cpython-37.pyc | Bin 2099 -> 0 bytes story/story_manager.py | 99 ++++--- 20 files changed, 395 insertions(+), 160 deletions(-) delete mode 100644 generator/__pycache__/__init__.cpython-36.pyc delete mode 100644 generator/__pycache__/__init__.cpython-37.pyc delete mode 100644 generator/tf/__pycache__/generator_local.cpython-37.pyc delete mode 100644 generator/tf/src/__pycache__/__init__.cpython-36.pyc delete mode 100644 generator/tf/src/__pycache__/__init__.cpython-37.pyc delete mode 100644 generator/tf/src/__pycache__/encoder.cpython-36.pyc delete mode 100644 generator/tf/src/__pycache__/encoder.cpython-37.pyc delete mode 100644 generator/tf/src/__pycache__/model.cpython-37.pyc delete mode 100644 generator/tf/src/__pycache__/sample.cpython-37.pyc delete mode 100644 generator/web/__pycache__/web_generator.cpython-36.pyc delete mode 100644 generator/web/__pycache__/web_generator.cpython-37.pyc create mode 100644 old_main.py delete mode 100644 story/__pycache__/story_manager.cpython-36.pyc delete mode 100644 story/__pycache__/story_manager.cpython-37.pyc delete mode 100644 story/__pycache__/utils.cpython-36.pyc delete mode 100644 story/__pycache__/utils.cpython-37.pyc diff --git a/console_play.py b/console_play.py index 1a4543d..4ca22ab 100644 --- a/console_play.py +++ b/console_play.py @@ -7,6 +7,9 @@ from generator.web.web_generator import * import tensorflow as tf import textwrap +CRED_FILE = "./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" + # Set the key def console_print(str): @@ -15,8 +18,7 @@ def console_print(str): 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" + generator = WebGenerator(CRED_FILE) story_manager = UnconstrainedStoryManager(generator, prompt) console_print(str(story_manager.story)) @@ -26,9 +28,9 @@ def play_unconstrained(): result = story_manager.act(action) console_print(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" + generator = WebGenerator(CRED_FILE) story_manager = ConstrainedStoryManager(generator, prompt) console_print(str(story_manager.story)) @@ -47,8 +49,28 @@ def play_constrained(): console_print(result) +def play_cached(): + generator = WebGenerator(CRED_FILE) + story_manager = CachedStoryManager(generator, 0, 0, CRED_FILE) + + 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() + play_cached() diff --git a/generator/__pycache__/__init__.cpython-36.pyc b/generator/__pycache__/__init__.cpython-36.pyc deleted file mode 100644 index 4dffe2fd29e460638fe085e95089bb5116ec81f0..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 137 zcmXr!<>gxbwmg;r2p)q77+?f49Dul(1xTbY1T$zd`mJOr0tq9CUuOCl`MIh3d6~)C z<%u~Z`FREg`stY^`i`D1rFrS8`FZ;3sd=eIi6!|(`tk9Zd6^~g@p=W7w>WHa^HWN5 MQtd$I6$3E?0C(>p8vpg`kg5__^V?p#|5CH>>K!yVl7qb9~6oz01O-8?!3`HPe1o6vEKO;XkRX;B? zIlDYDrzAhmz(7Aevqay~)1@>oJvBd1KRq=swJ5P9zeqnmJ~J<~BtBlRpz;=nO>TZl OX-=vg$h=}8W&i+zk|8Pp diff --git a/generator/tf/__pycache__/generator_local.cpython-37.pyc b/generator/tf/__pycache__/generator_local.cpython-37.pyc deleted file mode 100644 index 8c7f27944b6f6e546d63b1e6dfaa7d39868f8fae..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 3102 zcmZuzTaO$^6|SnjPS0hp_ImBa7zlw$VApm)e!#K~Sh0Yx66^!eQfQ5*duC_0r@LKU zL*h}-(;5*Ttbmjs(EG^$;1Tiz>NO9@Qhq_6_)hgqypE|+pE-3uRp&e3sh@Yd0YmxB z|Mf zzb9K~4DuVHm@S={FWp!yy^2>XVv5(I^3F-OlwCXLZ0Rdk zc~ie`RbFhJi!=V|p->ZPSrZ+=bW?OW3At^-?2Y1)vbHyxEr-mZ>o*yR^`TQ zEY1a1+*Ui&-TG4H&{@?dLYe?qlxN`Wp5RjP?duwQN^G{U*k- zs(miVf4BnY=mp1F7MtTj-99ufC=YMly7QYw=cE|O!K^%(mU-5)R z&s6fU(!t|_8lFVuWRdi*Yr9O!QhP}@gg>Qjj}sNGD;5+fu%F(DlhI(Fs%X8cofk=_ zMP6#V7+`gKnolylo&PzJauO%eNanMMV1XVOr0p~x#M&z2$GTNw3nd4`L<=?2?IMMG zlBcx2Gs)C@w`H656j`BNImqICMph=V-k_OHOJo|{axg2>1RdrdjZfyVIni!DSH)au z2m8P-=Xe0a^j11hNvWdYNisYwV6f7jL_;odEmfA#+1Nj^&sv&r!Eqd}_j?B>nG@kAZ| z?3WMc**M9w!|@uYMN1<5$34$|Nw!m?o zvXwJ+YY+VAtT@)MyxLc-I;i~DIF{9I#ipGK8IQPi*;3tUuimI^!>p{{TxnZiRIg%H z3nSgQ_u!BSB_{tR zX)}Ys))%2W$v|&q=ruf7_OHon=+WLW*nE;_hP5S8h1@1V0Va1zTq1Fq#2$$&B;FwL zCJD;I@>?WskoY!Juvp6b&tIe@%*l z`=W>54(|y0KE{6yXQFTdx`Fc)3c*={jn25umK?AaO95zG=FwmQI<~SvilQ>~w{$8z zAOL4R#uI}8Sh*8^#_`m!YZs3RZy!%BLxR5YD*p_Rg34F^NPrMqN8|TQ`QrYcWXr2xR<$x?2Rs`Q%+&=BIRpe#V+(cBW xfmEg8LW3Ygw%Q4z}k1dl-k3@`#24nSPY0whuxf*CX!{Z=v*frJsnFI)YL{M=Oiyv*e6 z^2D5y{5%5#{q)QdeMe81(!BK4{5<{i)V$Q9#FG3X{gO2O;-X~z`1s7c%#!$cy@JYH T95%W6DWy57b|7PmftUdR0T3d= diff --git a/generator/tf/src/__pycache__/__init__.cpython-37.pyc b/generator/tf/src/__pycache__/__init__.cpython-37.pyc deleted file mode 100644 index 4f376a5f534d369102ee9771d53fef9af4d0740a..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 154 zcmZ?b<>g`kf>j(9u^{>}h=2h`Aj1KOi&=m~3PUi1CZpdPhpPrhRT9jClU!-4>7N3)!oS36uT$HRIAD@|* bSrQ+wS5SG2!zMRBr8Fni4rF*S5HkP({aYnK diff --git a/generator/tf/src/__pycache__/encoder.cpython-36.pyc b/generator/tf/src/__pycache__/encoder.cpython-36.pyc deleted file mode 100644 index 22a67d44ac68f683d5becdae6f3e73f33d7684a5..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 4982 zcmbtYTaz0{74DuHjYe8|vm{$~0w$SU$Xa-_yCH-GYHXC`FYv1uv;U@e6pQisC8nRKdUKSDyR>cyUn$-|5k472_8`_VmnjpYA?=&Ue0Z zW?#8*q3-qf?xqNML1 zs3?w96lV9aH`{A+U5n*GyBl-|VVAn-=iqX2y^SmT9F1a%E2uT2)|x@BOKgVwyx%eX zmGk~fkE2Za zX%ANDz4W@@A@c-@@B_b_N*QL8G%-e6h$Rus7W$czuwAxg*81@v^ksM$$}H-{q2HN= zm@UH0KT74uPZR&((X*n=ABO>yI7tuuBn)ZWIFS9&f8MAxpVb`%GUzHS??(yAO8+qJ z1|2`@nVCmHqI{M5!C@K+|0jFz&DFX+2#xrL#F>x+uAC||*LGVBDv#6i+qAuw(2_=a_ zh>l{EPEt7z;!QsbPB#5lzVV(hlsU`9WiPZIY%G+lb7AN5cm-NKcvqv-?q_ZomH zL6rZ_+IMcTKQct1-5toLrLC;&u|L)IA%0m?NGO_$N+n4JC#DaFfWj!9T9bvwoju;)6rHZ(^fZ=k71VP zo%3@@@E_A_5ARFHGIn7rd%VFNzREpbN87+9FJp9JN-`r|Z_||TpiyiF?f{Q)aENvM zBA&ns#3Q*B3;?a!$+!bh5KWqp3n-Aq zR^0^qZ9Ed%WjFyp0OAz3nXe8e2>r$6%)e5sLWm{Y9R@NY1iuuANnZ_00%;b=GKi5C zoMR7ex1oo&10i5Thql%gqf4uXEf3oSR_##xutjVAP_>Qd(_|OI4Xd~?;^W4Y@;SVB zsk8jqcz6?6wuL5S$TtkRW`#_5##?rgU66H(+)^nrOUo5?$|kz{J<1`W)KwTGMm%-N zCG?~FB3&eS$>Ex|$!Z{HHe^FSgJWjY3#Xv9l{3RgyCQ618JWtGjH{~3%Pmz?^&Aew zKeUiuiD*tkT^Ozkq#QgyCJhE(3E?tcdaR(T2RaSy$QQ>Q#Nh6 z3Ng(`njfBb8_Xhv2;e4!UV_jfgyxpmYPX{#QtkFq2rVZ3-zv$kVN3ZUH492GYRL`s z?JVi@?nahxs2t3nZye$p=NoaHZ;0G~AkD@-R$a> zH=ezFt4UAIHC>%^h^~M*$C<7o4H&<>p_Hl2Olax`g>$U0JfqA{59Y zl!{C!bnpPkTE}-3JOlrLGiMfZ>i4I%KoxEvPwB$T?V&>?Otfx-!x~o1C@B5R!%UZI z_aQ!1N*T`f*O7zC%hf%7p4=%WY4cZbWyBgF?6s+hM~cIr0{{D#;dzeF6}btCre#Q5 zQe9Kp%+LcF?m=yr%A(_myw=65A#I19K3g6>G-+%J6{YlP#gv%~;{GocBi5QL)-r5o zfbbvc$#)?q&J~m2rFouw88=<+0bW5IOEQMMPF-%HuOx8GA5u>xo8S|!=@fW#6rk8L zY-u*k$AVyGv@wt2lt9-$_wEh#cK82Afk5|tfUlhpYLB0-&((QD> z@@xyY%P*mKfUC^X&y@!&($bQ{7iI4Mstgqab7eH|K3bJ%NtFUvkfU4$#~H`gD1)M^ z3r@5<&zxr1F8u1wy}5@er4bDA8E<}o-ahR>br(ardv3{-#cC8RFJQ1)rCgZ=%xlii4P*8}0V~K@b-s@>>u!_%cY5w1L@nsWCf0j{$N<=Al8s-td}URZ{Q5nnWkW z@8qZKV>I+Yozy8hfXWc@%98voz6Mp4N&_QE==$?2!~a zBZT^WaT1~c0_iO}32~DFDd~QAydB80GkivRpNx_&DCp<6|TkHJGNIaL3hAb{S=xeU% zYP&7cZo4f%!n?eO249e2KRhmouGTw9x++cMOp=XFO=Z%CwR+KOtX}^&%~pcyUn$-|5k4W#bn>cK6KmobEn-&Ue0Z zX0M$)*I@Yh{q1|{5@Ua-hw5kJ=3TV>4m!cwOmJZ(yv@zkYFp-Nw;f#V#2tBUkF&@- zV!{#bF+1Yzn(##Jn6>>6nD9mYn2Gw4)vk+%SiopQShrZKdHDRb0~H11FqT1-b~6#D z{a{DMNvz^1e}KK&W{c~3B6m97usevlG(okQJ+}*?{3H{Yo5vlh??%H7LwKUQn+eRzE* z?MGR<){jznMP+hL^*RXo!Kwg^c0}A&+CGfOEk|2hxwg8wT*T^{Z=Ib% zK1Cx-kMHJ_Dr6J3u*sX;;fvho4fIVk`5eY)W+W53Xz!!t-$JL@5hw%n!9g9?{;RkH z84!_xI?jVS_z5qp!tPmyC~T0#^S5Dgkb!gf>yiWlf<8c+AB?tu1!6}NX(7e0SZ$aX zzl~c0xs1l31t6QEH4oIz7!g04oClZ7RS2jAx7|?21l-q>DDA63#T%^>PzETH4zvlE z+t5SXp%Ad4LtAUt=+drZ%iRuPR6Eo^Zqa%_QXM1uBH5*hj=3I;__Q&l4Dj5g!Rlw@ z=3TUW6h&6!f-oj_{`A3R_r4#+{S#R89GXrRoY?e}p(89*6Wo zOmmv*+{vPFPHE&RfzWEXn=_=)5oZB+9uG9--r`+r!Y3Bgvi@K~?#z@eTb_fM=0nYQ z&xQ?V5jX^J6F9FyXcjnAORRP}aT=>mXD)DN6aH_N^_}U);^jD7wwby(`7t{L1C;T)ca$ zMR%UU)GU-glmrtY9HiAIcL9QK8 z$S~KgsWG%ABV8lg#tCR;94f+p9wj~bD$Q!pgi(aFn9CNe^{5fFUm5=jUgoH*7~>v4 zkLII)hI{tm(&K*W67$iIN;2Xff=hk}Y+dkAfzzQUm+4x7R&JrvA&#C{$|4j}J(8~w zh5U~1FDX3zmi0`bcI*rLGi%}$4(2+jSZ>~Ag#*aB6R*!EHQ`MB{Vxh1H3`-kV!m8H zc(t%jyuvN~Q)^N$Y9}?&m0#3}!VG2AC;n}g1{~=Gb}M>s?l7A5EGyih zCI4CNAlzrzkl`8cgnNPEy=UInPJHa>O{*tgc&MKI?Jca1Xrpz^&X`k<-zE#7Vw=>8 z8mwTn7uX5TB!4;SXFG7AM_ z%Nm*qg)Q6wvey0`1<$}g;LKACIrV!JTcG$hkf(Iv7xsxmBuumpx7&cTYu zz>Sjq1m7o>3}@TR$id|0+NM5B?qqw~{7tm{8FV1*rHP40io>4*|J#<~d5*6X`E^LN zEJNDzS-O854-e}bR1)ov<<%}NN$Wf8@M3lN$fU73R8-Q3uBk3(#Qk3?MyxeetZmrN z0O3E>lOIA(oGB(ZXr3?EaM86M;1wo`d>v1c@C>AtY4=?k`G7{;e2g!*mQ&)*QGjCG zu%+2Bm)A&=*s|K2;3RLL`z2a#K1TtJW&R?+%nuhIwaW|_qAeSLgxt$JJ zoy}wgb3jS~1GFkjzf>KpOiObPpB1S*ty<)5IVzYg5D@D&dbozmY8tj`hQ$~kdCFUT3lAM8i$V|Hg=(k4l&Nk6kc2P{z1 zd#7w-pUIYWDu}2i{8Rk>Ks3q^0AwtO>c_$n3!f46nwy8sX%wyA&$Co>c~~#MK*6!o zut!qxjS%Yh*-3~3=rp(JB*aYyq@+7)1gco8Z@AJ!mi#tk&gvj};!90N zp)iUWSG8Q-01l!A-w$e_d0BI{eX6#}w8Ja|$M;IGuo@*jnb#A@MwqG-PoZ zl12-<*6E0>+v&(Qp50wZBoq X_bP{&QNXc~74ffWEuQzAi`V`Q^6i$& diff --git a/generator/tf/src/__pycache__/model.cpython-37.pyc b/generator/tf/src/__pycache__/model.cpython-37.pyc deleted file mode 100644 index a59f4e76c70b7669cfcecb07e2a2f1b0741b92f2..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 6724 zcmZ`;OLH98b?)1bnVts*4?-kJ(NbHHEsab8q>{)knIP*FdT0vaaiypW3>iXCK*XRXwNYL9gireF68n zUer(hROyTQ5_)I!lD-UTLtoKP<36je>QCW5r=QW!;y$kx{sOG_Y5g3=FX+G0&x3wJ ze@1^6_eK3V-PAAO`%h_g*J@rn#{R9YeSl8x7-?ZWKo*OowJA?X>#% zPJbma?Ulh$NB!l|!+cgpop8IKwl+q>lZUkv?#xp5^v~4r-sblC<1VOFUdc;nGHcU% ztd6X)O6|vZJ5oQiKD0U%ul>kRolG65JJtb(=^gngrg|v z#X%ay$bM+%DQ_mZ?d9E@&JfeD!`S^Lk(-dgNi zzA55MZsFl%Xw%I-8IhO6L`LOy+Q}n5Qttz`L#7 zj`j-Ci6sgrw+71siISY&?tisSasDG$rBSRXk|yei0CWUDI|P?WFdm`!x2S zS=t5F!`(Zw7p+Seccj>gFLf~)zu(c<@#nXXDX3cg!q!Z~GSD%3Pd z0GQGjBSYLW*GN*^hQ4>+pxqeUT}6igx@Ff@T`l2nPF+zK@oYYizNv2#nWhoovv^4% z`UCZ2^<8UhkDXM}E9r|F(C~`YwV@!J__Y6p_o4#-Anc$YCMje+01ozsaWL!zd%#5~ z=Y#lKu(vD&KY01Z{gwFsGCYzPQOL4vidEps!nlgTdm6&v)!;m5^<>oVrG*?4+IPa; zTXsX4o48NCmDDNToZ-m^Pr#QcS1`iNS+u(F&IwOnygLY!4Vq8EQMx-kQgr zpY>faB8=71PMc6c1flazk?CW+Oda^2l$Anv6J5IEst)ls^D?{Z&~Y>1)?wkubXZ5T zr*QBSM~34*7RTism{!ibu6;z@D`h3T*?c>)v+^Y?b2*!`9K)5TE=%0=S97&{;%Epe z8Y}?FK8O7lVj(NykWF9g!cZ24nXKR`zYMq4!$FI-9rn+9lex~3ly_blgy~?rKk*^+ zH>BBd$K2)!(SjZCnb-MNI8YJT3Y|zWD8XNMZfM4+q{t=es-d2N_q?ji+n|cyaI>^8 z+L-_+txx0*@PdTZ4*;7T`?3W~o51N9vW9%r61?k@wdGM8@GQt%EYIHmw;L}cl1b3? zhP#mgQuRWb*tMWN+>TS2ThKAXL1COS!t}ZuDS8burR8ch!0&E2h=MQ)(#|r3VS0P5 zbqKD76e)cXv|$`2FMbf~YeAgcZ(a*t23E$&%dy7u`Zhdx7$xxwX)rQF2GpP%n)R?7 z1#szaz5DHCS-!Cf2bQqwK`GrRciV9OW(lI!xRTt{=|iC641|)aP4h+E&zpTCOp!9B zfC|LyWQ)2zlvwEaR#Y5T(ooAd~y}iXWq66Id9p?Txx$Ge*6)v_afpk zD9;JSL0(lcuVs#mbU=TNbjj|J**@~eNY-`vku@$!vTprQeW;N6F=6)+u3SWLhQXi1 zi)iC;2BLuCJ7(wYBfWIR1@48}>UXk~kPh+6Ct}J_! za#bXdj)qAj*lW1%BsF24I!VgUxK^zIZ&n7$*SP^|*0T36!xZAv6!pxgC-zP2xXWpb z+No8O6KcWK@7-&?$021R!+LLJ#sQjyz(oyM)w{335i|x+vq!H?Y2F zq>G23_1JzBNuz{Rx=qo2flow5^HnrWPhc-sE%P^XSmBG7RE*(Y`OXLH1JIpstw+)YBsBg&k~`^Qd3gk&^9zK;acL9*$=XG@4FOamR% zMGkD(a!qlgfYN?`lSM`pEu^_*Oun6g%C%T#eFd_Z!$ z0P7t8FN&ocy+GZ~W_^-Q#w>91#v2Jdj6nGnV}?fD>`Z`s8`E=V5RQ;EF*lOjslAn# zx0!b0fdQ=@rp>vcXkbV*f6Gbk2%v4g$@llzP|aqI&0B1~!{&W9|9~bhkC3JTuu+9< z;o1=)iFA-4%OWB!iJ>lq@DMQcHQ+|->+Bepj;K&0@?-0jj}Tvm@&<2 z=t`8mMAtX*K%^2pUOrcrxG5_wIMD&^!^0b$#|1VbKY0Qd-Bq!SHMiaBw9P#%Iw`e^ zT!Tq`VH(UVCD=F1=cHD!O5SpX+HB*M^Ks1?^hXK06LA=@MP|X=7HHF&i%>RYZiJU9 zbN4=(IY2X$2jB#1E)T)O19Pa4SR17!5rBs=CE!UJHIsvi%N0F4YoQR?y zw{D2!p+=3J2!m@b$gjwrkD21YR;gl0gKm6YTGX4A6fyzn8`F5ny-rs&6j@<25uZLO zCGgg9O0o4lbSJE8ED_#S*v1uHW&jH6|LS^}wl`V{D!<%Mq6gbi+>Y8D&fI_(Ek~?| z@eNfAG6mM^Kpn)nOQw&weUp8P$7k9F{?v%BNiQ#I%Fd#a0&%gx=u?_+kp%s-*m823FwevmK+;Nwt>Un%z{303TJEO3 zew5d6aKc{xw!zGfqjt8-JNtBTw zQvy&A(&i#vh6F%CSm2QfoSLPkGf6}OOrE44#pd}P^NvvSQaap1@tM~|IsGW^rW=MT zTq^+zaqMxPff?j3sP|Eyw{RSAh>?@ip9V61^APX3Cx?x1wYh!%c)m7u+#O_7qExu1 zR)_YF)sNu5Jqs8FCJq$bv?RwK4$z6U7S!AYyPG9ZHPnnWz{k4-$F(-$uG9W9Orhwn zOwI>m?Zb;Zy;sv&((YLIUjhdWC}=LLA^UAy)sVdsoPWqPDx1qH zJPL;-=4I48M^PGG2V>^5d91Z?6HxaAl3(Ciwr8U4B+8xLQ8cNNAE8P%Z$L-5v9d&8 zdeZmtZE+35xNU-B32PQkcjcc4a*$u-xJ=sQbh8lDO}_H1 zL_)R{juO!%9Eo%VnFW*zRd7dVI6Vb+|0W!MyC27y9y{mfVkD%H!CCSxv+Oy?5VYe` zfWu`dqqF3kfr^Q`44|g9t2xdYV)%znCJjHsFn`2k*%3d&DJxvzLVPKfY{_f75c>wY zzzs&h4qJkH9uSnXHxf(-9KYK3zk+z($ zBPYtzMYft}Q`aihasG5-Ty%V>v~umLR37zFrGbl|P^soO1Zv$MeEljcsgPp$lYGqd&U#*VPamAj!Lg5xTjRmAo zlW{-C@Ve4Klb{QSRViJhzFOo3yz6Zw7dJ)++MnSUJj+|&yJ;rH%WPm`JQ1PTNJ8Gh ztABx&7nFc$%BwUBI@=sR_Y-?sU^MfyJDf8%A}uO#N(s};*pB&f@=sKupNhH}%F2vunNTLEq$ zqf7qX_nfUz1|;oM9~Juz%58@|6U!Er-g$5yjzV9<$uf~`**Rm&w&h1*-L8|G*DcxI zVb}68$=2aGfQHD%zQB4H;*!(%ijDadCD|CO<3Qe$8}jx=>&sx-!P~97d+g8bQ}zYy zK^ncbdpg*gtKnpH(pOLOp)F4OsF!{<>fi8dKQCo}?asIK>u}VNTt*KCnugELzksl; zW-uErI1;G+a2eErgwfXyAOygkqX3Q|3<$xe->8nBhc>E_a5AJTaa^>HWb^}2Dvfv3 z(6tWdxz1;(+%H@>&y97FRtuwiYPr^2E2Bz_RJ#j9^bV$V+PP%SGimIiStV7`q-f{i zXOI7sefsFLy+;6!>Up|NV-2WyQ1TcXbsz9j@d{Z)e%nL(APa;ScU0g`dIUt((lfLA;JdvGWp zYU&qoGC{D<7J}g5r7#4|9twbmd${`zgu(`OxIzwJAfq84BJUysW*OrIqGAL$kgXkL zmqfN+kd_t&Y0TGOaV_cy+=#x%)GRsxaPO_gMRWr!_`0W5N`>o|i&@s()VN;b&&=bz zCWpTGZv*`y!Frnb$mup1F0oZMDyrOiNO0ZRqOg;)noaV;MU&Eg@V<+Yx5M1JmVtxS zsX`bja>yU8p18PC=6WTNVjW4dPN%g&Y1pLGJCM?E&_oz=Gr9+rJozow)=+hq|C;v@ z-EF?{Q}HD!k3^7U1ZFfo$1vm}-_6DD;|43a3v)d-jZ%;M9(4)b?~{b@^gpGHkl6e3 zW^sJ7pFAq%=ZBL*r7cZ|>EB^;TQfL%hLB7w!RDkYjcXSby6cl%8@K68Bb$`?n&hPg z^B^Biv$?M3E-dqzLU%%6-gszGnk4>6>q@VC!e~+yu0y*rGpwk@j>!OYK}dHWtI||@ zk7`f<2DYTVbDSIvj3e4Nq6qNEd0Ybc?9BX(Gp^Q6ugV;-!#! zpQhME3JfvPcM!-;z9}}kz3`6M>Lx*m_ANw=;D*5_y3vN%3^Ylh8Fj3;A1QgPosXFQSx52u@f zJ7rv?MgA<-`z{{>!Amgl890UlOc9{Olzc#d5YWQAMGS9?I)KYn-dlFOyWBd~JLtz) z*fzYyC*T?^(HK+w0`5K|k~CzF8x$H8!7j#7J>ESy*xla=w_d;A-`n0g+B+DA+i!MW z?F^6hw)PL}M+5)*0=^vilY+8L_?w&0zweuh6hb*HpQgOXgI>a9$nqF6xH2H<8_EF;kr26jSFnGVaXL z%%W~%o!b0Fb{mM-(H#GbA?mDa9lAqBCYd7B&3d)u#hArX`0r;+vfs~Y&tDXd9uGQc zE@YI)4A_s0T(S$?(!{?;1uCB3G;h!M52pTj7l^Z8sw>@9p{Kv8qZF`T39qLO7L>nDp!WgDEG zRZPf@Q3QQJzu1ucGR>2XD2dV>D0A9Xi!RbPS4PHiWmBg6N?9dObJ@49sI@_hNh$dw zooDJ>^T!b0B2c-@!#{;d@cUf+6qI}4MV7hD$}tS zhG$ijUCnUabgfkUiHE{y6FJwl`hkTJC+W$NFxt_FpH<(LnHcu%IAVzg4q? z>yqX7VQuhHS%@%dOFV?8;8O5a80WfJ_{ZAmBW^%xk!d{2eQ|Y{Z?#my1LtCwXg=Ua F{spfAus{F+ diff --git a/generator/web/__pycache__/web_generator.cpython-37.pyc b/generator/web/__pycache__/web_generator.cpython-37.pyc deleted file mode 100644 index 6ebdc86271262aff467b9214bd7f959b3127a1cb..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 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& diff --git a/main.py b/main.py index 3ed4448..4aad724 100644 --- a/main.py +++ b/main.py @@ -6,19 +6,25 @@ import json from flask import Flask, render_template, request, abort from story.story_manager import * from generator.web.web_generator import * -from other.caching import * +from other.cacher import * app = Flask(__name__) app.secret_key = '#d\xe0\xd1\xfb\xee\xa4\xbb\xd0\xf0/e)\xb5g\xdd<`\xc7\xa5\xb0-\xb8d0S' GOOGLE_CRED_LOCATION = "./AI-Adventure-2bb65e3a4e2f.json" +# Initializes everything for a session +def story_init(session, seed): + pass + + +# Routes to index @app.route('/') def root(): seed = -1 data = {'seed': seed} return render_template('index.html', data=data) - +# Starts an adventure with a specific seed @app.route('/') def rootseed(seed): if seed == "": @@ -29,104 +35,23 @@ def rootseed(seed): session["seed"] = seed return render_template('index.html', data=data) - +# Starts an adventure @app.route('/index.html') def index(): data = {'seed': -1} return render_template('index.html', data=data) - +# Shows about. (Should also link to paper when published) @app.route('/about.html') def about(): return render_template('about.html') - - -# 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) - - -def story_init(session, seed): - session["seed"] = seed - prompt_num = 0 - session["prompt_num"] = prompt_num - session["generator"] = WebGenerator(GOOGLE_CRED_LOCATION) - - first_story = retrieve_from_cache(seed, prompt_num, [], "story") - - - if prompt is None: - - prompt = prompts[prompt_num] - response = generate_story_block(prompt, local=RUN_LOCAL) - cache_file(seed, prompt_num, [], response, "story") - - session["story_manager"] = ConstrainedStoryManager(session["generator"], prompt) - session["initialized"] = True - - -@app.route('/generate', methods=['POST']) +# Bread and butter of app, updates story and returns based on choice +@app.route('/choose', methods=['POST']) def story_request(): - print("****Generating Story****") - seed = request.form["seed"] - prompt_num = int(request.form["prompt_num"]) - gen_actions = request.form["actions"] + pass - if "initialized" not in session: - story_init(session, seed, prompt_num) - - print("Session Seed is ", session["seed"]) - - if int(seed) < 0 or int(seed) > 100: - abort(404) - - if gen_actions == "true": - - #prompt = request.form["prompt"] - choices = json.loads(request.form["choices"]) - #print("Getting response for seed ", seed, " prompt_num ", prompt_num, " and choices ", choices) - - action_results = retrieve_from_cache(seed, prompt_num, choices, "choices") - - if action_results is not None: - response = action_results - else: - last_action_result = request.form["last_action_result"] - prompt = continuing_prompts[prompt_num] + last_action_result - #print("\n\nAction prompt is \n ", prompt) - action_results = [generate_action_result(prompt, phrase, local=RUN_LOCAL) for phrase in phrases] - response = json.dumps(action_results) - cache_file(seed, prompt_num, choices, response, "choices") - else: - - result = retrieve_from_cache(seed, prompt_num, [], "story") - - if result is not None: - response = result - else: - prompt = prompts[prompt_num] - response = generate_story_block(prompt, local=RUN_LOCAL) - cache_file(seed, prompt_num, [], response, "story") - - - return response if __name__ == '__main__': app.run(host='0.0.0.0', port=8080) \ No newline at end of file diff --git a/old_main.py b/old_main.py new file mode 100644 index 0000000..4ced5f2 --- /dev/null +++ b/old_main.py @@ -0,0 +1,264 @@ +# Copyright 2018 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# [START gae_python37_render_template] +import datetime +from flask import g +import os +import googleapiclient.discovery +from utils import * +from google.cloud import storage +from google import cloud +import json +from flask import Flask, render_template, request, abort +from flask import Response +import requests +import pdb +import sys +import gpt2.src.encoder as encoder + +# 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__) + +# 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") + + +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): + block = generate(prompt) + block = cut_trailing_sentence(block) + block = story_replace(block) + + return block + + +def generate_action_result(prompt, phrase): + 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 == "": + seed = -1 + else: + seed = int(seed) + data = {'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') + + +def cache_file(seed, prompt_num, choices, response, tag): + blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag + for action in choices: + blob_file_name = blob_file_name + str(action) + blob = bucket.blob(blob_file_name) + + blob.upload_from_string(response) + + print("File ", blob_file_name, " cached") + + +def retrieve_from_cache(seed, prompt_num, choices, tag): + blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag + + for action in choices: + blob_file_name = blob_file_name + str(action) + + blob = bucket.blob(blob_file_name) + + if blob.exists(storage_client): + result = blob.download_as_string().decode("utf-8") + print(blob_file_name, " found in cache") + else: + result = None + print(blob_file_name, " not found in cache") + + return result + + +@app.route('/generate', methods=['POST']) +def story_request(): + print("****Generating Story****") + seed = request.form["seed"] + prompt_num = int(request.form["prompt_num"]) + gen_actions = request.form["actions"] + + if int(seed) < 0 or int(seed) > 100: + print("Invalid seed: " + seed) + abort(404) + + if gen_actions == "true": + + # prompt = request.form["prompt"] + choices = json.loads(request.form["choices"]) + print("Getting response for seed ", seed, " prompt_num ", prompt_num, " and choices ", choices) + + action_results = retrieve_from_cache(seed, prompt_num, choices, "choices") + + if action_results is not None: + response = action_results + else: + last_action_result = request.form["last_action_result"] + prompt = continuing_prompts[prompt_num] + last_action_result + print("\n\nAction prompt is \n ", prompt) + action_results = [generate_action_result(prompt, phrase) for phrase in phrases] + response = json.dumps(action_results) + cache_file(seed, prompt_num, choices, response, "choices") + else: + + print("Getting response for seed ", seed, " prompt_num ", prompt_num) + result = retrieve_from_cache(seed, prompt_num, [], "story") + + if result is not None: + response = result + else: + prompt = prompts[prompt_num] + response = generate_story_block(prompt) + cache_file(seed, prompt_num, [], response, "story") + + print("\nGenerated response is: \n", response) + print("") + + return response + + +def generate_cache(): + start_seed = int(sys.argv[1]) + end_seed = int(sys.argv[2]) + + # Generate story sections + prompt_num = 0 + action_queue = [] + prompt = prompts[prompt_num] + for seed in range(start_seed, end_seed): + result = retrieve_from_cache(seed, prompt_num, [], "story") + if result is not None: + response = result + else: + prompt = prompts[prompt_num] + # print("\n Story prompt is ", prompt) + response = generate_story_block(prompt) + # print("\n Story response is ", response) + cache_file(seed, prompt_num, [], response, "story") + + action_queue.append([seed, 0, [], response]) + + while (True): + + next_gen = action_queue.pop(0) + seed = next_gen[0] + prompt_num = next_gen[1] + choices = next_gen[2] + last_action_result = next_gen[3] + + action_results = retrieve_from_cache(seed, prompt_num, choices, "choices") + + if action_results is not None: + response = action_results + + else: + if len(choices) is 0: + prompt = prompts[prompt_num] + last_action_result + else: + prompt = continuing_prompts[prompt_num] + last_action_result + # print("\n\n Action prompt is \n ", prompt) + action_results = [generate_action_result(prompt, phrase) for phrase in phrases] + response = json.dumps(action_results) + + # print("\n\n Action + cache_file(seed, prompt_num, choices, response, "choices") + + un_jsoned = json.loads(response) + for j in range(4): + new_choices = choices[:] + new_choices.append(j) + action_queue.append([seed, 0, new_choices, un_jsoned[j][1]]) + + +if __name__ == '__main__': + if (len(sys.argv) > 1): + generate_cache() + else: + app.run(host='0.0.0.0', port=8080) + +# [START gae_python37_render_template] \ No newline at end of file diff --git a/other/cacher.py b/other/cacher.py index bac80fe..3c65b42 100644 --- a/other/cacher.py +++ b/other/cacher.py @@ -1,35 +1,44 @@ from google.cloud import storage - -# Model/Cache Info -storage_client = storage.Client() -bucket = storage_client.get_bucket("dungeon-cache") +import os -def cache_file(seed, prompt_num, choices, response, tag): +class cacher(): - blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag - for action in choices: - blob_file_name = blob_file_name + str(action) - blob = bucket.blob(blob_file_name) + def __init__(self, credentials_file): + # Model/Cache Info + os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = credentials_file + self.storage_client = storage.Client() + self.bucket = self.storage_client.get_bucket("dungeon-cache") + pass - blob.upload_from_string(response) + def cache_file(self, seed, prompt_num, choices, response, tag, print_result=False): - print("File ", blob_file_name, " cached") + blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag + for action in choices: + blob_file_name = blob_file_name + str(action) + blob = self.bucket.blob(blob_file_name) + + blob.upload_from_string(response) + + if print_result: print("File ", blob_file_name, " cached") + + def retrieve_from_cache(self, seed, prompt_num, choices, tag, print_result=False): + blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag + + for action in choices: + blob_file_name = blob_file_name + str(action) + + blob = self.bucket.blob(blob_file_name) + + if blob.exists(self.storage_client): + result = blob.download_as_string().decode("utf-8") + if print_result: print(blob_file_name, " found in cache") + else: + result = None + if print_result: print(blob_file_name, " not found in cache") + + return result -def retrieve_from_cache(seed, prompt_num, choices, tag): - blob_file_name = "prompt" + str(prompt_num) + "/seed" + str(seed) + "/" + tag - for action in choices: - blob_file_name = blob_file_name + str(action) - blob = bucket.blob(blob_file_name) - - if blob.exists(storage_client): - result = blob.download_as_string().decode("utf-8") - print(blob_file_name, " found in cache") - else: - result = None - print(blob_file_name, " not found in cache") - - return result diff --git a/story/__pycache__/story_manager.cpython-36.pyc b/story/__pycache__/story_manager.cpython-36.pyc deleted file mode 100644 index 36811e193c3099eab6ca5740b3d9de52b70b2829..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 4112 zcmbVPTXWk)6y9B3thi3vrfJjAmIi1kE+r{26dp1mrDYgyW}q+~GITVmS8`mUw*>leKoul_>XPe*sF@EsZEMtGO ziN`{J10}tMN;1h~_C5QINg=JLOj`S5lVfB{2O|d~fsrd47&S1mq<4?CXP#j#Yx9C% z?^3UN1Uefi=?zpN+vJjM3dyB-%CORsHd^c~U9^sDNDr+mXRu?laPDP&^`uK%&_gYi zJ1FTnR3koOKI<{mT-QP}EaaWGRh&p^t)FIr%8EwN&7yvi77Z1q!#GRZqOj92-cmNU z`I9LR&GQdexBGkHY7%vKe+lBOpWM8;+KaN)+n?VZCcUtqtm-{Y?e7Ii&pf9Lj-$yx^Rk*=m z5GGR55p7E~%sYLz&GNzz<9>Hn-^WxlG^Z6v>1Tat13Ec{H!0nl@fOeLjfj&Ix}5&R zLv1bQG6oq52O%Zjwbgm*OM3|_mlg`C-svb>^o**4J^Kd_RBgp!Vr1r8=A+`x#z7XQ znQx4JQEmM`w;jpdvIkeMtsM{ z8)qbznAWrzTh;i{`TKDG9hdBY%ZQFTi5=Sw?Iney0(gNsjaT|I>IC()UF@BHl#~Y9 zY-%n0ag=5SkB;88;lm2b_g}%{w26vwk1z03d_m;1=BDM0PP;*#l>~dC?-xzq-|Nd^ zO#PPcKN<$H`A5+v3Xo=Mo+<*y#G&Nbikv`Is!HjVC%mR-d#--l#f$nOgKwa`3}Baj zNV7rwHjXMH36)iO>!WeMBS#B8O~G!MxEj+$9SiXCGiv9eJ<@ z6a$FK`gwc0_(#kJfU8B?_dHcgRB0@ePZ?S0go#rRjwbea1LXkAu4gHF+eN7#(&z?C zx`yh-gRg8N2sh*j*+T2dS@_*tv3TbNkX=*pS2)lhn_@l1`5vg(2YP{Xp-oWZ8z#|s zuW{~_ICmx7%5RjBC2#)NA1()37Ib%WkNQ~{$9aQ#!!*oY>h$_;Pg}H@Dd3(tmatWFI}T1*1oLrE!+Qd;z&HAZS7 z*%JF*JhF*Ibdc7Lh3G6+ZWa|^9Lp2Wn`;!zy5Ph0Z)jsW#bx=%w0rf!%s-;kbY%^R zNf4*&=v14vg=%A~VxATpQU39)v2fhhUy@9k!k}#=3^)FdTs4bOU?N)Ay3f#g(n7kU zdH6m{<~fYRS%%a{=VeHEMnZ~R7+YhVc_6p2;M~OK zY^R~^iJ*qQx6ihS-(g%d zU%(uTFRsWdCeNw*`Z#irusV+M{Q8R@SLbttJ7OGd)eUsO;!qEL%!~wuUSw`nMrR`k zBZJd$a7;Egb~2}Ae{1Au2;lS0oe0WA*~LOSi1C576)AisR0*Aq(lyVQX&u&@&mxn* zsrRtlxcA%CJI=@7#UI-!Wgdq(ZZbKMuS{2n&N`2Tj81#z*nmz5RD>iG)io;V&{(Arb diff --git a/story/__pycache__/story_manager.cpython-37.pyc b/story/__pycache__/story_manager.cpython-37.pyc deleted file mode 100644 index 098f93fb1943cd6d4038be1c494d8286e8b5f08a..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 3590 zcmb_f&2Q936rUM?th4!mLLj6-Nhy>Xp(KZv_JGn*kg8Vgs#4k_q9V)K&a%X@m$6p` zYxfkSa_yn_je6?8xJz zP@Gs{ng!G(vH$(L+zmJFU$B=3jk^Ej+&$R+&?Xh&9-<{rU`My145YNA4+D^doVtvJ2QjS@^KsfqDqZY0 z6~JEH$?}qSrtjMDpaAK4@8aW3c_VlDGC#?eMX{)FTCHfg4Kk|4KZrc9YUQo*(NONt?(;D&C>u3>9R`V+HA75>%w|j3yo7G#%5i)o6f$Q6gP$BmWm)UHu^m z;9@zg9|B~9Yy$=ki4N608=(>;o3LBWY**6c#+pXFo z>~gLhC6V-D(xp>1qw-A9%RLez?j&t5i;_G_f~Z{7ijh$__Jc^>ytLGqYiYYArIZVL zSy!e9-MqB4FA;FlfD-Umk+TI9jN4)Xs>DwemuKQ|Oap{7rCKct5KRL^fO5oPukfuQ ztQIsNKnbFzwG;JCE02-yVIO4zH6JW$uJeLNX$LpAwP#ITpfs>*1D{rG)r7wl7iWu( z)g!-~UC_;F2PkT%qC|70Q>dP)o@|Cybq{&R0jTV&VhjV_&u%L(*t~JW3;9{uo1V zA9uqh;%_6IhrKP7D|g<)qscv#0e^?azw_nV|A#h|#k>YR3+PGE2332G@^q_JH1DUq zH9yb&V83vvokww8G^o|fqQa&|J9X#Av^i3w@^od8<*9a*dzEiguYhu;n`T+(A((Y) zVyWtPA4xxpvT>4-mvB#Q6tkG&HdbU5EGWWiUTAWgH+iu(1JT#Hf%YCjyNM1#+vlt7 zCovEMqv8k9qv8iPG!%raPcv+<;tNxfPaJcEyoRo8YR0=vQ_8~0oW^WTnb{SLeS@qm zTr5#Q=EeD0^8cR%p&Vd`&aPa0^INr8xo64gePGkQ||!uIfF!IRT@0dVfJLkjUm zYs6@D3_Yn%BWGkg^q73GJm4?bkRfy8k|Xh}cw+8opzyLeJ96m>+fl+}TC8uDwl<#$ z8YFOEHE+=qAwZPbx<_Yf0gA0z_k7)S*KgsS=f~L=8ueB<$Hi{r$JK46Qd0hp^>K_B z*WdoQQ4A-TEM(>J+h`y{>-ddOMRbd3T}N>f#EYm6gEa=0j^L=RDzY1Wt2R@6%5}@- ztQ+G!z1NX=8rDHvF>_gaW);d^yLtudm@RKWBxTfP9S5uU_a+;^j>Y$pt0)c@*Ku4F z>$9OzZ;mO#Xt~Zm&^ci*>AW#NJLdEcG4VMGLO}rUvw5fC7+`f8Z@FgmXJ);Zcj8P_ c8u7iFe$#%Xb4x!Mx5yP#h1l?HDSYBVox~m6;T%r_^N1%efXNVCtC2fI5=VL z!};%yj3!ACM&l?|`E$o{@#HF#ag@7Rn#`W-5gI+3WmES&PLd;7HPrp1>*yeJRkjmy z%Ykymmq~Ixb23dsH4RjprR7R#=H?Teq2mTZxIx*uo-dC8Qb74FB{`cHDl z2|4(n$HZ5~_XWPZ149H6IYi8$y5JY@7W@@EZ7qy~{m2SK)Gmz$`|06ZRxnUu_dKgV z4M(2!THAaWE6> zX{=@~0*16;jJLSMTZXKly~R&F&`ugZB?+M<4hW@gn>%f=vK_CMOz1hb*rRd93wFur zKQQC;=EfU$Zjt%Xw<0z+ZBWwu_kzvs5ek~bo7=;wRMcAW z?_~|=Tk%tHN>biIZG)12gZm@Q%PB{hv6WHFDVrGE(oFKxr-c+~lq?|Gka9>^11+O4 z3O-`8gCaFYmiE>owlLH83fy-Xc>!OJGzx|yul#B}Mx5G(N$r9areM_S8(^f`Tk_6| zzkh`d*%d*$x(|aIw)$}@+Lq4JWuqUHB4@{3*w68l_?MDRwDYLCHD>6|N26Fyuo-wyaQIY;)rWbpC;rXKPlT0i!`IlsFeG z{b0?`H;8C{wh>~fQua=;wq`_9&B=RYNH%#u<`4#IzeCO}OP7d}56I97cZ#@4MSM0@ zzJ~lIf;W&J0_hzPO{mGMMz-F@6F5`&&>CI^WGQDro-_3$u3Tn&DbB* z**pdsPx0kn!wBXv!G-ZN^SCgD1#gIosKT3~CTw_1>)b!hNW% z{~&jqkb@6;Ong;*U*gLUI--ltmv3U*KkOcJ_I$B7hqlm)U-f*0b_pk?z+_2St5Yk24>~@Tq zIjzLF+=Klec4%tlHXA(H*?`jd+6G&JzC+}{2&DgYM13n_W77sD&3`Z0+#VsJNxZo| zoJvKl75_ojaK05k1*atC4b(R1(r<8of_XXJQD$sq)N;xu#-4vjzNiQ!#50h7m+P1R3hlp*ZW5_{(+UpYgQfu zqd_c`I2SDaXwA;Ih-iMc5n`!Q_D-<2W<*lX$y;PdHhDni5C&?$N6su$mxz*g$j}LQ zlDJ7pd^T0ShWr(R*N`3p=^-Jd1iyxKtN%BM+=!HPCp|SrtuISa!nOV{bVOq;BN|hs z_ji~eaz$#DOF&*dK{;QG?6q`fN8@)<-U)$F$skVRbl`7vJt4GZ(IQ=e6lhVB?_<2z z-domuUu2>0>msPjldd+BDuPZ3t-q#5|Fg5JW|=5j+MkUxF-@Y+DT6e(mTA}QhHc?( I*!%YWUpYj$jQ{`u diff --git a/story/story_manager.py b/story/story_manager.py index 67b079c..49f47d8 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -1,5 +1,9 @@ from story.utils import * +from other.cacher import * +import json +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"] class Story(): @@ -32,27 +36,36 @@ class Story(): return "".join(story_list) -class UnconstrainedStoryManager(): +class StoryManager(): def __init__(self, generator, story_prompt): self.generator = generator + self.story_prompt = story_prompt + self.action_phrases = ["You attack", "You tell", "You use", "You go"] - block = self.generator.generate(story_prompt) + def init_story(self): + block = self.generator.generate(self.story_prompt) block = cut_trailing_sentence(block) block = story_replace(block) - story_start = story_prompt + block - + story_start = self.story_prompt + block self.story = Story(story_start) - - def act(self, action_choice): - - result = self.generate_result(action_choice) - self.story.add_to_story(action_choice, result) - return result + return story_start def story_context(self): return self.story.latest_result() + +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) + return result + def generate_result(self, action): block = self.generator.generate(self.story_context() + action) block = cut_trailing_sentence(block) @@ -60,16 +73,12 @@ class UnconstrainedStoryManager(): return block -class ConstrainedStoryManager(): +class ConstrainedStoryManager(StoryManager): 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) + super().__init__(generator, story_prompt) + + self.init_story() self.possible_action_results = None def get_possible_actions(self): @@ -95,9 +104,6 @@ class ConstrainedStoryManager(): self.possible_action_results = self.get_action_results() return result, self.get_possible_actions() - 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] @@ -112,17 +118,26 @@ class ConstrainedStoryManager(): return action, result -class CachedStoryManager(): +class CachedStoryManager(ConstrainedStoryManager): - def __init__(self, generator, prompt_num, seed): - # 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.possible_action_results = None + def __init__(self, generator, prompt_num, seed, credentials_file): + self.cacher = cacher(credentials_file) + prompt = prompts[prompt_num] + super().__init__(generator, prompt) + self.seed = seed + self.prompt_num = prompt_num + self.choices = [] + + result = self.cacher.retrieve_from_cache(seed, prompt_num, [], "story") + if result is not None: + story_start = result + self.story = Story(story_start) + else: + + story_start = self.init_story() + self.cacher.cache_file(seed, prompt_num, [], story_start, "story") + + self.possible_action_results = None def get_possible_actions(self): if self.possible_action_results is None: @@ -142,23 +157,23 @@ class CachedStoryManager(): 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() - 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) + response = self.cacher.retrieve_from_cache(self.seed, self.prompt_num, self.choices, "choices") + + if response is not None: + action_results = json.loads(response) + else: + action_results = super().get_action_results() + response = json.dumps(action_results) + self.cacher.cache_file(self.seed, self.prompt_num, self.choices, response, "choices") + + return action_results - action, result = split_first_sentence(action_result) - result = story_replace(action_result) - action = action_replace(action) - return action, result \ No newline at end of file