From 95cc78d940b21d534b3d93afd3086895ad6fe6b2 Mon Sep 17 00:00:00 2001 From: huangfu <3045324663@qq.com> Date: Wed, 4 Feb 2026 20:20:13 +0800 Subject: [PATCH] =?UTF-8?q?=E7=8E=AF=E5=A2=83=E4=BB=A3=E7=A0=81=E4=BC=98?= =?UTF-8?q?=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../expert_replay_env.cpython-313.pyc | Bin 19967 -> 20984 bytes .../expert_replay_env.cpython-39.pyc | Bin 9234 -> 7990 bytes .../inverse_dynamics.cpython-39.pyc | Bin 2008 -> 2008 bytes Env/__pycache__/scenario_env.cpython-39.pyc | Bin 7046 -> 7469 bytes Env/bc_env.py | 100 ++++++++++ Env/expert_replay_env.py | 182 +++--------------- Env/scenario_env.py | 46 ++--- Env/utils.py | 137 +++++++++++++ scripts/README.md | 2 +- scripts/visualize.py | 3 + 10 files changed, 296 insertions(+), 174 deletions(-) diff --git a/Env/__pycache__/expert_replay_env.cpython-313.pyc b/Env/__pycache__/expert_replay_env.cpython-313.pyc index 5feb9b13eefc149f70fb899287dcd0bcba42634f..a3f53e1b4bf6d9fc09c00f72b4120c65b335315f 100644 GIT binary patch delta 5989 zcmc&&Yj9h~b-ovG0vEsqF5V9kqzRG$MUfQW(1WBNCMB7Q%=* zDzRv#os8VZW6B#pR8K3?@ziWPZl$u5@sBo9Q+Miirmb$tF%UFOl>SR++L2^uQcr%Q zXZPX)5KU+DqcerXgR{Hm?Ai0}o^y8btGDU@`3`OU#A-ED@L6m9?R4Re^&MJcpwS)r zLsX|}59)cUPlHD(N}w_pg`=niDrP`SnhtZxu$axH=aym&TGwRl5(^SBZAq6)WzI@& z)To`LohYwuqA#I0wMqY$$$#Gyc=P!i&%c?ykzGA^Tes=yyRLi1f_|wDGP`XQ^*Y7V zV2MKC($zQ5JxM{6*BsXz(~6p7I#C;;1WgmwS6&AQ(umgdShdrjs(dC{bf<+w|E3So z9B?6>Kux;QolY0pWw6p7^pN9bz>jIas=({|XoXMFUD!)&WhBJw1JrZ0s3-ju6Fio( zSEO(Y`%OMztqCQ20+!?=7>SZT)Yav8@dg#hwz6naK;2z_^v+PLK{SA#2dy=6&~%}H zgujS}YZVNS!Cx`iNr}czY}TXBA5it^PPEn^Mt+kGd{yRa^(3Y!Dhb;-(GJ^K#~bya zHr<76#t@5#=JC*B7&1=wp*u%hsLN({ft(4;@pYJrX2Co;0C)z_1CzrtNbCZcOuXi- zZxb!k`_&%7tis8gJK$hVKPKw|kKmsLiFLsNrUML%)~UZ1m`<%|<7v^(TSc~F*@i6( z7GmEa=$k0PQGw~2iQB{5d3H#v!a2X2IHxRHQMD3PCsB3QsxEI+g>~?b&sjE4j;dZ& zzo!qXyy_=0!c|u7B&+`Qx}(%%$}Vu-4S${*7&%pfHz&Rg+jr(MC=cVkrN>ZrGZ*7{ z0+GP+G&YSvS9L{q4gBM;IbBRmyJn6~7jNddAx%Zk18{f+H*Zfk!L|vI_3$pjW2mg? zLrd$j8(ZT d3p{qm;qbgzbtsK-d{H!J;3nPCTM<_l(Fc$lRhE5J0S%0#{$ zCi2(xnfbu(I<%tqpzktvI)sn9rju79CDuvq zppTcM*j?^lGjnL>unC=H_4Jq2q)5^OeUc;zhovf@qz#BK&EzxWlzo}dDRS1ROow{> zA)cO9`Rey9PBAb-ru`dRco`?lNhMEJrGU#oxGEo#T)WiP^jw9TS2^`qLdNS+`)cIQ zseLtaYKhQe^O^E$BD7fHUQlt!VFsm#Jm`A{lWvSocH^{6=ZV)B(bI!1xP)9rjT+8V zTMVlyAKe{uqpJ>MuwGqA^?SBR@TkY?jxS26ra-a*!aJ?1)dVTz*O-wOeXs;FhqpGr zkpnOKw~ZKgBF_)sJBqwDoDBy)_CclPvmhr3LMW}NBi?|0~>>Z9@-GX^UR_bf&nH7Z0@>&qF$nk8@~vC(?Hj6Ty6TQX3t41BVaOy z$q7tynB;+!nA~FWk~ou{eJ(YZmvp%Q60${uodGVM>1ixP7CNn}N~bWF1QOFrWLb6j?2ME;pS_qWX_K(3w^4n}4;wKUo1p8^H)AF6 zJsE#VbLJ1RlJyHMOSkaMzt&{QiZ?|n1C)yy$O@!Y)ZO8Utzq@IF?k2b=Ro-m9(fl? z3@vv29ALTHc7;E}_aj9~H_`tPn#-!qscu4gevka`_+MoeablgaW& z#DJNO;yCPKRILa#0moUe@ub}HI#TZltz3*a7}rh`J-$}&a*dvr=;;Yn=Tcga_Zg~k zym7ZyG)%Xk_lH6TxSUxg+ws!O%ydiHRdmhi4mVMJs4gAN^i|3i7vs08V5PfEwF-9a4DMx7ims#qv^g;Bf+SQ>_>x41=;@`wCLar5@+2d4GV)$k7%uI!qo)F3N`3} zaxW-*jgY@VsF&2>0p(OxmVo}9@zDt~w2zc?n8hkE)u&Q~NDN~&aAyQ{BpM8N$vVlv zL85v@t^Z>xa}+b0+5q;Dh5jnxj(T&NE-I%1Cn49v&DhH`Gnxsw7vpm!`k)8bXOL|C zyri{8(l2_E+~c7IB-eQ$sQ+MZ6jBO)u-B5lta#24JZgX9S%A4;=XWa02byEe1cyH^y-?yfe8GG-m^P$F%JF1zG0gb1*?Bd zm-O@brNtC7?eANTwJ4bw*J!0$le$f65+oO@WSh&LU%Ze{%_QeYBF2wv=>jGfF}Z|g zoS9@UKa;w=n3D3U9E(RG4N)asCY^&CkJ}Jkl=P>VkR*PZT842WTnr>Bxg_DMWJ!M| zC1rCZT6j-uthXru*Jt!P_Q#+cv%N z>*K|6)6d>D|KaN-Es2-G`P^a+SMcLt!pQ4C*n~~v7Zn)@bF1Xs|Wmo(9_?D~h zhptf38+z;1jjO8*vUm6TQo%b^7O6_yk3?4&)@gLt;wA?WG zfnRPoRNOA{(5S0-u9!bHQ-RQ}!0!iEPZXN_<>vm|mp^p=z3&IU<&PTwHn`b5u_c@; z2s1ML_y5+40lY7S_sHQrh48Q(9=>b7ck0Z&@bJCx^oqI2dJAl`%r>u{xowuCd$-tq zq$?(~vGw@v({g-xiyhfXs-tI%jThUy-@3fTwtdo0wH%`#EO${Y@j`SjOjL-D%F)q6 z^w9qStX0g4zVKTMH*KqzWMAj)I@y<4(f`;RszAG3K|dH-_1-+N9+!iCx6^WP1a@|Z zuI+%YFrV84=-NSXhpzaZ55DfZcRDGDX8~5+sx(3o^z2B8$5_d1TNTWyEp~Sm?Oj{! zz_ziDo&phdu)95(tB6duB>K0 zW_lhv(2ouVTZ<;9V2aA7=$1+0+hkMQmMMPU<}BD6WLv|gt?|0Hh?)*vql@U>Ly6v^ z-MKRMmaux{o;`NY6#MK^9mNfQR$>o)_8;tG4t`CcKJK)h9@G^W>*@Wv;{Kj98XdaU z!LP`NCFk`q$DhR1EKID7qKm!Jj-1?R03rW-@Ws~H6 z8C`g4m-6TX=%^7Xgx-HDIAPcKf7&Ez^iMoGrytaJJj&Dcdgr4N+M^$Q)D_b|MSuEb z`Yf%F|Kg;!OYeI?0U=0B7~+^+ItE`Qy?EjL;*xY4xhLKT6!0>zdn9LGUo1frKStFbI&r2K<3B#&wLrHvZ7{61yY;0W#poCwwP=;B-vf zv?*yPjop`|PSYefA1#x#jWf(-(q!668vko&f?`_H>ojpXX{Y(o)^?f?Po{PEy+cAK zb|?ROSbO{S?e5#%_ulTFKYNzB{sd$B$YL=ecs@r5M$yZbXBhSZ!Y?xyw8JQ>!>_Y3 zHfqLK*|aIDQM7V0bwG0CF-@Ft;I}jtOalK^lhFLm=*RuqKuz3+(4(lAA^ON8Af9Pq zRlG)hp~NYK@ReE*zOLi&ea+r-!@@+-tG6US08TG{C+foI*BS9YT8yTa1?6F)#EZi$ zonpecXN4E{x4H?VGmLmvKq+Hhi-7RqT6g_mCxQuHv%|FAEVA1Z2TY*{@_&{WF z;rGoRHeQY+tV^TqU4~!Y5ySyqIkOyZ>MqlQ9Ks#zOL0nP#h*93`UlMj^;YNC<8^s6 zlT0J9rrd>@a$aAC4l<&S(n||;!sWSP-b#6&+-Nw(|hs#&2Ic_ zgNuzf<3q*6HJB-JG3EG8eZ*JL+JHnu16e0XH4{DfO@o__tK1vFvmoAYurXR9o!5KD z5HT2aAko`~W7Z)4UaSPaYtZ9A154V%xrGwGgiP}$0!5QxifUoY|BqXx%C?YaeDUM7ENs2y(btrUKqqCmOw9K5;p5rh(7yCGP6 z=(e#|7bm!QjxOlpS$x>!Vs68mtwC!=acB$Nje?su^X~O*Ub*N27rjC$Z*O(NB5;cY z4_``g;=UE@%+(VfqM!HDRm@8QY3A8xmd@N5U^8!4>CJ3G!l!Du2{(5OEgaF(N(4ZS zv7pF*6Gdf3G!q|KSprw71}=pQe`dBawFJ`|;nw1Z+I-|dPk^@nr(-RP@%hq@@j?7= zZuw#+P%rZx)XOo8hfFa0$%WLDq!sU_ck`2GKXWHNE-%=^kigt-NAmdMd!PSBav_b}5d*w3w07 zsZ?^PFEN}HGxE%Nr;B9*_=3BYX_&e0eude6n_?JA$caowl0x91qECv$qheCA45pF^ zsV_ZnFgci!6>Vx%&M4Y@lH;`Rb;t}99NU&YUa3FN~yAdP>!q1a&>cD?9$AuWe@0~ z?5>LWXZA(5va4?=)kT-IY3qDWBP06=jFNPn4Yn5alGOze|^< zHYp}KlaMm@2*U0v|ICKyFgr0tG+1tFhIEcB^FP;@NUsppQzN8j09Wh_ z2dxh*qQ^^mmS}mBfad@d1<9^P>3K5sV*oL{q5flt)2i5U){Y64G$&FA0gvEcGz44r z6O_h+9C3Bmen2n-0rarh;Ar&?CFRVhn8=K#MahLtjiD~m2B={x?D9hFq$glXG17Zg z*i6IGa=LYkU2yV-l^Rhm*aQbOMg%wDthP3|ZuG}vzo;)I%ud0?qx_c@ ztgyhs#E5o{84AMtT`qiY8)u$cup!UUNLBgU=x39jinc7)#I=j)dV)@frNqVsmx^94 zbh}4p8>^KB1x^c1iKr6qC1h4(sHy+^5c6u+L`C@YaJ^% z>E~qrB?3-j{pOySQA!-^8x;?v6-|0Trs-BONXcUfX-Jkts)x+Vav9+Ga9LmrEly;8 z#l5sN%bDcZ7%3PRpz7U+E*nva@ zBWrV!joHY?Tx3(0E1&F|cL#FrnykBKa_eOd`OBX+J#L!pp7#X_GPU<{>ty#Oo9{zdsw;G`VyB zCZwAPevt&BnJBNC4^+w>G!W{6{Hs(0{iYtQ* zs_0-o8exz>I@yKq>a22+hKRIGIj%m-)lcu6q< z1^=-#PzO3oEIDIU)>t)Xte!V=Ib%)MSTkp=g&CVGXANbop^MhYq~;PX-*$?*grDBl zy5^G2d3xJfVQTBFEjnwAUb|{Ru1(hzd;7JkdgM6F$mF8Ez1G4b?X26v>$LCant0O; z*VUkf_M~m+#)3dh-jF5~-_EM$tfz`XTurFlh$Bhj1PVRX3CwA4^6Z zw8>RFxDAA-4S-@9NhXFy#o@lpND?}z8@oDsr=&r`JVXGwwu*MWoR%`sEFaG##UZJm z&~yn&N{lOp1F3W()3iEjS9>fY+^)W%(ZS5Z4aBjYyEeW>1dL*Gba-SSEh&aWqp}PJ zNWTHPYzIJ_ha#Z%Jmj`3LA+yUkon8Z2R-jG-PCETMEg?O=8&5!k*_Ohi~uO;YSxlc z1lvsjO>DaeMl%=9ZeHxz(;jjVBI$?J=8k3L@a8P_QFfiiAN5N vjZ1f9>)yx8vOrN;NYYV^ZHLlBM^edl=@a~`y*19AM#eYmsrwiag!26tLk@9j diff --git a/Env/__pycache__/expert_replay_env.cpython-39.pyc b/Env/__pycache__/expert_replay_env.cpython-39.pyc index b0935b989b1fb9f68a8afe2d92a3cadecceff11e..fcb7475c31ef6d2fd36ffd8652c0ae08218072ab 100644 GIT binary patch delta 3436 zcmZWsU2Gf25x%|Skw^0QBT=H{AK8+n*ruJTvK_mLEX#Hh18xh;O&t_1WA}XGo~Wab zcl6##mW*TBRt^FbNUHjP0?88zv`CE>H5vmgiazwELHkhjz57%I1qweT56xSP0-f0- zB|)LUxw+Zd*`05`ouwXM|LtPl%;ho?eExW0%l`58-{ec=@HgSrlB$GCRhbswkyZ$m zsq!dYRp?N3i5!~BK9Z^_I!tw%dqk=l&C>!{((gzKshYVZRZ7th$lHWRDY->5(T04M z6rx|sOGJX?9D-H{pA%G2PCDQU`W&tTEF~h8dhDP#{!y zUT7X~XT^5Jcv}y%p+1->o+q}H+e$kJGji>Gn1k&dZx_M>%^i^1#jqIa>#zm(5zHzc zzy$MQaRE2=QHXO+hIyD#Zg+go+ZY9Q}d#aC1wsE@q;*Nfz@ukLW z^kZ(ept*#);J>yd^R{2ytGBY`#fkDg!UB}AMfmkJ(RjH_ZNPp zzu4rq!mePj1{Na;+I%0YBIrf2bnF2Pf^7`+`~is;0FB{rxO=fD!S+50E(8nROW|;j zY!O$cq}?e&hyq1dgsBAXShd(P)VGZ8ScdzSGA`U*3`Y<#p2m{^Ffc6w7CvXxMtJhvQ2rm=D7j$9LJ&tSC;F>n-x;ct$3_v z8PxJ&Xp|CPbu-wtfg<*8N6q_$h<+PdXyvkMPt+|v(H?mWJlO(Z& zC<-Ww(f6`%b@HeVp%@0ivzVy2+^}k}u4)HX)8{F$nL)s~0(tNZtj_}ffgP;!q}gg& z5I$FPER*fb8`=W~BR+o)f+Ubl|lX*z1G% zXN-;DSi-j~%65@IgsbY1GuAl9MDd`E3SiC-$I)zBO^I}D8t+mc&;xT|S40T#AB1qZ z3W>5!s1M9I2E%M7fuW**wi(_jpu5M3^!&$B-i#z=>!{Z;RQ z5!g+O9fft-1PGpVEVm9cQXkkfSbZY1@8D1shty@OZoBL_=uZ+cF!^S@ako6eVU1+}w=C$bQ(7ot)r5yQ1^l@ps^D}5Nq^nbf zCLT$l{5o7n*kUNQ7j-Gd;x5dqVWSy*WKeKxx8UpJv8+R{x|3;>kTj&9$oCVs;)+&? zz!CbBJUho~;J9x~yqsCUqgC^ot;LUF{B!t94=+cZ(&^6Z0(1$}@fSfE_*8NXv*^Qy z8lo2Cc@DKj6o@QWpm^*qzP^Ow5{mO6`l#_V5S&>}Z{4aEfvnh2283noO-Y)N#Ey9y z&zj)}R%`F(qxsUL{8WnmJ2D-$OY_m`(a$^I#wCuRK=$mCg}sHZw^0Z>+ySj3@A88y z(4)k;s1(HMiK%e`CkTERs3lS0>407WJrJ9=?V8wG8TLr7M?Wtgyzw%6iG=4-Lv$)? zOp|5w(&8RAAW5pv*}bXZyz3zTEIrf5{;DiK$I`QD6oS%7Lma!L&|;p?aPAXQ&08BSivm*c#Ew3kI8kn*-37r_aYKt{!=q0oZckkkR@+?x z6aU@JrE?#gpJN|~ z87d7nTUKR=?Ita_f`q8%`WL?utinaTKVwm*qyt=b>hOcbV@&}S+)N$|c+%Pk@N^*d zTs0XpH+gEs@yy`hxwDmGe2Hn^vTq=7SPVq;>!}0q&%$4(WvKl+t6IguCFe8 zjHjBm@53(<)&cvzR2gDij~0$S?&vsr6@}GPkvPZCpE@Sz|Z`XZ{pCxazJoTbnU!;n$kuL~l%Ag#Eue{YFs~ MZdZ&)pH0vFA4MKyw*UYD delta 4715 zcmb_gU2GiH6`ngYyF2^uzpQ`bKz>NJ0mlSVN-#+ZBv4esPf`j4VLIMBUe9{Hv%WLV z51M7GCJ3n#s>M7tLe*OIp?Rrx>sBl!V3=e2edDlgcKWP&$P$9~QvLAKelrW#xMJ1@~_p5w|o+t9>kg!hgS7v)zYZ_T-Me``K?I5>Wt0eb}4qY^s-hHod# zz#iMedf(|E{CNMc;SF8HB&>0=@dT1lX)N^9kAh~oA29Itf z*LBuX5u+0jQI-}W?<0gdlLHDx2T`U56pC}AOw+-=C&NjZ#)O{=rv~`xa5}+zXEv1! zno`{teUN^>`#EY7ea_Uh1d-4HOyjwL$NvdHh#I-B_#+`}z=-RLsBm^cdHXebZ?Fc+ z5J&BNUq!jmI(n!{o`oMOT)UuzT2Bj+Q&>x(fN2fAlL=B`8urd~xjwCg>bSDPL<8&) znrve57E@lA0_Jffifroq#VRxN!U@(z%WB#!=iFTtx>eJ0zP2ue(+aFryB!yVXLGCN z1Xnzytwd!2#(~8-*SEYx?JC~= zDF7wi9bvkvsugCiv{u41#?tcpcJ=Jlwj+X9oOaW`cD8kS4ktUo+=5iwKQ*^4E-nYc zURrW%13mTNTHA?Jet=uWsmpfL<#!p5yZbdhg4)r=*ud@FpzA_u(nL|l&QzH=0OVY3 zTy}7b8(f3Eb^EI8_gTI;Z5pdiU~@aL&8)ZBaB6|&@;I~X*sx(ePB)!a9Y!JpRI{;) zwJhv)#R}Y2XtLUZOm@YY-0jNZ-MKQuu+lLSfx%F*r28{HzILj;BiWC5VO?b|^h#38(y zWuMn9e_JbaY%8?3j7*o0M{BmQ7F`i6!~08~)pY$JPPu`z>I+^LYm&9?wWDL% zLwT_m{9=|iMeXc=x5o5k>Q%x3WV2eI^066wgNwT#~6?PkUbIIG-gv`1Y4p? z?o@Y1i~@~y-*LFuHWG!SI7o6c(gH3wCoFi_3Or{UnFKTg3QwG~0W zGIL9=Kzfm@&1@65Nl;_ezlcLgu0Gzjs#v{H!be^NR^tiip!f>Sk3a=xTBecWXc4$W~#yb zys1OVylM197Vi)>zbjkCIoyhZ$Nx0S<)^w@ZYHf#VvevJ8)G_4>l3Vm)Ge``S^!N| zbL>7=&~jMQ*${Z>(e(=Dx9edaVpg(}$629t}g}kceaLTTguX1_a zU}M@RPJTb@%6y2m#k!TG#~=%CsoNQ`oiDd^gb`d8mo7WYE*yz^*IIJx;wUbcWE2$! z>BVYZJWlumfkz2EM&M-vuK>gazx|edrA1+Z#23^s#1Cj8mw0Z~_E!)-(s4mGEd$Ne zd=vo<>0+j!G}8PM?O>O8Tl^)E{}{E4E8P?@dgQqLQy}Q~3bkF#xoFnZZW>xorwRAc zGiWMkIicE1akk0CRDcGe`Q~3Ek4pi`Su&umGwBL0a!JleAx)lV-S^y>nj&Xse0@u@8 zq8R;dc;5@>i9`BrkP?08i63DpHe}r$gm|W(al3E#B6v@;PaYciR68YT6cC3A$b_U~ znq=jJz^dAgOAw)c^XK&W?P>I7l?TzED#uTL4!Zovb8e}zIyax#_g<^Uxdralf||G5 zKJ{lr&A*lCmlgYfWLkxQ-D-NE{RMoY9UU1LT{$&Cp)#?a1~5}^JHqn??Y+g!k|W0z z%dWv>EkBV)gO4CVf0crORA&?8Q|C>3uD9r)r39c640WAu67 z#}p~9l4|!xuk^*=Bfj^8Zk|@DKPUA{dP@3CdXIb-7cI;zwvREc*|e-3qQ+>ObcX?q zg3zbI$Y7;%uB)6dP8#r delta 20 acmcb?e}kVpk(ZZ?0SG$N(>8LSWd{H^BLzhO diff --git a/Env/__pycache__/scenario_env.cpython-39.pyc b/Env/__pycache__/scenario_env.cpython-39.pyc index d924eb81edd35c1f77dc6010c435cc7479a0b53d..918470ddc3723cbcd124f61b61ea199f37c24524 100644 GIT binary patch delta 1757 zcmah}O>7%Q6rNe{di}Sy6UT||rm1P%5T{WnQma7KMuHX!r9ZUL3QWSX@vM{GIO}wF z(~@dU3J&FhP|+Nyph6po3x`Ue-Z(%na0OLdSaITlgg7H4F1)vK5v_!<^}hY)n>TM} z-g~>(XKx?Q8k&|O@Vh?Uvs&X{X5%WoCW`7L9Taow)Ip}P>=tRH&5Ws;aWiY??0icq z6Ev2)MOmH=Y{`v*_lQgygD;VKAuNgylf^rqtD{mb4uMs$7??Xx#g9pQ)6;yy;uh(} ze4=N466ChYJn}q_yD>j@oGc7=WuMq-d%zwH3ckEeDY^7MQ(0n*fIbuyx5;tRqg};U z_;=_T4EqXtMYAK@Bx0rDK1&DX9WB_b^e5Byu3+>ZzS|zN_t;f9!Xpr3FC**p2IUjT zcY-hso!0kS@YKh>nzyeI$`vba_SMGna;anb9)w6gv*%VV?%7 zAB$2bGwiWgnP8QJ#5Uo}UeZeiX{>aQIuet+abMYtrfzx}U+Gb%l*oN27Kgw5E(AdL z5{VY5z62Y{Ooh;iGV$ZUahoJ^mcq7vKv8y@y|1Ow>j4jyZ9cq6W&K4gb1UjexT6GZ zLt8Ou1nsK#G}f=$GRV-|^xz4%ic=rpeGs81xs%I~ooLkJ? zS8a88wPm$ElbftM*)|;WF&7%fYSyc+**2S=)p0~4vxlA$8<|~s>_c7!s4M&^JjAD& z(Q;kmCs03ya2g;~EYDnZLm_7;EBqByBxfy*8+>Wicr%Q(O{bpV1#vpJyTD(-Rj);> zpbMKT{B?0HS2@N@s7Fvnb`rn=b87%ZlBrCGX_iW0DVm`bdH|FZj3+>u_$#GeJp|v9 zk+J`2___Sor6vSyq?;YbeF#_9P$(5 z=OLYLh$F=(Y8+i%gav?x)_+N%+)K+Ie+wtp#m(X&r2+JlJNJrR3Y)$-@+_SbAB~L1 zzekU!#NCnh*cQ@7%Q6rNdo*K4oWKlyR&{IyL&h(qX~Kq(Q8kSK+sQlbJLkkU$=Qq8HThqc7D5>YrHAr2h40qS|_l?z;uKwRL!d$uVm2wU^*H{W~nX5KgR zVBz+B%GC7)f#1TLd-k^zhbbjSZ;HX#49$qe*zBC8SxHN=bSvegD^h{rVETQ^QY?Kd zTFTrZQBumjMT*&QP~24qj}Bv_QW`62Fd{ITS0wQ$?(9@aD^evsF?63engjKMOq*OF z@BX$dwW2NAkM{KJn#4cZ*xlImWu_E}*XQ(uvg=2J0YBQNlz<+CQ>X>fQhX&CY!m*_ z%LPNe(x$G&;;mRu{nN|)%06X@Jn?oq35>u${!yer+B8o($trBCLur|5P&8Q}eyq2F zj+{|!`3WV_HKM1bFzbm7^*?J=me}boV(I%b_5j^Qp!Zzs@2=B9q1)X7-PmRsmc0dD zIm5vSdbCelaX)@Q`4{L>FzUxUK45&TO~4W2j6Y^taIC8Z$9pl&IT7qkJf2KCQ_fS) zG+f2guyBTv8}vZpdF0(93=3+;jE&b-^k%(h_9?JAV9$2gN$=D?36A-+dyL?8Z^JWf z0#Bju3a4=KVZ2ta8a3;Es}{;kENK&BIWs1eJbT>=hvE8Riy)w0jh|4>XP_5dzilfg-Tkv?s~jlt69wGw9j7=pZA@ay$Hc# zieH5fpGR22>}0cHZn{Q&#j(mB*HOPN(*3S%0o^MAmG3a)b@5)IHok+jW8yzf?t%7wH|3|o<+uMY zfHH2$Klt2R8sp39|wXLZ4QHqYu0yJ{}%F{{hlYFB8aG0AY+l!}D&|Aky|Lcd;@89vBcIYY#O~;WzX0Q(R?+|f diff --git a/Env/bc_env.py b/Env/bc_env.py index b802fba..59c080c 100644 --- a/Env/bc_env.py +++ b/Env/bc_env.py @@ -1,13 +1,113 @@ from Env.scenario_env import MultiAgentScenarioEnv +from Env.utils import filter_traffic_tracks_to_birth_lists +from metadrive.component.vehicle.vehicle_type import DefaultVehicle import numpy as np + class BCScenarioEnv(MultiAgentScenarioEnv): """ Environment for Behavior Cloning Evaluation. Uses the same 45-dim observation as ExpertReplayEnv: - Ego State (5): x, y, vx, vy, heading - Neighbors (40): 10 nearest * (rel_x, rel_y, vx, vy) + + Spawns background (static) vehicles so that observation distribution matches expert data collection: + expert data is generated with ExpertReplayEnv which includes bg_* in active_agents, so the policy + was trained on obs that can include those neighbors. Demo should use the same scene for consistency. """ + def reset(self, seed=None): + # Clear background vehicles from previous episode so engine.reset() passes _object_clean_check + if getattr(self, "engine", None) is not None: + ids_bg = [ + oid for oid, obj in self.engine.get_objects().items() + if (getattr(obj, "name", None) or getattr(obj, "id", None) or "").startswith("bg_") + ] + if ids_bg: + self.engine.clear_objects(ids_bg, force_destroy=True) + for aid in list(self.engine.agent_manager.active_agents.keys()): + if aid.startswith("bg_"): + self.engine.agent_manager.active_agents.pop(aid, None) + obs = super().reset(seed=seed) + self._spawn_background_vehicles() + return self._get_all_obs() + + def _build_birth_lists_from_traffic(self): + """Same lane/static filter as expert data; return background_vehicles so we spawn them (match training obs).""" + car_birth_info_list, background_vehicles, obj_to_clean, stats = filter_traffic_tracks_to_birth_lists( + self.engine.traffic_manager.current_traffic_data, + self.engine.traffic_manager.sdc_scenario_id, + self.engine.map_manager, + return_stats=True, + ) + if stats["n_controlled"] == 0 and stats["n_total"] > 0: + print( + "[BCScenarioEnv] 0 controlled agents: total_vehicles={}, off_lane={}, static={}, no_valid={}.".format( + stats["n_total"], + stats["n_off_lane"], + stats["n_static"], + stats["n_no_valid"], + ) + ) + return car_birth_info_list, background_vehicles, obj_to_clean + + def _spawn_background_vehicles(self): + """Spawn static background vehicles so they appear in active_agents and thus in obs (same as ExpertReplayEnv).""" + for sid, car in self.background_vehicles.items(): + if car["show_time"] != self.round: + continue + bg_id = f"bg_{car['id']}" + if bg_id in self.engine.agent_manager.active_agents: + continue + vehicle_config = {} + if "length" in car and "width" in car: + vehicle_config = {"length": car["length"], "width": car["width"]} + v = self.engine.spawn_object( + DefaultVehicle, + name=bg_id, + vehicle_config=vehicle_config, + position=car["begin"], + heading=car["heading"], + ) + v.set_velocity([0, 0]) + self.engine.agent_manager.active_agents[bg_id] = v + v.valid_mask = car["valid"] + v.start_t = car["show_time"] + + def _update_background_vehicles(self): + self._spawn_background_vehicles() + to_remove = [] + objects_to_clear = [] + for aid, v in self.engine.agent_manager.active_agents.items(): + if not aid.startswith("bg_"): + continue + if hasattr(v, "valid_mask"): + if self.round >= len(v.valid_mask) or not v.valid_mask[self.round]: + to_remove.append(aid) + objects_to_clear.append(v) + for aid in to_remove: + self.engine.agent_manager.active_agents.pop(aid, None) + if objects_to_clear: + self.engine.clear_objects([v.id for v in objects_to_clear]) + + def step(self, action_dict): + self.round += 1 + for agent_id, action in action_dict.items(): + if agent_id in self.controlled_agents: + self.controlled_agents[agent_id].before_step(action) + self.engine.step() + self.engine.after_step() + for agent_id in action_dict: + if agent_id in self.controlled_agents: + self.controlled_agents[agent_id].after_step() + self._spawn_controlled_agents() + self._update_background_vehicles() + obs = self._get_all_obs() + rewards = {aid: 0.0 for aid in self.controlled_agents} + dones = {aid: False for aid in self.controlled_agents} + dones["__all__"] = self.episode_step >= self.config["horizon"] + infos = {aid: {} for aid in self.controlled_agents} + return obs, rewards, dones, infos + def _get_all_obs(self): # Implement custom observation: 30m range, 10 nearest vehicles obs_dict = {} diff --git a/Env/expert_replay_env.py b/Env/expert_replay_env.py index ce516f0..d850244 100644 --- a/Env/expert_replay_env.py +++ b/Env/expert_replay_env.py @@ -35,132 +35,44 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): if self.engine is None: raise ValueError("Broken MetaDrive instance.") - self.background_vehicles = {} # Vehicles that exist but are static/background - - # Helper function to check if a position is on a valid lane - def is_on_lane(pos, map_manager, threshold=2.0): - # Check if point is close to any lane in the road network - # This can be expensive if checked for every point, so we check sample points - # or rely on lane index if available. - # Waymo tracks don't have lane index, just positions. - # We can use map.road_network.get_closest_lane_index(pos) - if map_manager is None or map_manager.current_map is None: - return True # If no map, assume valid - - try: - # Use a larger search radius to catch slightly offset lanes - lane, lane_index = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True) - if lane is None: - return False - - # Check lateral distance - long, lat = lane.local_coordinates(pos) - width = lane.width - # Allow being slightly off-lane (e.g. changing lanes) - # But parking lots are usually far from defined lanes in Waymo converted maps - if abs(lat) <= (width / 2 + threshold): - return True - return False - except: - return False - - # --- MODIFIED SECTION START --- - # Capture expert tracks before they are cleaned + self.background_vehicles = {} self.expert_tracks = {} - # Capture SDC track for ego replay (MetaDrive default agent) self.sdc_track = None self.sdc_vehicle = None + + # 在加载新场景前,必须清除上一轮通过 spawn_object 生成的物体,否则 engine.reset() 内 _object_clean_check 会报错 + # 从 engine 当前对象中按名称筛选(与 manager 无关的对象需在此清理),并强制销毁 + ids_to_clear = [] + for oid, obj in self.engine.get_objects().items(): + name = getattr(obj, "name", None) or getattr(obj, "id", None) + if name and (str(name).startswith("controlled_") or str(name).startswith("bg_")): + ids_to_clear.append(oid) + if ids_to_clear: + self.engine.clear_objects(ids_to_clear, force_destroy=True) + self.controlled_agents.clear() + self.controlled_agent_ids.clear() + for aid in list(self.engine.agent_manager.active_agents.keys()): + if aid.startswith("bg_") or aid.startswith("controlled_"): + self.engine.agent_manager.active_agents.pop(aid, None) + if self.replay_sdc and hasattr(self.engine, "traffic_manager"): sdc_sid = self.engine.traffic_manager.sdc_scenario_id self.sdc_track = self.engine.traffic_manager.current_traffic_data.get(sdc_sid, None) - _obj_to_clean_this_frame = [] - self.car_birth_info_list = [] - - # Pre-filter: Check tracks against map AND check for static vehicles - - for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items(): - if scenario_id == self.engine.traffic_manager.sdc_scenario_id: - continue - else: - if track["type"] == MetaDriveType.VEHICLE: - _obj_to_clean_this_frame.append(scenario_id) - - valid = track['state']['valid'] - if not valid.any(): - continue - - first_show = np.argmax(valid) - last_show = len(valid) - 1 - np.argmax(valid[::-1]) - mid_show = (first_show + last_show) // 2 - - # 1. Lane check (existing logic) - points_to_check = [first_show, mid_show, last_show] - on_road_count = 0 - is_valid_track = True - start_pos = track['state']['position'][first_show] - if not is_on_lane(start_pos, self.engine.map_manager, threshold=5.0): # 5m tolerance - mid_pos = track['state']['position'][mid_show] - if not is_on_lane(mid_pos, self.engine.map_manager, threshold=5.0): - is_valid_track = False - - # 2. Static check - # Calculate total displacement and max speed - positions = track['state']['position'][valid.astype(bool)] - velocities = track['state']['velocity'][valid.astype(bool)] - - total_displacement = 0 - max_speed = 0 - if len(positions) > 1: - total_displacement = np.linalg.norm(positions[-1] - positions[0]) - max_speed = np.max(np.linalg.norm(velocities, axis=1)) - - is_static = False - if total_displacement < 5.0 and max_speed < 1.0: # Relaxed threshold: <5m move and <1m/s - is_static = True - - # Decision logic: - # - If off-road AND static: Skip completely (don't even spawn as background) - # - If off-road but moving: Maybe keep? Or skip? Usually off-road moving is weird, skip. - # - If on-road but static: Spawn as BACKGROUND (visible but not controlled agent) - # - If on-road and moving: Spawn as CONTROLLED agent - - if not is_valid_track: - # Skip off-road vehicles entirely (both static and moving off-road) - continue - - if is_static: - # Add to background list, but NOT to car_birth_info_list (which is for controlled agents) - # We need a way to spawn them. Let's add a separate list. - self.background_vehicles[scenario_id] = { - 'id': track['metadata']['object_id'], - 'show_time': first_show, - 'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]), - 'heading': track['state']['heading'][first_show], - 'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]), - 'scenario_id': scenario_id, - 'length': track['state']['length'][first_show], - 'width': track['state']['width'][first_show], - 'valid': valid # Need validity to know when to show/hide - } - continue # Do not add to controlled list - # Store the full track for replay (only for controlled agents) - self.expert_tracks[scenario_id] = track - - self.car_birth_info_list.append({ - 'id': track['metadata']['object_id'], - 'show_time': first_show, - 'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]), - 'heading': track['state']['heading'][first_show], - 'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]), - 'scenario_id': scenario_id, # Keep track of original ID to lookup tracks - 'length': track['state']['length'][first_show], - 'width': track['state']['width'][first_show] - }) - - for scenario_id in _obj_to_clean_this_frame: + from Env.utils import filter_traffic_tracks_to_birth_lists + traffic_data = self.engine.traffic_manager.current_traffic_data + car_birth_info_list, self.background_vehicles, obj_to_clean = filter_traffic_tracks_to_birth_lists( + traffic_data, + self.engine.traffic_manager.sdc_scenario_id, + self.engine.map_manager, + ) + for entry in car_birth_info_list: + sid = entry["scenario_id"] + if sid in traffic_data: + self.expert_tracks[sid] = traffic_data[sid] + self.car_birth_info_list = car_birth_info_list + for scenario_id in obj_to_clean: self.engine.traffic_manager.current_traffic_data.pop(scenario_id) - # --- MODIFIED SECTION END --- self.engine.reset() self.reset_sensors() @@ -260,38 +172,6 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): v.valid_mask = car['valid'] v.start_t = car['show_time'] - def _update_background_vehicles(self): - # Remove background vehicles if they become invalid - # Or spawn new ones - self._spawn_background_vehicles() - - # Check validity for existing - to_remove = [] - for aid, v in self.engine.agent_manager.active_agents.items(): - if aid.startswith("bg_"): - # Check validity - if hasattr(v, 'valid_mask'): - curr_step = self.round - if curr_step >= len(v.valid_mask) or not v.valid_mask[curr_step]: - to_remove.append(aid) - - for aid in to_remove: - self.engine.agent_manager.active_agents.pop(aid, None) - # if aid in self.engine.obj_to_id: - # self.engine.clear_objects([self.engine.obj_to_id[aid]]) - # Instead, we should find the object by ID and clear it. - # Since we don't track obj directly, we can't easily clear it without obj ref. - # Wait, active_agents stores the vehicle object. - # So we can just clear that object. - pass - - # Re-iterate to clear objects properly - for aid in to_remove: - # We need to find the vehicle object to clear it. - # But we popped it from active_agents. - # Wait, we should get it before pop. - pass - def _update_background_vehicles(self): # Remove background vehicles if they become invalid # Or spawn new ones @@ -314,7 +194,7 @@ class ExpertReplayEnv(MultiAgentScenarioEnv): self.engine.agent_manager.active_agents.pop(aid, None) if objects_to_clear: - self.engine.clear_objects(objects_to_clear) + self.engine.clear_objects([v.id for v in objects_to_clear]) def _spawn_controlled_agents(self): for car in self.car_birth_info_list: diff --git a/Env/scenario_env.py b/Env/scenario_env.py index 8488854..6f1f531 100644 --- a/Env/scenario_env.py +++ b/Env/scenario_env.py @@ -76,28 +76,9 @@ class MultiAgentScenarioEnv(ScenarioEnv): if self.engine is None: raise ValueError("Broken MetaDrive instance.") - # 记录专家数据中每辆车的位置,接着全部清除,只保留位置等信息,用于后续生成 - _obj_to_clean_this_frame = [] - self.car_birth_info_list = [] - for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items(): - if scenario_id == self.engine.traffic_manager.sdc_scenario_id: - continue - else: - if track["type"] == MetaDriveType.VEHICLE: - _obj_to_clean_this_frame.append(scenario_id) - valid = track['state']['valid'] - first_show = np.argmax(valid) if valid.any() else -1 - last_show = len(valid) - 1 - np.argmax(valid[::-1]) if valid.any() else -1 - # id,出现时间,出生点坐标,出生朝向,目的地 - self.car_birth_info_list.append({ - 'id': track['metadata']['object_id'], - 'show_time': first_show, - 'begin': (track['state']['position'][first_show, 0], track['state']['position'][first_show, 1]), - 'heading': track['state']['heading'][first_show], - 'end': (track['state']['position'][last_show, 0], track['state']['position'][last_show, 1]) - }) - - for scenario_id in _obj_to_clean_this_frame: + self.background_vehicles = getattr(self, "background_vehicles", {}) + self.car_birth_info_list, self.background_vehicles, _obj_to_clean = self._build_birth_lists_from_traffic() + for scenario_id in _obj_to_clean: self.engine.traffic_manager.current_traffic_data.pop(scenario_id) # Clear vehicles we spawned via engine.spawn_object() so _object_clean_check() passes @@ -126,6 +107,27 @@ class MultiAgentScenarioEnv(ScenarioEnv): return self._get_all_obs() + def _build_birth_lists_from_traffic(self): + """Build car_birth_info_list and obj_to_clean from current_traffic_data. Override for filtered (lane/static) selection.""" + _obj_to_clean_this_frame = [] + car_birth_info_list = [] + for scenario_id, track in self.engine.traffic_manager.current_traffic_data.items(): + if scenario_id == self.engine.traffic_manager.sdc_scenario_id: + continue + if track["type"] == MetaDriveType.VEHICLE: + _obj_to_clean_this_frame.append(scenario_id) + valid = track["state"]["valid"] + first_show = int(np.argmax(valid)) if valid.any() else -1 + last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) if valid.any() else -1 + car_birth_info_list.append({ + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": (track["state"]["position"][first_show, 0], track["state"]["position"][first_show, 1]), + "heading": track["state"]["heading"][first_show], + "end": (track["state"]["position"][last_show, 0], track["state"]["position"][last_show, 1]), + }) + return car_birth_info_list, {}, _obj_to_clean_this_frame + def _spawn_controlled_agents(self): # ego_vehicle = self.engine.agent_manager.active_agents.get("default_agent") # ego_position = ego_vehicle.position if ego_vehicle else np.array([0, 0]) diff --git a/Env/utils.py b/Env/utils.py index c19bf24..8dfced6 100644 --- a/Env/utils.py +++ b/Env/utils.py @@ -2,6 +2,143 @@ import numpy as np import torch import random +from metadrive.type import MetaDriveType + + +def is_on_lane(pos, map_manager, threshold=2.0): + """Check if a position is on a valid lane (within lateral tolerance).""" + if map_manager is None or map_manager.current_map is None: + return True + try: + lane, _ = map_manager.current_map.road_network.get_closest_lane_index(pos, return_lane=True) + if lane is None: + return False + long, lat = lane.local_coordinates(pos) + width = lane.width + if abs(lat) <= (width / 2 + threshold): + return True + return False + except Exception: + return False + + +def filter_traffic_tracks_to_birth_lists( + current_traffic_data, + sdc_scenario_id, + map_manager, + *, + lane_threshold=5.0, + static_displacement_threshold=5.0, + static_speed_threshold=1.0, + return_stats=False, +): + """ + Filter traffic tracks into controlled (car_birth_info_list) and background lists. + + - controlled (car_birth_info_list): 非 SDC、类型 VEHICLE、至少一帧 valid、在车道内、且非静态 + (位移/速度超过阈值)。用于策略控制或专家回放,spawn 时机为 show_time == round。 + - background (background_vehicles): 同上但在车道内且判定为静态(位移 < 5m、速度 < 1 m/s)。 + 仅作场景占位与观测邻居,spawn 时机为 show_time == round,按 valid 在 step 中移除。 + + Returns (car_birth_info_list, background_vehicles, obj_to_clean) or, if return_stats=True, + (car_birth_info_list, background_vehicles, obj_to_clean, stats_dict). + stats_dict: n_total, n_no_valid, n_off_lane, n_static, n_controlled. + """ + car_birth_info_list = [] + background_vehicles = {} + obj_to_clean = [] + n_total = 0 + n_no_valid = 0 + n_off_lane = 0 + n_static = 0 + + for scenario_id, track in current_traffic_data.items(): + if scenario_id == sdc_scenario_id: + continue + if track["type"] != MetaDriveType.VEHICLE: + continue + + n_total += 1 + obj_to_clean.append(scenario_id) + valid = track["state"]["valid"] + if not valid.any(): + n_no_valid += 1 + continue + + first_show = int(np.argmax(valid)) + last_show = len(valid) - 1 - int(np.argmax(valid[::-1])) + mid_show = (first_show + last_show) // 2 + + start_pos = track["state"]["position"][first_show] + is_valid_track = True + if not is_on_lane(start_pos, map_manager, threshold=lane_threshold): + mid_pos = track["state"]["position"][mid_show] + if not is_on_lane(mid_pos, map_manager, threshold=lane_threshold): + is_valid_track = False + + if not is_valid_track: + n_off_lane += 1 + continue + + positions = track["state"]["position"][valid.astype(bool)] + velocities = track["state"]["velocity"][valid.astype(bool)] + total_displacement = 0.0 + max_speed = 0.0 + if len(positions) > 1: + total_displacement = float(np.linalg.norm(positions[-1] - positions[0])) + max_speed = float(np.max(np.linalg.norm(velocities, axis=1))) + is_static = total_displacement < static_displacement_threshold and max_speed < static_speed_threshold + + if is_static: + n_static += 1 + background_vehicles[scenario_id] = { + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": ( + float(track["state"]["position"][first_show, 0]), + float(track["state"]["position"][first_show, 1]), + ), + "heading": float(track["state"]["heading"][first_show]), + "end": ( + float(track["state"]["position"][last_show, 0]), + float(track["state"]["position"][last_show, 1]), + ), + "scenario_id": scenario_id, + "length": track["state"]["length"][first_show], + "width": track["state"]["width"][first_show], + "valid": valid, + } + continue + + car_birth_info_list.append({ + "id": track["metadata"]["object_id"], + "show_time": first_show, + "begin": ( + float(track["state"]["position"][first_show, 0]), + float(track["state"]["position"][first_show, 1]), + ), + "heading": float(track["state"]["heading"][first_show]), + "end": ( + float(track["state"]["position"][last_show, 0]), + float(track["state"]["position"][last_show, 1]), + ), + "scenario_id": scenario_id, + "length": track["state"]["length"][first_show], + "width": track["state"]["width"][first_show], + }) + + if return_stats: + stats = { + "n_total": n_total, + "n_no_valid": n_no_valid, + "n_off_lane": n_off_lane, + "n_static": n_static, + "n_controlled": len(car_birth_info_list), + } + return car_birth_info_list, background_vehicles, obj_to_clean, stats + return car_birth_info_list, background_vehicles, obj_to_clean + + def set_seed(seed): if seed == -1: seed = np.random.randint(0, 10000) diff --git a/scripts/README.md b/scripts/README.md index 0dcffa1..349bd1f 100644 --- a/scripts/README.md +++ b/scripts/README.md @@ -35,7 +35,7 @@ python scripts/visualize.py replay --data_dir data/exp_filtered --num_scenarios 1 --horizon 500 ``` -- **policy**(BC 或 MAGAIL 训练策略): +- **policy**(BC 或 MAGAIL 训练策略):与专家数据生成/回放一致——同一套车道+静态筛选、且会生成背景车(bg_*),使观测分布与训练集一致,便于在训练集上公平演示。 ```bash python scripts/visualize.py policy --policy_type bc --model_path models/bc/policy_best.pt --data_dir data/exp_filtered --num_scenarios 1 python scripts/visualize.py policy --policy_type magail --model_path models/magail/model_50_actor.pth --num_scenarios 1 --deterministic diff --git a/scripts/visualize.py b/scripts/visualize.py index 7b411db..5a3e532 100644 --- a/scripts/visualize.py +++ b/scripts/visualize.py @@ -177,6 +177,9 @@ def _run_policy(args): continue print(f"Scenario loaded. Controlled agents: {len(obs_dict)}") + if len(obs_dict) == 0: + print(f"Scenario {i} has no controlled agents (all filtered out). Skipping.") + continue step_count = 0 episode_reward = 0.0