From 1624e1a349cfb50f708470c73e51bae58ca76f3c Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 23:49:06 +0100 Subject: [PATCH] Add SHAIL --- checkpoints/sgail-options-setobs2.pt | Bin 0 -> 15019 bytes evaluate_models.sh | 7 +- src/core/gail.py | 8 +- src/core/ppo.py | 4 +- src/core/trpo.py | 8 +- src/eval_main.py | 13 ++- src/evaluation/evaluation.py | 3 +- src/safe_options/options.py | 115 ++++++++------------------- src/safe_options/policy.py | 11 ++- src/safe_options/policy_gradient.py | 4 +- src/util/wrappers.py | 4 + 11 files changed, 79 insertions(+), 98 deletions(-) create mode 100644 checkpoints/sgail-options-setobs2.pt diff --git a/checkpoints/sgail-options-setobs2.pt b/checkpoints/sgail-options-setobs2.pt new file mode 100644 index 0000000000000000000000000000000000000000..5bc3e99e565dce52f269dde8af132de183d67fc5 GIT binary patch literal 15019 zcma*O2{=~Y*FS8|JS9X_GDe|+xX)g>2XYaK>dw)FKrDY@}$x%u3FeO57jnt_W#qGbMMqBwDBWMHhDWZc&6Vyu z2QQC{>NoAUN^F$;Om~UCKPgcP?vg1{is_60c?_l4C}qzmm6(Jm)wv3N>i)|@VgvuF zRg~I)$kol~_T@@MX(U8x{>Qyngha>m)hMdJyJXBC#HwfMhZ zNIDpax5a+|_ZR~R(_{?p<+I5@gSqxI@@%cE1`g{W75uLLL622suYW03q?leY^^r}>JL;l1Th z7_KpdyvOYmB$h6P8QSw`Dm1ss~%*nmeD*tdtx zyj!mum(Opa;f{giN1uwgum3 zqtUP{W)56gvgk(SK0bBHHr}Zxm^}2e&|pm)Wf;m( zYk@6%klRoFW}k4 z=7UJrFc$+n4lqZ>pZNZoD#`61O8s8+koxcj$a5IZc8pFXwUIGotF;8)UlgF?v6d6wY0y21+~n316P-fh&$lZ1F)Idb)8s>Dey9TQcbcrbV!2 zmMa#lzf8w&AB9QtO)%iiP0&0#gZ?(k(z9nq;Lx+2%5zH50yMySq+6`6Qi0Z=M@pt_~*{Iq3+1Bo%?92zDa8h(0$Of8|RCX-vjzn@Q^T2}R z>O_|+;IFYW@JHve=Jq9+n5jgE-e(Bi?8O+qSO|&6580Cg`yey1mz+v9pwwg^6g*ae zF&>@Vus@U8;2WFZ>T^}nS*;+JlDf}dFv+3is?rd)AOma@=YakBdXoC&iz%Cn>C(%a zu=}wU+?Sdr{&hxOm|A3jEB5y=odE_UGT({AhSriOVImEFoJ)qgI=I4s6Rav<8rcX} zCOm$QEm`M}UMKE?-lBNe@zV?TnmxpbhGckHxgQLCPr+(^Nn!KC*`U8!Ae7qbPf@ey z!ulu9p!lQ&@Q@JVZ(Bi+-Yl4At4Y64yHa^vCt6<{32t)96m?_~B~+iLI@x0+JXwd{ zt#?F63Z8R!INXhJEF(5^I{^WM?P&KC2dE!nNe@2Nvb^&X#dzs;@8e#gx+7)mNvpR=Y4eK;TLOqr`&KyrE+9khQ+^R0|nhdQCx zyK&U>#|T{oZt%iVm$XZpC?Ge3QkPGGZ>dV8@9NE43MbR<;WfC6iXd#r2ORStgALYu z&Z-vn@M+8EaYHNu=-HZ~@KZ2EEGIRC)i>$;A`ZsV<^I1yR`j(qIV-WB8u^skY8-ZW-4`Hk66=6?d8a=!9nA1HWgnF9} z2y%C!9eX6j7aJ9zSVaXU>fgbDb~E^}OaWwD_rsoPuRy-87(I3OFkpQ_SxiYEENp z&e3i-dvhwj_dP&smCuM)_g@9wYjmis=OOD(R)^OHnbdfr9gI9T@@<0an=KEQfr2v! z#7B+h*QxBZg66~nG_PHaeyto08H3$0_HqS*_yCU5-b!~}c0;+}cpw@05C?gzhB`1ihzLc%?=e zw!|L6#4ZmQD3TFMWTn?hg{6_tyI;66a6LryixNItng}Y1lPOr+o5ZKnp#8NZv)6V9 zar;Q&mqH1)_4E+T?3u!650QldWzv)f8_DpF2l;)z&kV;dXKyD)z?03{IK^uMiLaev zp#__1eR49aj8x%`We(%nln1!=+IJi{pp#wAPyorKV3@s00k`RFNX8Z6jbDLK}Z|=D=g|`v#j4S8!UuTo{@(RkY+stRU9)r_CJ*0R17#}BTC2rzQ zQstx`eDWup?taN4KW_>gT=@PeaF5_EdG7dw0ZH%WPQk@mJj za4p(b)blq@yec4{DR1Z}tQ&h8%#2LLf!0qczL}@5s;g=K8&@#f7f9t6h2q|6d&tXw zG)kvk#9;9iSb0H9c>CdWN_y!oyrk_0)6PmkXh=1MYkr4<=RM@{-d^0}*a?}kt#EaY zli1ovMeJ?*h~E!Nw11S8aP*8E=sj!z->01+^2wwo6M4Cf&3L<%7x`L z(Q+OoI*xEK|2mA;XDNcITO9oSB`LOV+CVIC2o}HH^b4OX^nslPCNRc8SNy~M7Ns_8;Dy3`EW3UUhCn%7Z#alY zCaDlEmZn#844Budt0=xNPZ77H$bOVQ&6zTjO3s{OrbA})+NKq(Dclo;D`m`Ev$tob?RGTh7$&hXONnU^GW!FAxh@+ys{FGLcdx&~PEG=@M|_Em z7$*eMeab!C`-FK{Y-jf_>42`+O_mzb%hu`bWTnj>Fml;P>~VA!h$qYk!D(swCtf%^ z-Gn|J(dGviRWY4)mvF83A2uP_56;UqqLa;8&gAU_Jb2g=zKN&9<)CQb#!Q3XYnDUx z&^YF$dlU+{*fAGQhZz=~#F_hs)0^Lu*sr<;bidY+dXGqhmboK)v~esnEX{H`Mx-Y*PT+rY3S9Z<1a?h+?D;rNu!u0FSs(9^M#N8;x#v4|MW5mHJU)>BuhsN7 zZ7yt2@J8q58?fn!B*a(*uwb*vbo}0Q_*!H{>ow}=al%$C`Er&knqS@LJIat|hF+d zgjHi}V9ZZ>;lxT0PB88wG~Ubs`)Y=1ny&Euq7rzjkAZ=49kl zkc=luuWdff7{3Q5jyud>)hlCJ3oTgFrx$oEQ5l9!tYSSOO=Q`fOB0)FLGs=dxbXQ2 z&9`gf%y-nWc?n^ZwtfPg3VsI81Cv4FO%6Va@@BbL-qD?=LTXnp<{g(d1wGLt2ImMae)Zm}AET%CR=|z9$!kzdFNT z{BZhn+ma=|0J@t-fcTfDqQLX*BMcP-` z(bWF%OHl(X#~Z`;+Qa3V(4@zUA-z%w{5mh;essqL)w={2W~tu3;DpzTyl!I`*=h*5mxe%^A3(P97FHXu#58QbZ@D}Pqz9nAK>{Db!XPB1_w z1KgvEnT-Bx*4BK7PtHh&q*pK4*;iA^%XtIeaL!buzpW76v$Wtzi3DtjmBPh(KKRaI zA!i-&NKn4jhq7@b6@E{H-zCyCYr_&6aKadd?Nwuyxjy6+)#ms-AcL(aZ(;r_m)XOM z4&+|AiTydg0hK$8*qI<1{$cGZ7Sx*!as~dVXgmS`wDIf^ufy)OXpzF0{VcUag$gG) zP?hKvu3W#B|8hziLu&+Zt-um_r+lVbHVqucuEyQPop|{C9@yRN4Wm2cV7g6TzjTqJ zt<8t{&a-XYmx_(-*#w@2R!G6Yqmk6@;ZS#T`x3BGFJ*HRzGKBlKT>xx1LGN|z1wL4Tf|3+EWu`B1r5}8rnT24gx-42Y|=Sole;fr z@WSh~XC6<7EfrBqNgiU1UbDsD6=3GadW;(UfKPnzif`W_4{eExkpAmF-*NFG?{!fL z{Y*>I=4>e{g;uk%E7Mu>r2({GT^}A6_+d>#-?)$=A<$Vf3I7cF%TF=%?t4E@4&K7s z7&%Id&KS4gErk|VmsrJhF8G8EYxjxXEUm#y8wEH_ay!N)x>4-9h5X8gnM}cJ87f8~ zTi9Uu4Jy&IpljKZefa@b}*5&fz?;gO9e9<)vc z*%8{X*Ue6_%KRBVab>9gvWUV8ytuT)WZ0-~OZ%qgi#9#tpzK8e`P7Xjlls+^H+2W> z$F(3#)1|+EYq`;r2SSqm6~=!kVLIjq!0f~u2>-T+-yoVzleG^+z^h2`2Ya^rTQ2(b zjb{;uilF}T1rq?k;axyF(dXYYh_m2b(@{5A`lFCkocFBx{6 zUqRKO-*~IC3{oL;@pOSX|NHy`{+AlhcREMHpvB7(o5FF`$q)F@XBo4!X~l*a3-N(% z8Nbq8g1_Hs&fV~Fgj-*R(hQ2HVVl>|HNO!sEmj&FjNbFbJ2bIT5X$BD&0QG@nrIO6 zh`Zpn9FBjn!J@n&WMsRDH7!$SmquL?^xTpozcu-ydzq&ob+$Ts{;{I0L!sbkG6BSq z0v7*DNE2^w=LcBlvxC=HqPO58|GskzT>4bT3dav+9;dT#{=*{FX|^Kgx&HKLXa?{5 zU4wVrJd)FkRYQ%zSNOQIGwCchfc4Mb#RsRmvCeFFG`*0-9BDM`##B@+@uc$bc*?su z7!2q3XO4&W@}5CzAh%o@zS)_;_d#W7F;S2Gt=FYE%grd7H58Oz@5G2k?Z8O*3I z<~G@F!ir}GAlH-1yyB;`?Z>p4t-U4`=B{J~m*S}B)?c;&ji}9F8ow+dgUj0@4Kahq z(TS6;BgCL@83EcS<&l>$RY4)E;wxb}5 zeXm2oN`DK{wYSru=2|jM?9LAB8pw&=}M-vy&x&_gE zxauiLs%N|fe}#f22h;Y~gJERb20EM3&OP^15RM)Elev7UX4hU7v-ea_X6>Cc?^YV> zwJ(OP*A&5eeI{5e-_0!NKE;4kA*p}A$(^Zpqi@^0*?7CnT>Uj)y0D}OV?TUBs!7D$ zIz`-hG@neSW<%GiYiyePUA*MtK+omOVXo;^wrHd&?Tgw42;EfAH-IZGTUghlC*uy&fFIQd%##!o2c$7B}agP=7q<6AZcn|(tz zwgX=2=YSt=riVkvQe*2{*gr*qmPJQW$68Z3UG$CF*GP+t%MVh$**(lTqeL1{&w%6n z*N{11U1b0D63X??ql3$SVOT=}J{`G_t@GEWi%ZRD-)>vFGs2&(Hoqb$9;C>JPv)_= zPKr9UKB8sfb2i|a4Rqae!!wNw$#r9a;O>fQ_AS#2LihX7fjb%OQe+d`dtfYOONWDJ z!FYHvppceMcg51+N1TWCNH#Ea8H>~3&WU#ga)EOUKwJ|Jrf)f#9&1FxCn+E>F{F8y z%h38+5EPBAV#CL5r~Ug=MWfZ`u#tad(P4OZ{ar!fn=id3n;^;A|j zp#}|yk0jq&4{+4-b#%-0Ebp!~9OijXXUnpF@w1K1`Y`;yIqJ)YKA$j`iQUIznZYT# zF~*ZEJ-rMHZEILsK?9aaO=IbzHC#fv7M+}K2j+j>SiH_7GI%0Ohv!^p*VFraMV&6@ zS?D7wzbi*)-lRfiO*GpuZ#nHcJQl5=hC%1mMQDFEkXDfi2$SR3h0(jI<@$8G)nNo* zM&vO0b(L_j&JXrvRpNoLvAAb~1B*TSnFW^Q;la1P{M^2=PS!A!^E$8^w{n9~<@Gx> zJ9UrwM!Rbum_Mk8z2CZlKcH5J#VQk- zc7!Dr>?*=^haz^eLX~bE+KP2{X@XFZFZ(^%24`5TWrg!YIos2+tm>LNX6EW}6V-!Q zv4t%W=VZEnL{%ey$~a$j7be@f8bLLT;S{lzx(MxftMgYBDI$Vy*-#P`Q! z$XrvKoqBi;Dc~V%7-Y$QD_QXir1#T-m^3)Ao5JQ)#6e!Rh#QdG#g;_LP-^8gUZQ|e zoYYcwD|Z4U3i2^*rytjFq6-FB&m{BAWmKAYl5S4PfpZT+VX;C_-Dk;KG?|eJZhscT z`@{@#{VER@*K5J#PhX$Uy+hlLFT?Qt4v;>@7D5I#p~EF5`ZQEcsNoldsUB6N-sy_k z8Ct@F7uDg`<^AM5CWQLwNkLin0Xn?y8H;fkBQyvOfW|LbEU?RkNep{L_dnT-ErJT! zi@bMi%dt{cv(Ozrw5h`0Db=uU-+JgjJs1v3I^tPHpcl^P+4S-6u;79ruCf}y&Pl|= zysU{Z%pK_3fgLd8;d!W49Y-#_9vMCVhPxK2z?s*{;Hu*Vg|GIpim$eqs4NTdQhLIM zUCTjc>U(~elM-xuu0Z9NE;5y-eHf!~m?9N)aO{f^woiExv^ZG_9_r15rpLLUz4tn; z`5jDe)ewH!^wQx>H90tcKQN2U?UB=cl2gg5Nv-bzk^Decm1^N_o&u>UmZcq)%| z9eGT%;u%wZ;ltY{603Vk_ zG>e9eQfIRw4PpJG9jsVS0xxE5z_=ODu&rAST0acIRY&qT$#1vd3m%1vC`)ve6;Yc;c~?whL`Z!)>u(Rgk&VD+4-e9WOLEFvYyjL}DF*p2p` z@ocrLE3FJKM$=i6q7%`^eQV0)`1pP=I(w;7^X4>1)uM?|Ilz}>5?-;G;Uc`GB|+yE z)7jWo#>H6G^EzN>*;qMb>Gc!_2aVK>WE=T(g5R+B~dhO~W$SY8yp?XQUntj_*dlwU)5WwiK^S zxr}2T&%nhq?4V!fCoZsM5Y1ICN9%qHp!@a=SM5{6_4_*oss^ajW=nbc+N6VV<=`u09f1Rk@Vl9*y_r;(;Dye5MwMfWj zT7j{9%s|!E3L3*=xqw$oFnZJgs5-iyMS62|{_Ay={?&%2Qj+v=n+_(t0QNjMnVlCN z;zK8Xy>!L4_jP@US3m3>+wEJOmSwW=5-wR#Ge)KDB~Ab zKjPD_-o(&n@A-Av`|#?tQqKKIKU#BChQdXTP~YvsZ8HuK=qKAT$rJtHxRePTd@EwU z2D?o0z@^@?rPDah3Ft{FQhAX$_vBUXhu-{&trr4Zf(;pn?^bZ(AO|&+96naB+ zZ}CUG=H7_gWiR17gEpLZppz5kda=r7%Q4iZQB=3)G(X0}7@nIOLDsBV{1|wL%k<{p ztGNXnvz(24r8QA;=M-wUk))TyLLg?U9*pk-mOE)CCO3G{J2@fScI31mXGjyC?LCL3 zDejQkeuqEV)Qq~>CRnnuhApU^%L+fQWpiL1Y+Ny%Z5DiB8TPVd^W2$cKRLqIy_BJ; zL*wCvt~|^Cd6cCzD)9$8C*br+zxd-0k5Tb*JH8&8iT8BsnT1z1-_ZGqzwl-^+?rs_ z^)u-awHhYVQ{i6Tq`Cn5O)~<`I7Sl-0`Y2VDh!KCX2H)#aOyo<*pW0pd>!O0TH&}7 z*2K*QO`oHjVed%xYgag_2UWuG*ea3f*+~=|{F&dP{Djrt|IOUjoT0?$D{;rmM{G^k zYs{VTj(c=P1Q&vLvN5XG{8J?r`e?vWKieIY+M+`nLOAMAXR&6Y7b?CPOIxIVU{$g? ztzY25EQT(kx`U}~&i)VFx^pLa&52qtaO7!RvTPpFH$7VLg+tL!AGE$y$Ndf3$=kIa zXRqH42HkhMtaZvHbQty>4{FtO`Dc^alf%O~t7igCdNzm}cF96SANG6PcNL2}_mQ2R zoX_e`xxwczX83!>ID9{PJ{=v~&JE=|Fh@y-0@w6|^Gh=Y=N?O8;&^qMp}3Um4*tuA zG@oSy)mLI##9n5&(}Pl%4#2xfWsc1$70jw25`Tr>!C}{RNV1`CJ!SntoMv3j7AmWO z+YD_Sx?wE!Sej6NQvlnnWrBMmPP5-7aX9+yOU{0lKluN!g}u(4;M4mye3dCpUQh4y z7VZLp^86X>>wPE6Un5IRp$mAq8xehRSUc?_dj^KV(9`4~(v+KX@f7SfISBUl<%fR@!&?As(=-p6t{1bP=z$c|+}!hT>la}|Dle1h#XX56p1L9`%A1IT_ay00~(OPm^$@0Yg^EuY%t zHd9A$5loQ2%&s^okeXP9cBCF*z6~eo!lWkB{Wg$!o*V@k_>N9CU&kmtV+!hvlW(f~ zg5=;<>WNKZ)~Rt&c~y=Ea>wbMVI>OHjbX|X4VZ9Bk90pgf%@6gvASDZT)bb6innF4 zGY^88^wj>~`O!jXJP_-eJU6kj8`UE;XznKyGTmXxx`pcmn|CVH zy;2)kv3@G%2Tb7AhbYmoo6hL;d^1#^aAt#I#Pmyh7|_t=pg77@FnsV(GF9;=r{`tZ zWu;D?cWglRiV;k%@1%Pxl_~O5I8~JIn5=ulS0s|@fx_M9b>vypD5yh2^Gg%;K~DGbm7Ms@Qiy8Z~IS$+JafwusDsG z{>Z1gSK;Jat1r+HBnsm0n&9&n8^9nv1dTe~X=#Ka$c_Jtty(q2S28f)xKi*(F$g&b zS=xxzENb)(8o2o?y8c)XFQuI+{K`ohx2i9$J?#%$UJn&yXG_tV*hw(dd?9-5Q5J`M zSdA?$RxBpwHXFAq7A!yaKt2XhchX&uuBc+$cJ=vVFZ24o?OQl_xi5Zq%I1}q4upL7 z3H&_Ep_Jn)q*xOa==ko02DhWURi!Lzah^wJOH&0(XIAz3iqW9?yQc3OmLW{r3clvj z@Z)9!Yl$)xz;&SJVkbImauOs3^I%_h2(vsmk&YzHqz~5h5bNyT_vFSa7VWfz*K%frdqGn6@aFI_C& z%`a3282;S}e(e|oj|yDzsgVT5xn5_vA7!A{-Hbkr*v~(a{J|9mr=#U0J&NqZnx593 zXEs@3yvJ*IIu>(*6PcwjOZPnL3G9oxW1p~`f|I<`l*g!Gl>wfux7hOLO)R5VmT51} z1>HWJe6&|Ix5;)m>Tzy(B>oiJp8FlEtkjp?mB4!<0&K@o^reMV&7GsgmOdm92 zvi(oaTzd*!L|1m-PnsGU=5>O?r#v z>j=-E_k_A3(JXC>H?6q6l_`i`WAn)4n9;BmwXdub1+B~C#$R%wpsk8@bVquf=>8j4 zXttACKbb*I6Fb=A*Z?~BaU&@moX9I5K*zt9kIV32I)AB|_xt^m zziB>`JsA0wKWCyr9kpk<(OYtmlTF6=4>qH3R}*tdFogMvLulHuW$f7SQ2gV37_%(1 zSm>L{v_0Sn`L%>sYRfcesXw8&(;D7@(cesj=ApHZzWX;nhC-eTeAPw`sSMZUFI8E)^nz(t7V(6tPp`hpv&KAkSQ-K+s+8u@JCw;uw<^XHgd zb1+4{OU1_%*J8_Be^MCTi`F;OSmyZmd<&;XFU~yci>D1CeUt>gHMeFl(b4Q^oe#vi ztPrVn>wt8k70x`~_f48fl1hCY$}oE@jmu(DFI|}O)O~pF%0|)}o`z93&Y()YAM|gV z#MI?l@kXN^oSP(%*TV#~^mhuy{xo5yx@Vx?nO$_&e+0kKIGyR$sDhG|5%#nGiN)7{ zqE!4h7L+3d>F0B4%dbd!@=TYG?-Mbp8xI8HFimhQ@fAF}+{N411T*oL9Ok&Yix&hex<8D# zH$$W;Tp@XwD09h)X8zV8r2zw(tX>i|uHHb4?p*;-PG0aUJBNJd zT_ELyKiQ(7Aoi}p5O!=k33*C?QN8>e6kqj)X)71OiY+CqviAg&NSMss>+0o(UVFfD zgZJV^a$@ZpZq`(VX&b-@`fBT|QK4!h9gv5Q&F zx)V(K784CIErd&bxQ%cA4;(1j1il{!Lc*F`?B#}ry7Zp0wDq+;4xDm=+0BqVCN_S6VSSibhf5M=eBMVUdt|`vrzcn1ok~zX zj!KK(h~$RX;rFB<{)WwHD*7e`N!z4g$7fRt)QF}p72Vh>EW>Tv-Ux(MTXEcK4!XXt zVJkjb(6&-(@_(1jb8$-a&Hb?8a)}GXeA9*r6{&RNR5iE1z?p*nj3OnC<6Mu5FLCPT zlo#2z#<(_vwHs&RwmI*aP2+vmX(+<_$0KOmuG^f^)2&dD--zcG4v=QAJbyv4fcs;! zU6j}mhSKZf$WFDCoGKzoDkg^;9x#oad)7$FT{7p7Cp-Q&OqJQp2T*hK+VE%-{|I*V{|#jM&Z+?KWH zV3%#zOBPPAnPb&4=(!c$S%BUAKj9R#WKbRhWkP+IEyvTnUBqo=0= zF|IWUop%qSmSYXPfyW-kHXgwGqm%n^{cY@VZyG#UmWoMnJNY388nM?ifh5Ls*oBN4 z4i;D9sO|6AzIMEX%NqJ{umzU*a_1WiKX{6@zTXb^M=R@o>qn4o>=EWOYXqbY`OFSi z>}Q4YBf-Tx0e1F!(B4~m)HtJx^N88W-gwwD@9I=|v8_K0&l(CV28?GvWPf1tf<&;$ zyn?)jKkiKH+q+f4!QL~6*k1Pv=m8rV6k|A01Gpzqo9{n+7;Ed= zBT(I-PT`}9L3Q?YZuBo%x-Gk(W%OY~DZ=|$oqYghUNoUZ6(%c732he#WBs9KjD0m3>tu(Dn;)AB9c%^*dlp}UHOlcqf4gH4 zcWVpyoLk1U^Qwgh4_=`4$|dA245lIdkKy+k19m@ZJ{!JgiunD-eW;!yi+)ZjwCnc( zHs!FhINSI#TQ*{>(0ImWu-6KQ#t1E1x?&KvUfU|JoOPI*d|cSRyD#wf!6;DGG!af- zwhiw}4)z+sacN8owFfw+fC5P$fifNlOQFj8=5~3 zAUQun;ed@Y!r`7NkafsWSk(Ok`Y-v+^2}!8$<#RZ&Tk{z@ZB3rv*!vGsxOmQ)iJ!e zlW4Msy`LbQi+;AyC&e@&jO&8)%Mz%TxdZ6w zVTRUh^1L@j_>|iTqdZe#cl&YXqO338W0ej*dwK%#78<--VJtB}5W#@z0`-PCT=@&UXd`ENrkgXe3AZ|Ogj3RqA zgj(0%vAq(rgm*3u6&Bp|A%*XR-t)DD8*)aAr%C6C9e?@>mq!?hyB?>DU&(W1SCC01 z+k9Ea#xmGy+X04GwIN@iEdGZu=h{3Y=8Bz!#6QRXoiL{+A^0!C+<&qJCr+^c|Ie8F zcmCXH#eeR6Yxg2=sXpG|KZl3A8e{#0uJW~=M)|JwfNJMeIqle79yB}3^E|E2y9IrE?N|C~?%>rx`OuZVwA jpWQ!PPl?Hn|NKfw{6p&WaF_jua@l7iE%A^3zjXf(B+Gy^ literal 0 HcmV?d00001 diff --git a/evaluate_models.sh b/evaluate_models.sh index a2d8952..c4e5f72 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,6 +13,9 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm +# behavior cloning +python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' + # GAIL python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' @@ -22,5 +25,5 @@ python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-s # options GAIL-PPO python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' -# behavior cloning -python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' +# SHAIL +python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' diff --git a/src/core/gail.py b/src/core/gail.py index d630bd4..2e34333 100644 --- a/src/core/gail.py +++ b/src/core/gail.py @@ -1,10 +1,10 @@ import torch import torch.nn.functional as F from dataclasses import dataclass -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.trpo import trpo_step -from core.ppo import ppo_step +from src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.trpo import trpo_step +from src.core.ppo import ppo_step from tqdm import tqdm class TerminalLogger: diff --git a/src/core/ppo.py b/src/core/ppo.py index c742061..c73db2b 100644 --- a/src/core/ppo.py +++ b/src/core/ppo.py @@ -1,6 +1,6 @@ import torch -from core.sampling import rollout -from core.value_estimation import gae +from src.core.sampling import rollout +from src.core.value_estimation import gae def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): diff --git a/src/core/trpo.py b/src/core/trpo.py index d91e363..67d0113 100644 --- a/src/core/trpo.py +++ b/src/core/trpo.py @@ -1,8 +1,8 @@ import torch -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search +from src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.value_estimation import gae +from src.core.optimization import conjugate_gradient, line_search def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): diff --git a/src/eval_main.py b/src/eval_main.py index 6b460b3..381b061 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -11,6 +11,8 @@ from src.evaluation.metrics import divergence, visualize_distribution from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.reparam_module import ReparamPolicy from src.options import envs as options_envs2 +from src.safe_options.policy import SetMaskedDiscretePolicy +from src.safe_options import options as options_envs3 from typing import Optional, List, Dict, Tuple import torch @@ -62,8 +64,14 @@ def load_policy(method:str, policy.load_state_dict(torch.load(policy_file)) policy.eval() elif method == 'sgail': - policy = sb3.PPO.load(policy_file) - raise NotImplementedError + policy = SetMaskedDiscretePolicy(env.action_space.n) + policy( + torch.zeros(env.observation_space['observation'].shape), + torch.zeros(env.observation_space['safe_actions'].shape) + ) + policy = ReparamPolicy(policy) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() else: raise NotImplementedError return policy @@ -181,6 +189,7 @@ def evaluate_policy(locations:List[Tuple[int,int]], envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict.update(dict(options_envs.__dict__)) envs_dict.update(dict(options_envs2.__dict__)) + envs_dict.update(dict(options_envs3.__dict__)) policy_metrics = [None]* len(locations) # iterate through vehicles diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 27817ff..96b8437 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -7,6 +7,7 @@ import os import pickle from tqdm import tqdm from src.options.envs import OptionsEnv +from src.util.wrappers import OptionsTimeLimit class IntersimpleEvaluation: """ @@ -35,7 +36,7 @@ class IntersimpleEvaluation: self.env = eval_env self.n_episodes = eval_env.nv self.use_pbar = use_pbar - self.is_options_env = isinstance(self.env, OptionsEnv) + self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit)) # metrics present on every step of every episode self.metric_keys_all = ['v_all', 'a_all', 'col_all'] diff --git a/src/safe_options/options.py b/src/safe_options/options.py index a8d0e50..eeb5bae 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -3,14 +3,18 @@ import numpy as np import torch from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv -from core.reparam_module import ReparamPolicy +from src.core.reparam_module import ReparamPolicy from tqdm import tqdm -from core.gail import train_discriminator, roll_buffer, TerminalLogger +from src.core.gail import train_discriminator, roll_buffer, TerminalLogger from dataclasses import dataclass -from safe_options.policy_gradient import trpo_step, ppo_step +from src.safe_options.policy_gradient import trpo_step, ppo_step import torch.nn.functional as F -from safe_options.collisions import feasible +from src.options.envs import OptionsEnv +from src.safe_options.collisions import feasible + +from intersim.envs import IntersimpleLidarFlatIncrementingAgent +from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation @dataclass class Buffer: @@ -156,81 +160,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones) -class OptionsEnv(gym.Wrapper): - - def __init__(self, env, options): - super().__init__(env) - self.ll_action_space = env.action_space - self.options = options - self.action_space = gym.spaces.Discrete(len(options)) - self.max_plan_length = max(t for _, t in options) - - def plan(self, option): - target_v, t = option - current_v = self.env._env.state[self.env._agent, 1].item() - dt = self.env._env._dt - a = (target_v - current_v) / (t * dt) - a = self.env._normalize(a) - a = a * np.ones((t,)) - a += 0.01 * np.random.randn(*a.shape) - a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) - return a - - def execute_plan(self, obs, option, render_mode=None): - observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) - actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) - rewards = np.zeros((self.max_plan_length + 1,)) - env_done = np.ones((self.max_plan_length + 1,), dtype=bool) - plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) - infos = [] - - plan = self.plan(option) - observations[0] = obs - env_done[0] = False - for k, u in enumerate(plan): - plan_done[k] = False - o, r, d, i = self.env.step(u) - actions[k] = u - rewards[k] = r - env_done[k+1] = d - infos.append(i) - observations[k+1] = o - - if render_mode is not None: - self.env.render(render_mode) - - if d: - break - - n_steps = k + 1 - return observations, actions, rewards, env_done, plan_done, infos, n_steps - - def step(self, action, render_mode=None): - a = int(action) - assert a == action - ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) - hl_obs = ll_obs[ll_steps] - hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() - hl_infos = { - 'll': { - 'observations': ll_obs, - 'actions': ll_actions, - 'rewards': ll_rewards, - 'env_done': ll_env_done, - 'plan_done': ll_plan_done, - 'infos': ll_infos, - 'steps': ll_steps, - } - } - self.last_obs = hl_obs - return hl_obs, hl_reward, hl_done, hl_infos - - def reset(self, *args, **kwargs): - self.last_obs = super().reset(*args, **kwargs) - return self.last_obs - - class SafeOptionsEnv(OptionsEnv): def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None): @@ -287,7 +216,7 @@ class SafeOptionsEnv(OptionsEnv): o, r, d, i = self.env.step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -303,3 +232,29 @@ class SafeOptionsEnv(OptionsEnv): n_steps = k + 1 return observations, actions, rewards, env_done, plan_done, infos, n_steps + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs): + return OptionsTimeLimit(SafeOptionsEnv(Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps) diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py index 4e0e9d0..0325f59 100644 --- a/src/safe_options/policy.py +++ b/src/safe_options/policy.py @@ -2,7 +2,7 @@ import torch import torch.nn as nn from torch.distributions import Categorical from torch.distributions.kl import kl_divergence -from core.policy import SetDiscretePolicy +from src.core.policy import SetDiscretePolicy class SetMaskedDiscretePolicy(SetDiscretePolicy): @@ -21,6 +21,15 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy): a = super().torch_dist(logits).probs return (a * (1 - z)).sum(-1) + def predict(self, observations, state=None, episode_start=None, deterministic=True): + observation = torch.tensor(observations['observation']) + safe_actions = torch.tensor(observations['safe_actions']) + if deterministic: + _, actions = self.forward(observation, safe_actions).max(-1) + else: + actions = self.sample(self.forward(observation, safe_actions)) + return actions, None + # def torch_dist_nomask(self, dist): # print('no mask logprob') # logits = dist[..., :self.action_dim] diff --git a/src/safe_options/policy_gradient.py b/src/safe_options/policy_gradient.py index 6e968e4..803af15 100644 --- a/src/safe_options/policy_gradient.py +++ b/src/safe_options/policy_gradient.py @@ -1,6 +1,6 @@ import torch -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search +from src.core.value_estimation import gae +from src.core.optimization import conjugate_gradient, line_search def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): diff --git a/src/util/wrappers.py b/src/util/wrappers.py index 3916e75..3a088f9 100644 --- a/src/util/wrappers.py +++ b/src/util/wrappers.py @@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation): def __getattr__(self, name): return getattr(self.env, name) +class OptionsTimeLimit(gym.wrappers.TimeLimit): + def __getattr__(self, name): + return getattr(self.env, name) + class CollisionPenaltyWrapper(Wrapper): def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):