From 6c8af8d7a81744cc7afac652e97c4fa679feada1 Mon Sep 17 00:00:00 2001 From: "zeyinzi.jzyz" Date: Wed, 2 Apr 2025 19:27:43 +0800 Subject: [PATCH] update v1.4.1 --- .github/workflows/publish.yml | 4 +- asset/images/icon.png | Bin 14682 -> 0 bytes environment.yaml | 10 - readme.md | 7 - requirements/recommended.txt | 2 +- scepter/modules/annotator/__init__.py | 4 + scepter/modules/annotator/dwpose/__init__.py | 2 + scepter/modules/annotator/dwpose/onnxdet.py | 127 +++++ scepter/modules/annotator/dwpose/onnxpose.py | 362 ++++++++++++++ scepter/modules/annotator/dwpose/util.py | 299 ++++++++++++ scepter/modules/annotator/dwpose/wholebody.py | 80 ++++ scepter/modules/annotator/dwpose_op.py | 203 ++++++++ scepter/modules/annotator/face.py | 63 +++ scepter/modules/annotator/frame_reference.py | 58 +++ scepter/modules/annotator/mask_aug.py | 450 ++++++++++++++++++ scepter/modules/annotator/outpainting.py | 10 +- scepter/modules/annotator/raft.py | 62 +++ scepter/modules/annotator/region_canvas.py | 95 ++++ .../modules/annotator/video_segmentation.py | 153 ++++++ scepter/modules/data/dataset/registry.py | 3 +- scepter/modules/data/sampler/sampler.py | 6 +- scepter/modules/model/base_model.py | 13 +- scepter/modules/model/registry.py | 20 +- scepter/modules/model/utils/basic_utils.py | 6 +- scepter/modules/solver/diffusion_solver.py | 3 +- scepter/modules/utils/ast_utils.py | 2 + scepter/modules/utils/config.py | 75 ++- scepter/modules/utils/model.py | 42 ++ scepter/version.py | 2 +- scepter/workflow/calculator_node.py | 27 +- scepter/workflow/config/scepter_workflow.yaml | 2 +- 31 files changed, 2134 insertions(+), 58 deletions(-) delete mode 100644 asset/images/icon.png delete mode 100644 environment.yaml create mode 100644 scepter/modules/annotator/dwpose/__init__.py create mode 100644 scepter/modules/annotator/dwpose/onnxdet.py create mode 100644 scepter/modules/annotator/dwpose/onnxpose.py create mode 100644 scepter/modules/annotator/dwpose/util.py create mode 100644 scepter/modules/annotator/dwpose/wholebody.py create mode 100644 scepter/modules/annotator/dwpose_op.py create mode 100644 scepter/modules/annotator/face.py create mode 100644 scepter/modules/annotator/frame_reference.py create mode 100644 scepter/modules/annotator/mask_aug.py create mode 100644 scepter/modules/annotator/raft.py create mode 100644 scepter/modules/annotator/region_canvas.py create mode 100644 scepter/modules/annotator/video_segmentation.py diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index 5da0199..bec77e5 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -13,7 +13,7 @@ jobs: name: Publish Custom Node to registry runs-on: ubuntu-latest # if this is a forked repository. Skipping the workflow. - if: github.event.repository.fork == false + if: github.event.repository.fork == false steps: - name: Check out code uses: actions/checkout@v4 @@ -21,4 +21,4 @@ jobs: uses: Comfy-Org/publish-node-action@main with: ## Add your own personal access token to your Github Repository secrets and reference it here. - personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} + personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} \ No newline at end of file diff --git a/asset/images/icon.png b/asset/images/icon.png deleted file mode 100644 index 035f686af9fc2aef150c7d2c8713ed9056dbc80f..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 14682 zcmcJ$Wmr^S)Hr%(sG+;NL<9tB=?*Ccl@39=LuqCJrAts!0ZD@tkY+|H3F%O3KtMn` zXMkb2=l8y!?uYx_FaQ5DbDlGE&faUUwf5R;uiA*vJ3?pw!pXG6et#_!9!a zh{2mpzy}xbM(C<;qz(Y}>EyU4MBslOCp}Xm00_GQ0Eia=Z~;C-Yy!YD2>|%#000WP z0Kn|`w(Wru_=4EcKvxU6{_j)VS(ybsA@$R<330iAoy9$t|FexNL^sLnERv#os-8e|35UfC8cS3*x7|gmNF3@}MxaYXcE-bO zGs>IE;3E*VxX8WLjcC@y-{rhe5HluC8E6Jn<4;go?yFom9I9CSQ`3KF1A6te!VGj7 zP+Gc5a^a-_rnF7glkz5y`q(azGM1Z)&Y@6Y6(&<28(JO#KD`HfqFle9Xrn+X|L_S+ z>pe+PwIy`9%UBvl6d73^4RWJbGwEO1N@X(2uPmhBQl+g|8)Yfaw&If%Hg{WVdIDLKB3jaLmp zv#f05lfVme`2I-sb|I*A%ogTYSI)tzSM%1|8&TMLyakDXw7b;<#_K)2Nb*vDmV*iHm`)OS-O} zllj;)2}CnYLT*eO`~Dz1;(%Km1&Xy~mUM{_qo zuJ=n5#PLMN6;XCJhFy*b+vw<&e{SH=$8hEs3|&BJi*{j)ydP$s3{l?epPkl5qF-%x z$0KNL*7PQy7;C}~-!hrj-)x2#qJ8O}xbH0RRt|k`~_3h)%`o8?&VNgdR|o%Zc|={|jo zwItM;n<&zjh+pej{-Ht9iotkK@SQ&yyo7U4FQX?;VTZBB$(#LT2xsI25goKpN4}~v zf5Cc&;I6NavaeRLD-K^)9C1h@r2LR{D2|=D|P|EBvCqBcGf1*TOC31S)#NY-e7M zGj$JMa-mp|+h^qVYx8B3OiC6v+)m<1ZE5}lmwfyh#aospW6z-sw;ppI?jcj#k|dC# ztpG+e>!k0vhJ-pU9SodDSa!<&RS{?|3x4*B4$3?C2H5;Z8c92>?x!90$AsJ=4MQYr z^p-2=u2enre#vVit9r%|zSeHsx1z*lG1qSiy(1;G8w%sEDDf&&N8=jREXPBQ`Ksr? zRFnZ*89X`c7OkU~)h^n!EbTn->afQ;_I53|2vGjBHadmD91?Sfr%_n*>b zR+1(Pu$3NLlwM*n4`X;d`SryY)BVcHo$I*SOI`3o6fLF_o^#u0n z?twNN;ls~;OTe%Wi34t_0}$A~3R}i_{A|hg@mot}M9@`wO^~Islj|ji%WTg2{<$+r zkYPB^F2q3>bgncN?nL!)KRuGFg?seadnimU*zD3nI{l-C07Fk?{T+Zjv9&v#=NC_z zeD2O2#K|wjekAXNYL1{p3&{r^vzHZY)t`q0ZBt9ws#WO-nC@0vs<@tS@@gb8XK)m$ z%=>n~igNEpSbws6K7W0KI3q0W(xza;r1YdW;NZ{WPUbud2n?i`1K{}AN%-}LUyyIc z?CV@rT%|_BzAIZ)fMbg2WmdPfRPQ%#cE9wZlL*F-zx`hmr^ZfjyRRq4)My-bDLwH| zHq2gmn9>Fl=AttV67v4x^Bkf#k2Zxav$dbGMlwNw@*>?W-zI0WCPu_#tzN#xJ441r zgcL-?1n~s6Um`k~unZrdGws%I_MbIT=vwbm^WTAM`SM!CGrT8;X7nDgGm@-arrARp ztqE>a_8zgSf5|>M+g>{A5xiv7^~+AVpvpu`C+xWNDR3{t%+_(Kp?5 zE`@3$fLKbBHla&e%V!A#qvomHkkmYQ4H=dO3t*XP2vpoi{r#BIa$nL^WMto(Y=K+j7p@@BT+iB-l69j9gncxV)HNnjb!?+1Y_6GD$uo`u zOxH(fqP^^*A79A4R*{qbQZn7C0S=vzq-_j|B297sQ!0aELT_`SzU7iu=0Edr;1pK3 zSXZfey|OtHL6woK`sY~Vh)yGDz#Py@b5W1Yg*;_}+KKHe&srVcx_52E#r^L01Txv| z?ma$jd){7BJ!7?`Z?Iz16=NMj+}HY8*woxB80pG7Z|s=oh+;a8@I_Pb!+V z;x;FL)w8hT%HMmeF(JH|D?Ci_<0~3X=(hQ9wggLm@A@8v4FmaFv$gUb3omO5gAgCf z3))0_+CLeME-6A!jY!q!C{THQCiw?jeP^t+xv774###nDv^wruKz^BkDA8jT{n~t+ z36v<6c)G`XtNlhaqakW=hvLi4&OOb^!X8(aJ0GT+jP>@P`KXy?Pv%N_YNB1)2qBR= z{F07<0wZodU!{wXOGV zaV0ECGnx-hWa=|ctkCq3U9_tf&nHILH%2xyxA+{0ct|{c>MtTBU&zDcf#b=2$k(Ch zKu;da;~avl2zy(YVb;p#WwcGOsA(K+G#-_RBHfcEoE-jfs7*~{Olj!76TAD&Jy*gEo&J9oVLm1WY!qz6B)MTV~xhr#@P36%7ncV zuequi-`SchS9~!uUqDZ+i&#SLd){ap5A!FL$bKEo?qtFWj{^iS7`5$HGdRs3vlkZ0 ztWMXwH)t?i2VEQz?H_LR8xFI{QiiZ(z={p|6v(xN+%E=(p6=<3XPHF&1k}F@6o(O- zKrqQm<%?kUHDyu(tN*dSc$k z6Juug5lk`Oxw64Tb%rNJjK7}dftyPWB&z7Elr^f>P)(0#p^kn!U-qm#R8)(wPN|W=agqZwQikgWP;z< zzp2C_x+k`|voCJ-Y7CGftfUS-mR*L6yfpcl`MMK2F%w!VYGeK7dnFOn^xSR-p}Z40 z@j=ZdYN9>6tTb1yU-c(+!R>z=LBiHC&eJ5PJ69XjTYM+@R}0>A{9?8&jPrT3OY4Hg zFb@yKNncRI9#qA;Y4?*sx^mlOvsQ(yr`q!xvv=uTP1`H3xLdRJr}_Tzd|1)#FAl8S z+EO|c78c|A&ff0#bbzUUiO_E*h~$Imku^bpuN!d+WI|->DR~{A9x?>5z$ z|C=659lX7+az=ddGMHo;&>^ zZu9EWSI9J(##V{)$2T*_gri+A94!yS{MG zEt>g(&wqV+ZVy08VA2WRD0a2hO^i>nP6T~JT=n;<6qD7M+Q-w7yRVEq-i*V~eqW1^ORCSL}O)Bg$*If6He-L{oEvt4d%WKmN< zzK$)?mRoUE2?~MS|15N-AwfTxvpDwe=39ifl`I?!Vy6 zhD3x;YvdVQqekht)l8rE=xW&6way%mKvTdlJ?!={Ex|8RiVPvq+ik+grgM2Kt%RNRv?GCrBLhQqFpMtLJRO=n*Ndpn}7M`{mRBtPB-@DjYzKEkZ(Rn zHCA7Kb0wbhLdhZ{vRFU|A{UKqs>zQitjS+^K56BCMUy{`?>IYYe5AH*bhESSXU%|a z?}95A=R&jBr#t@k@sV*bH{?oqy;`u0eYN^O+e%;iE3G-0arKGLZtjWgZX)9MC;t-y zA{&?S@o^!dP8ff{Zx8r!?$Sp3enZ*5Nz}*&;Xhjx+bmYmX|Ixyitqt3Ye`1YP|;7{ zpAkqQ7w01Jf1d`n5uqrJam+503Bw^Qwc1cVS$B!-&&#Q5+P$B5{=Q_Y9eiRi*I{IN z2a~63vN^jIFquUE(SG8xmC$F*zS098d}zFSafvpw{qyN~_KxK{;&L-1=qph%6txC_ z&dk?mq1KsT9K)v>ZN8XeG>w1ydqb{AVcE#xhVfOC+m#U}uy)Yzn2!VgjbmkSx=E;nS1Yro;-u6QT>Ap3t_i!EltLL+j@8;~CocOHDVg_2iL1 zuZ|H}5s!&9a^KXx^a;cu9rvWt-;N=N*8hIhH%t~yhPA{MRvj*4aMdjPofj=^(e?mG z=c`5)HOiQEl$p2fQg%f*8Q^j-w8FB%K=%3S>XHf3VnUVMNvhV%w{bo)t~@)SeVx~Y z659JZcwJ}*Z}Y3IriS?npw(V5e|D3MQ#)$aJNAleTTk2mC%o^}v#dP6QJu|RP>kh@ zIfv~U5Kj~*^7mMaUGdBIO|*{US9m)gV~!sYPydd}+S_+C9-CO{lnM3JL1>%`a6C0cp5WZA8Mm#DAi8 zXO4S#|Lx8`a^)c7$eGjP&V1zUBC7tBnK1Lw%6+-%Ry*?|R)WV+5$z(m;GIz;%GA#` z%c5-O&;$D@YT(<>*oX4!{KJiY8NA4g+E!mt1v5fPN)_jwJihzP)LXybwAU(OzHf~> zzYc9^aM0q|TvLGfzU&hZZ>S!e%MY@F&`<;xIm?B;GGdrW`KP?RM6pVeHSpKpE7>jp z9sJdS&;^13C7)}U>(#ra;>IKTe}dkiW4!| zh)#mB3U%?%h3oP_{Mi91LA(7Eukqyc2{ppT=MlsvbH|3XACe7}D+6ZVY}m^t%|3Sa_Ex9?a_87@pg~r&)h_8UOwln zIIkm9GdH2lHEhH0%}`5e-ghaMD$UEd5wJv|nH6o-xl$#AVr#Z3IHG+q_QIK4fm}xk zcG0dvy*xeBLXsz%cibFr!tv|sULFHi^Vd-ti1~V?x=Mq6`82Nhjcgy%%hh<5q{Br` zb$I|Z*u*`PzkDU)v;?T+s4YD%s$|^eHkbw>fir9JK7t>An8GHfa(^=s(M2vz#j!JDnuXkMMV~PD-L2npZNMT(|7Dcc!kKk8-|9F<0#qArL*sX)WLS0nHd1YuNEz7L=AL3uF?8U?mS$i^;cjPof&qrt+a%a?ahB z&bdp;TmnMOj6AR>pT^lzR`%)pJ7|4Rw36(3f5p#?jwgv09xuddrtOBlIf=T#+mM7f zyTCx!jWlc2Kc>hcW7_1yeAdt0(JgKltLHYFzsnfs5zh+_bCRhPnbfNDuv`gTA^P%o zUEYTmDrhXAuI%>(6fak6i(Wo^V!Oj#1*^DHZJBKo-6w}T8&RrjEu*$s0hf_h)n>{a zRwBWZZtE<~_N|v4x|mS!&Y9zE@1c3Dq9o^C#$N_61{kk=v4x=vbGe7@9jl5eYVGVK zOI!&fQFx4In(7)G?^R?ZUfTY6W4<^tU9S~4NQi$ubF=)nd3vF>dn&q=pBw?N3gQ15 zO}q!sdf*q#iOh*?zfIa!$CtfZZ@A`W^sMkW*hv^B_R!hZ4oVdGVQ%`u`q zm}vghci;Fp>Q33Mv*aDV47+*HU|UM3ok4Nvy@p!5?N-JjqozJ|29+rA9Vwx!vJ2l2 z)vH!Mt*c18-91jY)v#Sy>iCL2|DsS(F4?8?VE+9|H&RSH;EVdPP8_Udi1dc})V-R@ zl_t!ad(*gQR-?8Tzo(AH^YE3$L*^+#4iJaXw5!yn>&^{=Bc*=+TpuY3_dMoFB$G-!rd;ky7$j4K^D(ei zfBZ=z)>|nkZG-sLS6`Y?(E6%op*?81-^&+!d|o5i``uDpMZ)WVc6NZ+2i|9>f>$e= zZLY=nAkB}vyA-N79(jH3Llc6aY+NQ^uef4hI!u0PMUbU38s6olV)Dg83yJMsGzELa z)U(#bsM^06>Lq=mbB@KaU9|b<6I_x>)!~M^E;FJp&g-8}A@w6qlikcw(x>umiE;J7 zJbwnsu`B*cuA7AWic{D*+lsqret-QWcH)){XHk2;PyW39=aY^>qa4fC(e$+2Kw$Ku zw9n4(!eB@rvtMEOordl(rRACl&w+2^dlzTUMP3%6yB+7pA8VGD)BgPYbEc_-;Tn)d znzHqW2T6x}z9}?6JS=Cn37ag6y=@=+v?Ft+Y#666N)yA#;%ybtx-IoJ2tEh8Y4AFV}9cg+TWvM7U(y%d*?I?M7D z-TXAGN3@3l=^rf023{WDdwGe+%o(%=Wv08sLWwDbX~(P@ z{&Zq;rCe{Pl}0Q^_{^(romCP#_*R#0h%fqWy)>&nl#8bl2X0f#EUX{Hv}Tqdh`^IJ z{ZRHg3f-Y-p~&_d1iJwPAF}^0aodLb&yM_gFvfdr*pbJFh&z4sXt)Rotd~uMa`rRt z1{XlWJ?X=uv`+i_^+A5J zLNHyr`%S61A{EZy$Fv`6{&c<(Sp`6nYH=VYN$l@T2SaOk1DN2>rtf2 z)rpkZvEs*kWQFYe4RKreVQ&ulWG%`j|0NwQ0NnN4^BTT5d?{&BTOj?uIPhhu6e9kr zzgQpgFJ$37*7NqxfPl4sX6^r64M zehswXvQK)IkJj147TqJs&SNNanX{J2#~4E%^zxCAMABO&b1#LR!gF5QtO`ZO6;^1W z{29_l59rf0`Bn~oXYYnHi&}c7A{jU9vSMW8E)Le7po$iP%z{aMR*K}M&#lU`H)eRU zA!z0+SEiuJhoVQK%l2jmMX0-y9Nxx(Wp1sj;pXqqtyv)&ut6HAu;~y98RYxoM)jQ7 z;(&5lkDPyH-m;l6Nj1Zvmg8*pZqxPocD$VhOVGIjU->44hh)SSqLkZZp4?|oAVGnR zkTSkYc_9AwP;j}Xcwf^XpdVACpvpjJa?>cUdC5-lTYR5#w3bGzy6T&%(|0d4y=^0A zE3Gl_oKumfH}xTYp%iW-+%s=LUc@sW5wWK25swJUgu+nOx0vp#ZWfHuEj&Mdgl5PT zYMgP#TZSFcM!#5HmQTC<;76weoK5!zdoMl+pCN#U)ZbN>%T^M$gk%tG+*+1hb_!nH zjMySu_%QqK@%^M)3SDz(j>}GzO)uZeVdeWl3kp|*$1j%kcV}>UX3u19K>L=>n&3*W zc%FFR-&qGw6}p26g3tgE+zz2nPt}W6UFpNHn|-O`I}t0J$7Log6OdH? zKjm!BDVIN%9cFbwO`Q0ar*Ucbh%V_!KM$_#Fg$wVMQh~wWjy<+oj)jP*jq<1}Z znU7>Rlt2nuR8&FuURMh43__E5O}fcF{+F$QkbOZxpwhbzRrTO*2FA zBZfobaz8g(^WOH`pTISq3ndnIwe^yUMI)}%B(soViV_>2q@6}_p%P*fe%e?o{IkJo za-gWa=-=q*E@_B&)aseYAOF>h*=g%lapEdqO^qt|N@24o;~(e24?)-!!m}W>jT<_O zwd1DjC~(n)(eEIlADO(Fg?*6ZFSobp=RMD6OhW(tX4Pgy z-*on(kuU)_vB>BmXnYw7v>}bJO5HDh^KwM@Wua~_-xCL}1MaTHNoV{g$HvxJ*tMi!YJ|Wd(oT|wDCywk2 zF#iaIEtc7#-tkHT-KE@p$hKC+X|e3_XFb)6>0`7OUbaRRBYRo!m!~glkA#mP!s{l%8fU$`NH- zvafikF$x-(wFINDJSwIkW>4i6nwti*-Emy!2?$eTQ=sv3Q=as#mg9Wa9CmHv8vpOh z<7Bz+u|$*`%)=H-ObD3YgTN8~?r}5( zbN-wY+EL;7s@v*|k`b%w(>&zpVh3->(%)yN)3g1}dyCAf-E44wCuPKyR;}XhLqa(O z>{XuF&54-U(c~^)}FDR(7lhe4&HX}cETklWAb6Uj+{0EI)eqAM1B7q|V>KyL?iS?;M-} zu@{?|Jto7MWsZP!f!8R;HjkqOjdbzu@prk?NvbEiT`XAz_h~LYad=T=h?06F_>GJ( zu#Lo}oVnoIJT&CM1T}AGR0I#HxjHy+M6_2+)y!8*o!r>u_Aa-t#l?g*laI1YZk}Kz zIj$J^q4)>!XtVksl;>Qff+?zZ^K_NemAERqRJsHU6HKCZn~oU0{R*rpI{m-<>y_yF zGip%4F2Z6^WNIQj7oC@qF48M@a<=k;7pDm=|kWR6<_=XGV`wt0Zd2?KZxE5CL;}K%S9;)V+4Ps!rh{Z!KF-u_<8d}7y#RJpCFe4xQC85LbdxXW^mP=9G4~v} z>P}*g3DG&ZJ*x0g?`usLLAy}PRY}vyg^fquvY9VeZ#!drS!6Rmiz=ca=%H&_Zmd9N z%2D`Y#8PRy18K$*?e;%Ht5MCsCMA_N$<3oi#ub$L#MHR>=8CS@(Jia7=-@sNj`0TkFYi!+X z^4Kw%zH}n6&w4Lb$EEhFT@d(Ly>eNCkH;|sZwBD3qpE)|->2Z!H+TI+WhM*qA zTwc$lV`iW4O7!iEvfbqz7(JZE@m2q2X6^o>ML^B7MO;Q$#@5@sUMDCObKBn2k*GKU zR2E?|Ei^1z_Xs~ltJ>g8oo-~gW{x6C&-{?AvlfnXma48kR4uZ@9q|d` zqywrM8WpXVd!>3(QAI0qGyq!=YkmX-pEffauk*+Y)>Z{p&>{GA!O$meu=U2Qtf3wu-P>~>}7R_)^=_F#-41P8rOkP+adXugp?iukC#Q0tyM#=+!V;;Y? zqi&WslRPF$#1yY%9`S3fw=u8^Xr#wMkOXZaeN?{kE2miaR$^VDGxFjK|%4v`Zrj-{He5TR|BO$_Z8*GNDn z#2Fz7$E2jAelHc*-q8cCE8(Ac1G+Hzi#mCBIfH-O#|B3a>h*%b)L_T{edZs2C$;M6 z*leCVt}R;kqa8x#Qe%3Z>2tJWy1I(%jSzfgt&<_{l#jSKd~tvz$0B(ItHt{K)kYu4 z5tG;NzSz7!!F-X1%o)|{{!6Clul9|&|y=xiUA7~Iuid>r|_bV zW4me`!R?rqc8#zQgqJ`I@wLIJ=e|j}B#N}RO4wgn$K{EtiU1ann`}u>&g@vz@Z^&m zzc%(1p#0GIwc{(t3NX{^jLO;@N(DpGK|dDo8lY85$PjxN4%SCO{BGaSXBt*3%K`17 zri@hvVO7B6+YFuyRN)o5DQL~O(psqb#q)P5mZ2YPuRFnFx(Adt_bw;~1ygs4DKZ?T zqg272%YIq6%NCijviVaTEoU$^Uc=a_Y_`+4f8?~B-ZeW3?vjrJsa=Fssz_*n-5x9i ztUO?qAZQT1oBIpJ3|hV}W0>(T*PwmiG?#T44%(c6gY2@HQX6^h@)?`O+?~h3T*jFj zjtlg=p%%bNU~n}_UyHt0_Vc7vNF(j;Z!81^Lh3S>k#lp3I!nyLFI2QQmJuOa+{uoN zr_jBoj9_fYyn_;0mh*lVWiSLqCXjTz^Owk2727*MT7%A%U_*6=Bl zMBic)Ounhp@U{mig~?{O`lts6v_^M|Ho$`=Tp&>)L=uPm zioDAt)t*uQqYm_FdOdI;5c$pUtXOu$`rfhiLs8Qs1Aj*L$G}mUMJl@7xR=i<$r~TV znl=O7crwG9fm;14eBZ_b@*|c{E()aGJg(FL0RaW4BuTpm&(%=ww+d-)5jP?9p}xo$ z74{NWqi1M8bOn9RQOE|t5ZFUv^K!#33st2Sp>RKai!nrXbNdU(%-5h<9Q9B%fXPKY zJs$Mc#_?RWL^YMvb#(GHxyh_DZr507ADJ4}T*F~_OT+n6+`vyHnzNIc(sC(53f|q}+Q38*YnOjlSgoMmjOO=m%n#v6N7k zvGEjcFg);Jc{_JN$rY@iUf_YgtQmsKxcI;7$`x!EgE%S!+??ZBDgW3Ps>C+{3?Sb8 z%o>dDAR%3`?)e---Q(%8UE<|Qj6ja5#D6R3E>#ON3xIne%K!yMi!cJ# z`KSuGAJ(@Is;|}PA_9`u^wGxNw^1PLWDl$j8=e+Z=*sPE@QVE`UULGV_w>m@Q9^RFrgecyi9^(S?lfE$6gkW-Ce~KC z?nr2t{6X>M(N6J@{ioMn^e8s?D){yo734s6KbrM7#vF`-`FsncAjy6bM$|?;t!z2? zx(u*b9R*hm1(6S7*Gu=>QIz4O45&H6&yIf)AgC8}7QKzt=FI}24ynjc5SqmDt!2q# z)%)ZfNwl$Yi_@{>%752UrD|^|bT5f)`%KHq2vKa;_)*C`+lzM)_Z9-}o9(*^q16L(Lz?fLFtqxvm}3(|^Bn|Mws{|V(8s$cPg?oA9IlhMPh!n2mAgH09$gk$E%}!(7s<_ zO^p_xdODZ#AFx}whUFOZn2lX{CvLC)m4CBA-h@Q8Itwqz*J+8ng9%TCLD`H4qVD$R}c1#xKMif8$QTaS(vji0aDxA?8d)gbW~%(1tN zdM6$?Dndd|n~uW!(l~<)URM;^U*eCG)h4h_{?|>WH`g=#QqlE&zV(gb|JzI_J#pbi zIbroUY<616i14-PP4ASYmMaDKlRUovI6c@{P!CAii6LEwbANchzYpV$a58*J3z4{N zG;TVdT#ORh!i)4Lp!OEqr};pf=_+hP{S|&v-IyPK1ve$*m z;M8VQB#jXKZ)^4^9#Olk7l02v>X@)`{%2()Db{^-*CgaYln^VyhmUvq;H5HA$%3gR zbO>{o4*&hSmC{97+`jEH?B+W&1C%hc4)V}l&6>O*VMRGJW)gK~IB6ngAQsn6lx1A2|z zu1HACLtSLfN@bJr3~oT6$gmmdlK#oUBNgrFo}{LX00T@Hk(zNSZQ3wRtZx5 z=+=0@_p#9Bd6qZ)jn}cbE!)-K6ZZdxFrv*RIar$kyWCRw@Hram|iF%`j&jnI*|H`(3m8+D-tjPixcOT(-YKLd}^ z$m=U#S>L6UED76i(tchp%rs6WeCo?RQ1plJG`Rq7#rN5$aPQ~TMAW4S+I7Dxl5!EB z&(_otFmNts(hN@g4s|psYHm}e&84X8q5I4RU9=Y8w~N#5pQg5H{@4Fuz&He|Kv6OE zV)9tr=GlA=qIlU)iHK;J*cDK#LtzQ*Ajtf=^FxPri1s zeM~!+5AGFTf_+oNe0xp*YeP=g0sss6t05=F=UaEu{`7kB-M4NxuWkHK&B#Yn6A!K} zU5De6-zoYhzkW-hw=QrN_w?#i&IZ0%s%~GT0uraLFSH##2p0*iNU7f$p2RdZ{MfxE}{S59ojd8 RzJVP8eQjf{8V!e-{|7qWVMqV~ diff --git a/environment.yaml b/environment.yaml deleted file mode 100644 index 009e28d..0000000 --- a/environment.yaml +++ /dev/null @@ -1,10 +0,0 @@ -name: scepter -channels: - - defaults -dependencies: - - python==3.8 - - pip>=20.3 - - numpy>=1.23.1 - - pip: - - -r requirements/recommended.txt - - -r requirements.txt diff --git a/readme.md b/readme.md index c92514c..1a515db 100644 --- a/readme.md +++ b/readme.md @@ -135,13 +135,6 @@ SCEPTER offers 3 core components: ## 🛠️ Installation -- Create new environment with `conda` command: - -```shell -conda env create -f environment.yaml -conda activate scepter -``` - - Install with `pip` command: We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip: diff --git a/requirements/recommended.txt b/requirements/recommended.txt index faed02a..8eb5626 100644 --- a/requirements/recommended.txt +++ b/requirements/recommended.txt @@ -1,5 +1,5 @@ git+https://github.com/cocodataset/panopticapi.git torch==2.4.1 -torchvision==.19.1 +torchvision==0.19.1 flash-attn==2.5.8 xformers==0.0.28 \ No newline at end of file diff --git a/scepter/modules/annotator/__init__.py b/scepter/modules/annotator/__init__.py index 8262ea5..6bc81d3 100644 --- a/scepter/modules/annotator/__init__.py +++ b/scepter/modules/annotator/__init__.py @@ -25,6 +25,8 @@ if TYPE_CHECKING: from scepter.modules.annotator.segmentation import ESAMAnnotator from scepter.modules.annotator.sketch import SketchAnnotator from scepter.modules.annotator.lama import LamaAnnotator + from scepter.modules.annotator.mask_aug import MaskAugAnnotator, MaskDrawAnnotator, MaskLayoutAnnotator + from scepter.modules.annotator.raft import RAFTAnnotator, RAFTVisAnnotator else: _import_structure = { 'base_annotator': ['GeneralAnnotator'], @@ -48,6 +50,8 @@ else: 'segmentation': ['ESAMAnnotator'], 'sketch': ['SketchAnnotator'], 'lama': ['LamaAnnotator'], + 'mask_aug': ['MaskAugAnnotator', 'MaskDrawAnnotator', 'MaskLayoutAnnotator'], + 'raft': ['RAFTAnnotator', 'RAFTVisAnnotator'], } import sys diff --git a/scepter/modules/annotator/dwpose/__init__.py b/scepter/modules/annotator/dwpose/__init__.py new file mode 100644 index 0000000..cc26a06 --- /dev/null +++ b/scepter/modules/annotator/dwpose/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/modules/annotator/dwpose/onnxdet.py b/scepter/modules/annotator/dwpose/onnxdet.py new file mode 100644 index 0000000..0bcebce --- /dev/null +++ b/scepter/modules/annotator/dwpose/onnxdet.py @@ -0,0 +1,127 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import cv2 +import numpy as np + +import onnxruntime + +def nms(boxes, scores, nms_thr): + """Single class NMS implemented in Numpy.""" + x1 = boxes[:, 0] + y1 = boxes[:, 1] + x2 = boxes[:, 2] + y2 = boxes[:, 3] + + areas = (x2 - x1 + 1) * (y2 - y1 + 1) + order = scores.argsort()[::-1] + + keep = [] + while order.size > 0: + i = order[0] + keep.append(i) + xx1 = np.maximum(x1[i], x1[order[1:]]) + yy1 = np.maximum(y1[i], y1[order[1:]]) + xx2 = np.minimum(x2[i], x2[order[1:]]) + yy2 = np.minimum(y2[i], y2[order[1:]]) + + w = np.maximum(0.0, xx2 - xx1 + 1) + h = np.maximum(0.0, yy2 - yy1 + 1) + inter = w * h + ovr = inter / (areas[i] + areas[order[1:]] - inter) + + inds = np.where(ovr <= nms_thr)[0] + order = order[inds + 1] + + return keep + +def multiclass_nms(boxes, scores, nms_thr, score_thr): + """Multiclass NMS implemented in Numpy. Class-aware version.""" + final_dets = [] + num_classes = scores.shape[1] + for cls_ind in range(num_classes): + cls_scores = scores[:, cls_ind] + valid_score_mask = cls_scores > score_thr + if valid_score_mask.sum() == 0: + continue + else: + valid_scores = cls_scores[valid_score_mask] + valid_boxes = boxes[valid_score_mask] + keep = nms(valid_boxes, valid_scores, nms_thr) + if len(keep) > 0: + cls_inds = np.ones((len(keep), 1)) * cls_ind + dets = np.concatenate( + [valid_boxes[keep], valid_scores[keep, None], cls_inds], 1 + ) + final_dets.append(dets) + if len(final_dets) == 0: + return None + return np.concatenate(final_dets, 0) + +def demo_postprocess(outputs, img_size, p6=False): + grids = [] + expanded_strides = [] + strides = [8, 16, 32] if not p6 else [8, 16, 32, 64] + + hsizes = [img_size[0] // stride for stride in strides] + wsizes = [img_size[1] // stride for stride in strides] + + for hsize, wsize, stride in zip(hsizes, wsizes, strides): + xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize)) + grid = np.stack((xv, yv), 2).reshape(1, -1, 2) + grids.append(grid) + shape = grid.shape[:2] + expanded_strides.append(np.full((*shape, 1), stride)) + + grids = np.concatenate(grids, 1) + expanded_strides = np.concatenate(expanded_strides, 1) + outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides + outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides + + return outputs + +def preprocess(img, input_size, swap=(2, 0, 1)): + if len(img.shape) == 3: + padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114 + else: + padded_img = np.ones(input_size, dtype=np.uint8) * 114 + + r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1]) + resized_img = cv2.resize( + img, + (int(img.shape[1] * r), int(img.shape[0] * r)), + interpolation=cv2.INTER_LINEAR, + ).astype(np.uint8) + padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img + + padded_img = padded_img.transpose(swap) + padded_img = np.ascontiguousarray(padded_img, dtype=np.float32) + return padded_img, r + +def inference_detector(session, oriImg): + input_shape = (640,640) + img, ratio = preprocess(oriImg, input_shape) + + ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]} + output = session.run(None, ort_inputs) + predictions = demo_postprocess(output[0], input_shape)[0] + + boxes = predictions[:, :4] + scores = predictions[:, 4:5] * predictions[:, 5:] + + boxes_xyxy = np.ones_like(boxes) + boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2. + boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2. + boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2. + boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2. + boxes_xyxy /= ratio + dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1) + if dets is not None: + final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5] + isscore = final_scores>0.3 + iscat = final_cls_inds == 0 + isbbox = [ i and j for (i, j) in zip(isscore, iscat)] + final_boxes = final_boxes[isbbox] + else: + final_boxes = np.array([]) + + return final_boxes diff --git a/scepter/modules/annotator/dwpose/onnxpose.py b/scepter/modules/annotator/dwpose/onnxpose.py new file mode 100644 index 0000000..16316ca --- /dev/null +++ b/scepter/modules/annotator/dwpose/onnxpose.py @@ -0,0 +1,362 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from typing import List, Tuple + +import cv2 +import numpy as np +import onnxruntime as ort + +def preprocess( + img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256) +) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: + """Do preprocessing for RTMPose model inference. + + Args: + img (np.ndarray): Input image in shape. + input_size (tuple): Input image size in shape (w, h). + + Returns: + tuple: + - resized_img (np.ndarray): Preprocessed image. + - center (np.ndarray): Center of image. + - scale (np.ndarray): Scale of image. + """ + # get shape of image + img_shape = img.shape[:2] + out_img, out_center, out_scale = [], [], [] + if len(out_bbox) == 0: + out_bbox = [[0, 0, img_shape[1], img_shape[0]]] + for i in range(len(out_bbox)): + x0 = out_bbox[i][0] + y0 = out_bbox[i][1] + x1 = out_bbox[i][2] + y1 = out_bbox[i][3] + bbox = np.array([x0, y0, x1, y1]) + + # get center and scale + center, scale = bbox_xyxy2cs(bbox, padding=1.25) + + # do affine transformation + resized_img, scale = top_down_affine(input_size, scale, center, img) + + # normalize image + mean = np.array([123.675, 116.28, 103.53]) + std = np.array([58.395, 57.12, 57.375]) + resized_img = (resized_img - mean) / std + + out_img.append(resized_img) + out_center.append(center) + out_scale.append(scale) + + return out_img, out_center, out_scale + + +def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray: + """Inference RTMPose model. + + Args: + sess (ort.InferenceSession): ONNXRuntime session. + img (np.ndarray): Input image in shape. + + Returns: + outputs (np.ndarray): Output of RTMPose model. + """ + all_out = [] + # build input + for i in range(len(img)): + input = [img[i].transpose(2, 0, 1)] + + # build output + sess_input = {sess.get_inputs()[0].name: input} + sess_output = [] + for out in sess.get_outputs(): + sess_output.append(out.name) + + # run model + outputs = sess.run(sess_output, sess_input) + all_out.append(outputs) + + return all_out + + +def postprocess(outputs: List[np.ndarray], + model_input_size: Tuple[int, int], + center: Tuple[int, int], + scale: Tuple[int, int], + simcc_split_ratio: float = 2.0 + ) -> Tuple[np.ndarray, np.ndarray]: + """Postprocess for RTMPose model output. + + Args: + outputs (np.ndarray): Output of RTMPose model. + model_input_size (tuple): RTMPose model Input image size. + center (tuple): Center of bbox in shape (x, y). + scale (tuple): Scale of bbox in shape (w, h). + simcc_split_ratio (float): Split ratio of simcc. + + Returns: + tuple: + - keypoints (np.ndarray): Rescaled keypoints. + - scores (np.ndarray): Model predict scores. + """ + all_key = [] + all_score = [] + for i in range(len(outputs)): + # use simcc to decode + simcc_x, simcc_y = outputs[i] + keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio) + + # rescale keypoints + keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2 + all_key.append(keypoints[0]) + all_score.append(scores[0]) + + return np.array(all_key), np.array(all_score) + + +def bbox_xyxy2cs(bbox: np.ndarray, + padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]: + """Transform the bbox format from (x,y,w,h) into (center, scale) + + Args: + bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted + as (left, top, right, bottom) + padding (float): BBox padding factor that will be multilied to scale. + Default: 1.0 + + Returns: + tuple: A tuple containing center and scale. + - np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or + (n, 2) + - np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or + (n, 2) + """ + # convert single bbox from (4, ) to (1, 4) + dim = bbox.ndim + if dim == 1: + bbox = bbox[None, :] + + # get bbox center and scale + x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3]) + center = np.hstack([x1 + x2, y1 + y2]) * 0.5 + scale = np.hstack([x2 - x1, y2 - y1]) * padding + + if dim == 1: + center = center[0] + scale = scale[0] + + return center, scale + + +def _fix_aspect_ratio(bbox_scale: np.ndarray, + aspect_ratio: float) -> np.ndarray: + """Extend the scale to match the given aspect ratio. + + Args: + scale (np.ndarray): The image scale (w, h) in shape (2, ) + aspect_ratio (float): The ratio of ``w/h`` + + Returns: + np.ndarray: The reshaped image scale in (2, ) + """ + w, h = np.hsplit(bbox_scale, [1]) + bbox_scale = np.where(w > h * aspect_ratio, + np.hstack([w, w / aspect_ratio]), + np.hstack([h * aspect_ratio, h])) + return bbox_scale + + +def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray: + """Rotate a point by an angle. + + Args: + pt (np.ndarray): 2D point coordinates (x, y) in shape (2, ) + angle_rad (float): rotation angle in radian + + Returns: + np.ndarray: Rotated point in shape (2, ) + """ + sn, cs = np.sin(angle_rad), np.cos(angle_rad) + rot_mat = np.array([[cs, -sn], [sn, cs]]) + return rot_mat @ pt + + +def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray: + """To calculate the affine matrix, three pairs of points are required. This + function is used to get the 3rd point, given 2D points a & b. + + The 3rd point is defined by rotating vector `a - b` by 90 degrees + anticlockwise, using b as the rotation center. + + Args: + a (np.ndarray): The 1st point (x,y) in shape (2, ) + b (np.ndarray): The 2nd point (x,y) in shape (2, ) + + Returns: + np.ndarray: The 3rd point. + """ + direction = a - b + c = b + np.r_[-direction[1], direction[0]] + return c + + +def get_warp_matrix(center: np.ndarray, + scale: np.ndarray, + rot: float, + output_size: Tuple[int, int], + shift: Tuple[float, float] = (0., 0.), + inv: bool = False) -> np.ndarray: + """Calculate the affine transformation matrix that can warp the bbox area + in the input image to the output size. + + Args: + center (np.ndarray[2, ]): Center of the bounding box (x, y). + scale (np.ndarray[2, ]): Scale of the bounding box + wrt [width, height]. + rot (float): Rotation angle (degree). + output_size (np.ndarray[2, ] | list(2,)): Size of the + destination heatmaps. + shift (0-100%): Shift translation ratio wrt the width/height. + Default (0., 0.). + inv (bool): Option to inverse the affine transform direction. + (inv=False: src->dst or inv=True: dst->src) + + Returns: + np.ndarray: A 2x3 transformation matrix + """ + shift = np.array(shift) + src_w = scale[0] + dst_w = output_size[0] + dst_h = output_size[1] + + # compute transformation matrix + rot_rad = np.deg2rad(rot) + src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad) + dst_dir = np.array([0., dst_w * -0.5]) + + # get four corners of the src rectangle in the original image + src = np.zeros((3, 2), dtype=np.float32) + src[0, :] = center + scale * shift + src[1, :] = center + src_dir + scale * shift + src[2, :] = _get_3rd_point(src[0, :], src[1, :]) + + # get four corners of the dst rectangle in the input image + dst = np.zeros((3, 2), dtype=np.float32) + dst[0, :] = [dst_w * 0.5, dst_h * 0.5] + dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir + dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :]) + + if inv: + warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src)) + else: + warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst)) + + return warp_mat + + +def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict, + img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + """Get the bbox image as the model input by affine transform. + + Args: + input_size (dict): The input size of the model. + bbox_scale (dict): The bbox scale of the img. + bbox_center (dict): The bbox center of the img. + img (np.ndarray): The original image. + + Returns: + tuple: A tuple containing center and scale. + - np.ndarray[float32]: img after affine transform. + - np.ndarray[float32]: bbox scale after affine transform. + """ + w, h = input_size + warp_size = (int(w), int(h)) + + # reshape bbox to fixed aspect ratio + bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h) + + # get the affine matrix + center = bbox_center + scale = bbox_scale + rot = 0 + warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h)) + + # do affine transform + img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR) + + return img, bbox_scale + + +def get_simcc_maximum(simcc_x: np.ndarray, + simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + """Get maximum response location and value from simcc representations. + + Note: + instance number: N + num_keypoints: K + heatmap height: H + heatmap width: W + + Args: + simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx) + simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy) + + Returns: + tuple: + - locs (np.ndarray): locations of maximum heatmap responses in shape + (K, 2) or (N, K, 2) + - vals (np.ndarray): values of maximum heatmap responses in shape + (K,) or (N, K) + """ + N, K, Wx = simcc_x.shape + simcc_x = simcc_x.reshape(N * K, -1) + simcc_y = simcc_y.reshape(N * K, -1) + + # get maximum value locations + x_locs = np.argmax(simcc_x, axis=1) + y_locs = np.argmax(simcc_y, axis=1) + locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32) + max_val_x = np.amax(simcc_x, axis=1) + max_val_y = np.amax(simcc_y, axis=1) + + # get maximum value across x and y axis + mask = max_val_x > max_val_y + max_val_x[mask] = max_val_y[mask] + vals = max_val_x + locs[vals <= 0.] = -1 + + # reshape + locs = locs.reshape(N, K, 2) + vals = vals.reshape(N, K) + + return locs, vals + + +def decode(simcc_x: np.ndarray, simcc_y: np.ndarray, + simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]: + """Modulate simcc distribution with Gaussian. + + Args: + simcc_x (np.ndarray[K, Wx]): model predicted simcc in x. + simcc_y (np.ndarray[K, Wy]): model predicted simcc in y. + simcc_split_ratio (int): The split ratio of simcc. + + Returns: + tuple: A tuple containing center and scale. + - np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2) + - np.ndarray[float32]: scores in shape (K,) or (n, K) + """ + keypoints, scores = get_simcc_maximum(simcc_x, simcc_y) + keypoints /= simcc_split_ratio + + return keypoints, scores + + +def inference_pose(session, out_bbox, oriImg): + h, w = session.get_inputs()[0].shape[2:] + model_input_size = (w, h) + resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size) + outputs = inference(session, resized_img) + keypoints, scores = postprocess(outputs, model_input_size, center, scale) + + return keypoints, scores \ No newline at end of file diff --git a/scepter/modules/annotator/dwpose/util.py b/scepter/modules/annotator/dwpose/util.py new file mode 100644 index 0000000..232de86 --- /dev/null +++ b/scepter/modules/annotator/dwpose/util.py @@ -0,0 +1,299 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import math +import numpy as np +import matplotlib +import cv2 + + +eps = 0.01 + + +def smart_resize(x, s): + Ht, Wt = s + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4) + else: + return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2) + + +def smart_resize_k(x, fx, fy): + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + Ht, Wt = Ho * fy, Wo * fx + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4) + else: + return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2) + + +def padRightDownCorner(img, stride, padValue): + h = img.shape[0] + w = img.shape[1] + + pad = 4 * [None] + pad[0] = 0 # up + pad[1] = 0 # left + pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down + pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right + + img_padded = img + pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1)) + img_padded = np.concatenate((pad_up, img_padded), axis=0) + pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1)) + img_padded = np.concatenate((pad_left, img_padded), axis=1) + pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1)) + img_padded = np.concatenate((img_padded, pad_down), axis=0) + pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1)) + img_padded = np.concatenate((img_padded, pad_right), axis=1) + + return img_padded, pad + + +def transfer(model, model_weights): + transfered_model_weights = {} + for weights_name in model.state_dict().keys(): + transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])] + return transfered_model_weights + + +def draw_bodypose(canvas, candidate, subset): + H, W, C = canvas.shape + candidate = np.array(candidate) + subset = np.array(subset) + + stickwidth = 4 + + limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \ + [10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \ + [1, 16], [16, 18], [3, 17], [6, 18]] + + colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \ + [0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \ + [170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]] + + for i in range(17): + for n in range(len(subset)): + index = subset[n][np.array(limbSeq[i]) - 1] + if -1 in index: + continue + Y = candidate[index.astype(int), 0] * float(W) + X = candidate[index.astype(int), 1] * float(H) + mX = np.mean(X) + mY = np.mean(Y) + length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 + angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) + polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1) + cv2.fillConvexPoly(canvas, polygon, colors[i]) + + canvas = (canvas * 0.6).astype(np.uint8) + + for i in range(18): + for n in range(len(subset)): + index = int(subset[n][i]) + if index == -1: + continue + x, y = candidate[index][0:2] + x = int(x * W) + y = int(y * H) + cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1) + + return canvas + + +def draw_handpose(canvas, all_hand_peaks): + H, W, C = canvas.shape + + edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \ + [10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]] + + for peaks in all_hand_peaks: + peaks = np.array(peaks) + + for ie, e in enumerate(edges): + x1, y1 = peaks[e[0]] + x2, y2 = peaks[e[1]] + x1 = int(x1 * W) + y1 = int(y1 * H) + x2 = int(x2 * W) + y2 = int(y2 * H) + if x1 > eps and y1 > eps and x2 > eps and y2 > eps: + cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2) + + for i, keyponit in enumerate(peaks): + x, y = keyponit + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1) + return canvas + + +def draw_facepose(canvas, all_lmks): + H, W, C = canvas.shape + for lmks in all_lmks: + lmks = np.array(lmks) + for lmk in lmks: + x, y = lmk + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1) + return canvas + + +# detect hand according to body pose keypoints +# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp +def handDetect(candidate, subset, oriImg): + # right hand: wrist 4, elbow 3, shoulder 2 + # left hand: wrist 7, elbow 6, shoulder 5 + ratioWristElbow = 0.33 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + # if any of three not detected + has_left = np.sum(person[[5, 6, 7]] == -1) == 0 + has_right = np.sum(person[[2, 3, 4]] == -1) == 0 + if not (has_left or has_right): + continue + hands = [] + #left hand + if has_left: + left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]] + x1, y1 = candidate[left_shoulder_index][:2] + x2, y2 = candidate[left_elbow_index][:2] + x3, y3 = candidate[left_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, True]) + # right hand + if has_right: + right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]] + x1, y1 = candidate[right_shoulder_index][:2] + x2, y2 = candidate[right_elbow_index][:2] + x3, y3 = candidate[right_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, False]) + + for x1, y1, x2, y2, x3, y3, is_left in hands: + # pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox + # handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]); + # handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]); + # const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow); + # const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder); + # handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder); + x = x3 + ratioWristElbow * (x3 - x2) + y = y3 + ratioWristElbow * (y3 - y2) + distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2) + distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2) + width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder) + # x-y refers to the center --> offset to topLeft point + # handRectangle.x -= handRectangle.width / 2.f; + # handRectangle.y -= handRectangle.height / 2.f; + x -= width / 2 + y -= width / 2 # width = height + # overflow the image + if x < 0: x = 0 + if y < 0: y = 0 + width1 = width + width2 = width + if x + width > image_width: width1 = image_width - x + if y + width > image_height: width2 = image_height - y + width = min(width1, width2) + # the max hand box value is 20 pixels + if width >= 20: + detect_result.append([int(x), int(y), int(width), is_left]) + + ''' + return value: [[x, y, w, True if left hand else False]]. + width=height since the network require squared input. + x, y is the coordinate of top left + ''' + return detect_result + + +# Written by Lvmin +def faceDetect(candidate, subset, oriImg): + # left right eye ear 14 15 16 17 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + has_head = person[0] > -1 + if not has_head: + continue + + has_left_eye = person[14] > -1 + has_right_eye = person[15] > -1 + has_left_ear = person[16] > -1 + has_right_ear = person[17] > -1 + + if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear): + continue + + head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]] + + width = 0.0 + x0, y0 = candidate[head][:2] + + if has_left_eye: + x1, y1 = candidate[left_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_right_eye: + x1, y1 = candidate[right_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_left_ear: + x1, y1 = candidate[left_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + if has_right_ear: + x1, y1 = candidate[right_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + x, y = x0, y0 + + x -= width + y -= width + + if x < 0: + x = 0 + + if y < 0: + y = 0 + + width1 = width * 2 + width2 = width * 2 + + if x + width > image_width: + width1 = image_width - x + + if y + width > image_height: + width2 = image_height - y + + width = min(width1, width2) + + if width >= 20: + detect_result.append([int(x), int(y), int(width)]) + + return detect_result + + +# get max index of 2d array +def npmax(array): + arrayindex = array.argmax(1) + arrayvalue = array.max(1) + i = arrayvalue.argmax() + j = arrayindex[i] + return i, j diff --git a/scepter/modules/annotator/dwpose/wholebody.py b/scepter/modules/annotator/dwpose/wholebody.py new file mode 100644 index 0000000..1ea43f3 --- /dev/null +++ b/scepter/modules/annotator/dwpose/wholebody.py @@ -0,0 +1,80 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import cv2 +import numpy as np +import onnxruntime as ort +from .onnxdet import inference_detector +from .onnxpose import inference_pose + +def HWC3(x): + assert x.dtype == np.uint8 + if x.ndim == 2: + x = x[:, :, None] + assert x.ndim == 3 + H, W, C = x.shape + assert C == 1 or C == 3 or C == 4 + if C == 3: + return x + if C == 1: + return np.concatenate([x, x, x], axis=2) + if C == 4: + color = x[:, :, 0:3].astype(np.float32) + alpha = x[:, :, 3:4].astype(np.float32) / 255.0 + y = color * alpha + 255.0 * (1.0 - alpha) + y = y.clip(0, 255).astype(np.uint8) + return y + + +def resize_image(input_image, resolution): + H, W, C = input_image.shape + H = float(H) + W = float(W) + k = float(resolution) / min(H, W) + H *= k + W *= k + H = int(np.round(H / 64.0)) * 64 + W = int(np.round(W / 64.0)) * 64 + img = cv2.resize(input_image, (W, H), interpolation=cv2.INTER_LANCZOS4 if k > 1 else cv2.INTER_AREA) + return img + +class Wholebody: + def __init__(self, onnx_det, onnx_pose, device = 'cuda:0'): + + providers = ['CPUExecutionProvider' + ] if device == 'cpu' else ['CUDAExecutionProvider'] + # onnx_det = 'annotator/ckpts/yolox_l.onnx' + # onnx_pose = 'annotator/ckpts/dw-ll_ucoco_384.onnx' + + self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers) + self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers) + + def __call__(self, ori_img): + det_result = inference_detector(self.session_det, ori_img) + keypoints, scores = inference_pose(self.session_pose, det_result, ori_img) + + keypoints_info = np.concatenate( + (keypoints, scores[..., None]), axis=-1) + # compute neck joint + neck = np.mean(keypoints_info[:, [5, 6]], axis=1) + # neck score when visualizing pred + neck[:, 2:4] = np.logical_and( + keypoints_info[:, 5, 2:4] > 0.3, + keypoints_info[:, 6, 2:4] > 0.3).astype(int) + new_keypoints_info = np.insert( + keypoints_info, 17, neck, axis=1) + mmpose_idx = [ + 17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3 + ] + openpose_idx = [ + 1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17 + ] + new_keypoints_info[:, openpose_idx] = \ + new_keypoints_info[:, mmpose_idx] + keypoints_info = new_keypoints_info + + keypoints, scores = keypoints_info[ + ..., :2], keypoints_info[..., 2] + + return keypoints, scores, det_result + + diff --git a/scepter/modules/annotator/dwpose_op.py b/scepter/modules/annotator/dwpose_op.py new file mode 100644 index 0000000..2d2cbfd --- /dev/null +++ b/scepter/modules/annotator/dwpose_op.py @@ -0,0 +1,203 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +# Openpose +# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose +# 2nd Edited by https://github.com/Hzzone/pytorch-openpose +# 3rd Edited by ControlNet +# 4th Edited by ControlNet (added face and correct hands) + +# ``` requirements for cuda 12.1: +# onnxruntime==1.19 +# onnxruntime-gpu==1.19 +# ``` + +import os + +import numpy as np +import torch +from PIL import Image + +import cv2 +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.dwpose import util +from scepter.modules.annotator.dwpose.wholebody import (HWC3, Wholebody, + resize_image) +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE' + + +def draw_pose(pose, H, W, use_hand=False, use_body=False, use_face=False): + bodies = pose['bodies'] + faces = pose['faces'] + hands = pose['hands'] + candidate = bodies['candidate'] + subset = bodies['subset'] + canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8) + + if use_body: + canvas = util.draw_bodypose(canvas, candidate, subset) + if use_hand: + canvas = util.draw_handpose(canvas, hands) + if use_face: + canvas = util.draw_facepose(canvas, faces) + + return canvas + + +@ANNOTATORS.register_class() +class DWposeAnnotator(BaseAnnotator): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + with FS.get_from(cfg['DETECTION_MODEL'], + wait_finish=True) as onnx_det, FS.get_from( + cfg['POSE_MODEL'], wait_finish=True) as onnx_pose: + self.pose_estimation = Wholebody(onnx_det, + onnx_pose, + device=f'cuda:{we.device_id}') + self.resize_size = cfg.get('RESIZE_SIZE', 1024) + self.use_body = cfg.get('USE_BODY', True) + self.use_face = cfg.get('USE_FACE', True) + self.use_hand = cfg.get('USE_HAND', True) + + @torch.no_grad() + @torch.inference_mode + def forward(self, image): + if isinstance(image, Image.Image): + image = np.array(image) + elif isinstance(image, torch.Tensor): + image = image.detach().cpu().numpy() + elif isinstance(image, np.ndarray): + image = image.copy() + else: + raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + input_image = HWC3(image[..., ::-1]) + return self.process(resize_image(input_image, self.resize_size), + image.shape[:2]) + + def process(self, ori_img, ori_shape): + ori_h, ori_w = ori_shape + ori_img = ori_img.copy() + H, W, C = ori_img.shape + with torch.no_grad(): + candidate, subset, det_result = self.pose_estimation(ori_img) + nums, keys, locs = candidate.shape + candidate[..., 0] /= float(W) + candidate[..., 1] /= float(H) + body = candidate[:, :18].copy() + body = body.reshape(nums * 18, locs) + score = subset[:, :18] + for i in range(len(score)): + for j in range(len(score[i])): + if score[i][j] > 0.3: + score[i][j] = int(18 * i + j) + else: + score[i][j] = -1 + + un_visible = subset < 0.3 + candidate[un_visible] = -1 + + foot = candidate[:, 18:24] + + faces = candidate[:, 24:92] + + hands = candidate[:, 92:113] + hands = np.vstack([hands, candidate[:, 113:]]) + + bodies = dict(candidate=body, subset=score) + pose = dict(bodies=bodies, hands=hands, faces=faces) + + ret_data = {} + if self.use_body: + detected_map_body = draw_pose(pose, H, W, use_body=True) + detected_map_body = cv2.resize( + detected_map_body[..., ::-1], (ori_w, ori_h), + interpolation=cv2.INTER_LANCZOS4 + if ori_h * ori_w > H * W else cv2.INTER_AREA) + ret_data['detected_map_body'] = detected_map_body + + if self.use_face: + detected_map_face = draw_pose(pose, H, W, use_face=True) + detected_map_face = cv2.resize( + detected_map_face[..., ::-1], (ori_w, ori_h), + interpolation=cv2.INTER_LANCZOS4 + if ori_h * ori_w > H * W else cv2.INTER_AREA) + ret_data['detected_map_face'] = detected_map_face + + if self.use_body and self.use_face: + detected_map_bodyface = draw_pose(pose, + H, + W, + use_body=True, + use_face=True) + detected_map_bodyface = cv2.resize( + detected_map_bodyface[..., ::-1], (ori_w, ori_h), + interpolation=cv2.INTER_LANCZOS4 + if ori_h * ori_w > H * W else cv2.INTER_AREA) + ret_data['detected_map_bodyface'] = detected_map_bodyface + + if self.use_hand and self.use_body and self.use_face: + detected_map_handbodyface = draw_pose(pose, + H, + W, + use_hand=True, + use_body=True, + use_face=True) + detected_map_handbodyface = cv2.resize( + detected_map_handbodyface[..., ::-1], (ori_w, ori_h), + interpolation=cv2.INTER_LANCZOS4 + if ori_h * ori_w > H * W else cv2.INTER_AREA) + ret_data[ + 'detected_map_handbodyface'] = detected_map_handbodyface + + # convert_size + if det_result.shape[0] > 0: + w_ratio, h_ratio = ori_w / W, ori_h / H + det_result[..., ::2] *= h_ratio + det_result[..., 1::2] *= w_ratio + det_result = det_result.astype(np.int32) + # for det_tup in det_result: + # cv2.rectangle(detected_map, det_tup[2:].tolist(), det_tup[:2].tolist(), color=(255, 0, 0), thickness=3) + return ret_data, det_result + + +@ANNOTATORS.register_class() +class DWposeBodyAnnotator(DWposeAnnotator): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_body, self.use_face, self.use_hand = True, False, False + + @torch.no_grad() + @torch.inference_mode + def forward(self, image): + ret_data, det_result = super().forward(image) + return ret_data['detected_map_body'] + + +@ANNOTATORS.register_class() +class DWposeFaceAnnotator(DWposeAnnotator): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_body, self.use_face, self.use_hand = False, True, False + + @torch.no_grad() + @torch.inference_mode + def forward(self, image): + ret_data, det_result = super().forward(image) + return ret_data['detected_map_face'] + + +@ANNOTATORS.register_class() +class DWposeBodyFaceAnnotator(DWposeAnnotator): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_body, self.use_face, self.use_hand = True, True, False + + @torch.no_grad() + @torch.inference_mode + def forward(self, image): + ret_data, det_result = super().forward(image) + return ret_data['detected_map_bodyface'] diff --git a/scepter/modules/annotator/face.py b/scepter/modules/annotator/face.py new file mode 100644 index 0000000..bf444a3 --- /dev/null +++ b/scepter/modules/annotator/face.py @@ -0,0 +1,63 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +from abc import ABCMeta + +import numpy as np +import torch +from PIL import Image + +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +@ANNOTATORS.register_class() +class FaceAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + from insightface.app import FaceAnalysis + local_path = FS.map_to_local(cfg.PRETRAINED_MODEL)[0] + local_model_path = os.path.join(local_path, 'models', cfg.MODEL_NAME) + FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, local_model_path) + self.model = FaceAnalysis(name=cfg.MODEL_NAME, root=local_path, providers=['CUDAExecutionProvider', 'CPUExecutionProvider']) + self.model.prepare(ctx_id=we.device_id, det_size=(640, 640)) + + def forward(self, image=None): + + if isinstance(image, Image.Image): + image = np.array(image) + elif isinstance(image, torch.Tensor): + image = image.detach().cpu().numpy() + elif isinstance(image, np.ndarray): + image = image.copy() + else: + raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + # [dict_keys(['bbox', 'kps', 'det_score', 'landmark_3d_68', 'pose', 'landmark_2d_106', 'gender', 'age', 'embedding'])] + faces = self.model.get(image) + return faces + + +@ANNOTATORS.register_class() +class FaceMaskAnnotator(FaceAnnotator): + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.multi_face = cfg.get('MULTI_FACE', True) + + def forward(self, image=None): + faces = super().forward(image) + if len(faces) > 0: + if not self.multi_face: + faces = faces[:1] + mask = np.zeros_like(image[:, :, 0]) + for face in faces: + x_min, y_min, x_max, y_max = face['bbox'].tolist() + mask[int(y_min): int(y_max) + 1, int(x_min): int(x_max) + 1] = 255 + return mask + else: + return np.zeros_like(image[:, :, 0]) diff --git a/scepter/modules/annotator/frame_reference.py b/scepter/modules/annotator/frame_reference.py new file mode 100644 index 0000000..b41d54f --- /dev/null +++ b/scepter/modules/annotator/frame_reference.py @@ -0,0 +1,58 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import random +import numpy as np + +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.config import Config + + +@ANNOTATORS.register_class() +class FrameReferenceAnnotator(BaseAnnotator): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + # first / last / firstlast / random + self.ref_cfg = cfg.get('REF_CFG', [{"mode": "first", "proba": 0.1}, + {"mode": "last", "proba": 0.1}, + {"mode": "firstlast", "proba": 0.1}, + {"mode": "random", "proba": 0.1}]) + self.ref_num = cfg.get('REF_NUM', 1) + self.ref_cfg = Config.get_dict(self.ref_cfg) if isinstance( + self.ref_cfg, Config) else self.ref_cfg + self.ref_color = cfg.get('REF_COLOR', 127.5) + + def forward(self, frames, ref_cfg=None, ref_num=None): + ref_cfg = ref_cfg if ref_cfg is not None else self.ref_cfg + ref_cfg = [ref_cfg] if not isinstance(ref_cfg, list) else ref_cfg + probas = [item['proba'] if 'proba' in item else 1.0 / len(ref_cfg) for item in ref_cfg] + sel_ref_cfg = random.choices(ref_cfg, weights=probas, k=1)[0] + mode = sel_ref_cfg['mode'] if 'mode' in sel_ref_cfg else 'original' + ref_num = int(ref_num) if ref_num is not None else self.ref_num + + frame_num = len(frames) + frame_num_range = list(range(frame_num)) + if mode == "first": + sel_idx = frame_num_range[:ref_num] + elif mode == "last": + sel_idx = frame_num_range[-ref_num:] + elif mode == "firstlast": + sel_idx = frame_num_range[:ref_num] + frame_num_range[-ref_num:] + elif mode == "random": + sel_idx = random.sample(frame_num_range, ref_num) + else: + raise NotImplementedError + + out_frames, out_masks = [], [] + for i in range(frame_num): + if i in sel_idx: + out_frame = frames[i] + out_mask = np.zeros_like(frames[i][:, :, 0]) + else: + out_frame = np.ones_like(frames[i]) * self.ref_color + out_mask = np.ones_like(frames[i][:, :, 0]) * 255 + out_frames.append(out_frame) + out_masks.append(out_mask) + return out_frames, out_masks diff --git a/scepter/modules/annotator/mask_aug.py b/scepter/modules/annotator/mask_aug.py new file mode 100644 index 0000000..7da938d --- /dev/null +++ b/scepter/modules/annotator/mask_aug.py @@ -0,0 +1,450 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import random +from abc import ABCMeta +from functools import partial + +import numpy as np +import torch +from PIL import Image, ImageDraw + +import cv2 +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.config import Config +from scepter.modules.utils.file_system import FS +from scipy import ndimage +from scipy.spatial import ConvexHull +from skimage.draw import polygon + + +@ANNOTATORS.register_class() +class MaskDrawAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.task_type = cfg.get('TASK_TYPE', 'input_box') + + def forward(self, mask=None, image=None, input_box=None, task_type=None): + task_type = task_type if task_type is not None else self.task_type + + if mask is not None: + if isinstance(mask, Image.Image): + mask = np.array(mask) + elif isinstance(mask, torch.Tensor): + mask = mask.detach().cpu().numpy() + elif isinstance(mask, np.ndarray): + mask = mask.copy() + else: + raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + if image is not None: + if isinstance(image, Image.Image): + image = np.array(image) + elif isinstance(image, torch.Tensor): + image = image.detach().cpu().numpy() + elif isinstance(image, np.ndarray): + image = image.copy() + else: + raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + mask_shape = mask.shape + if task_type == 'mask_point': + scribble = mask.transpose(1, 0) + labeled_array, num_features = ndimage.label(scribble >= 255) + centers = ndimage.center_of_mass(scribble, labeled_array, + range(1, num_features + 1)) + centers = np.array(centers) + out_mask = np.zeros(mask_shape, dtype=np.uint8) + hull = ConvexHull(centers) + hull_vertices = centers[hull.vertices] + rr, cc = polygon(hull_vertices[:, 1], hull_vertices[:, 0], + mask_shape) + out_mask[rr, cc] = 255 + elif task_type == 'mask_box': + scribble = mask.transpose(1, 0) + labeled_array, num_features = ndimage.label(scribble >= 255) + centers = ndimage.center_of_mass(scribble, labeled_array, + range(1, num_features + 1)) + centers = np.array(centers) + # (x1, y1, x2, y2) + x_min = centers[:, 0].min() + x_max = centers[:, 0].max() + y_min = centers[:, 1].min() + y_max = centers[:, 1].max() + out_mask = np.zeros(mask_shape, dtype=np.uint8) + out_mask[int(y_min):int(y_max) + 1, + int(x_min):int(x_max) + 1] = 255 + if image is not None: + out_image = image[int(y_min):int(y_max) + 1, + int(x_min):int(x_max) + 1] + elif task_type == 'input_box': + if isinstance(input_box, list): + input_box = np.array(input_box) + x_min, y_min, x_max, y_max = input_box + out_mask = np.zeros(mask_shape, dtype=np.uint8) + out_mask[int(y_min):int(y_max) + 1, + int(x_min):int(x_max) + 1] = 255 + if image is not None: + out_image = image[int(y_min):int(y_max) + 1, + int(x_min):int(x_max) + 1] + elif task_type == 'mask': + out_mask = mask + else: + raise NotImplementedError + + if image is not None: + return out_image, out_mask + else: + return out_mask + + +@ANNOTATORS.register_class() +class MaskAugAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + # original / original_expand / hull / hull_expand / bbox / bbox_expand + self.mask_cfg = cfg.get('MASK_CFG', [{ + 'mode': 'original', + 'proba': 0.1 + }, { + 'mode': 'original_expand', + 'proba': 0.1 + }, { + 'mode': 'hull', + 'proba': 0.1 + }, { + 'mode': 'hull_expand', + 'proba': 0.1, + 'kwargs': { + 'expand_rate': 0.2 + } + }, { + 'mode': 'bbox', + 'proba': 0.1 + }, { + 'mode': 'bbox_expand', + 'proba': 0.1, + 'kwargs': { + 'min_expand_rate': 0.2, + 'max_expand_rate': 0.5 + } + }]) + self.mask_cfg = Config.get_dict(self.mask_cfg) if isinstance( + self.mask_cfg, Config) else self.mask_cfg + + def forward(self, mask, mask_cfg=None): + mask_cfg = mask_cfg if mask_cfg is not None else self.mask_cfg + if not isinstance(mask, list): + is_batch = False + masks = [mask] + else: + is_batch = True + masks = mask + + mask_func = self.get_mask_func(mask_cfg) + # print(mask_func) + aug_masks = [] + for submask in masks: + mask = self.get_mask(submask) + valid, large, h, w, bbox = self.get_mask_info(mask) + # print(valid, large, h, w, bbox) + if valid: + mask = mask_func(mask, bbox, h, w) + else: + mask = mask.astype(np.uint8) + aug_masks.append(mask) + return aug_masks if is_batch else aug_masks[0] + + def get_mask(self, mask): + if isinstance(mask, Image.Image): + mask = np.array(mask) + elif isinstance(mask, torch.Tensor): + mask = mask.detach().cpu().numpy() + elif isinstance(mask, np.ndarray): + mask = mask.copy() + else: + raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + return mask + + def get_mask_info(self, mask): + h, w = mask.shape + locs = mask.nonzero() + valid = True + if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1: + valid = False + return valid, False, h, w, [0, 0, 0, 0] + + left, right = np.min(locs[1]), np.max(locs[1]) + top, bottom = np.min(locs[0]), np.max(locs[0]) + bbox = [left, top, right, bottom] + + large = False + if (right - left + 1) * (bottom - top + 1) > 0.9 * h * w: + large = True + return valid, large, h, w, bbox + + def get_expand_params(self, mask_kwargs): + if 'expand_rate' in mask_kwargs: + expand_rate = mask_kwargs['expand_rate'] + elif 'min_expand_rate' in mask_kwargs and 'max_expand_rate' in mask_kwargs: + expand_rate = random.uniform(mask_kwargs['min_expand_rate'], + mask_kwargs['max_expand_rate']) + else: + expand_rate = 0.3 + + if 'expand_iters' in mask_kwargs: + expand_iters = mask_kwargs['expand_iters'] + else: + expand_iters = random.randint(1, 10) + + if 'expand_lrtp' in mask_kwargs: + expand_lrtp = mask_kwargs['expand_lrtp'] + else: + expand_lrtp = [ + random.random(), + random.random(), + random.random(), + random.random() + ] + + return expand_rate, expand_iters, expand_lrtp + + def get_mask_func(self, mask_cfg): + if not isinstance(mask_cfg, list): + mask_cfg = [mask_cfg] + probas = [ + item['proba'] if 'proba' in item else 1.0 / len(mask_cfg) + for item in mask_cfg + ] + sel_mask_cfg = random.choices(mask_cfg, weights=probas, k=1)[0] + mode = sel_mask_cfg['mode'] if 'mode' in sel_mask_cfg else 'original' + mask_kwargs = sel_mask_cfg[ + 'kwargs'] if 'kwargs' in sel_mask_cfg else {} + + if mode == 'random': + mode = random.choice([ + 'original', 'original_expand', 'hull', 'hull_expand', 'bbox', + 'bbox_expand' + ]) + if mode == 'original': + mask_func = partial(self.generate_mask) + elif mode == 'original_expand': + expand_rate, expand_iters, expand_lrtp = self.get_expand_params( + mask_kwargs) + mask_func = partial(self.generate_mask, + expand_rate=expand_rate, + expand_iters=expand_iters, + expand_lrtp=expand_lrtp) + elif mode == 'hull': + clockwise = random.choice([ + True, False + ]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise'] + mask_func = partial(self.generate_hull_mask, clockwise=clockwise) + elif mode == 'hull_expand': + expand_rate, expand_iters, expand_lrtp = self.get_expand_params( + mask_kwargs) + clockwise = random.choice([ + True, False + ]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise'] + mask_func = partial(self.generate_hull_mask, + clockwise=clockwise, + expand_rate=expand_rate, + expand_iters=expand_iters, + expand_lrtp=expand_lrtp) + elif mode == 'bbox': + mask_func = partial(self.generate_bbox_mask) + elif mode == 'bbox_expand': + expand_rate, expand_iters, expand_lrtp = self.get_expand_params( + mask_kwargs) + mask_func = partial(self.generate_bbox_mask, + expand_rate=expand_rate, + expand_iters=expand_iters, + expand_lrtp=expand_lrtp) + else: + raise NotImplementedError + return mask_func + + def generate_mask(self, + mask, + bbox, + h, + w, + expand_rate=None, + expand_iters=None, + expand_lrtp=None): + bin_mask = mask.astype(np.uint8) + if expand_rate: + bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate, + expand_iters, expand_lrtp) + return bin_mask + + @staticmethod + def rand_expand_mask(mask, + bbox, + h, + w, + expand_rate=None, + expand_iters=None, + expand_lrtp=None): + expand_rate = 0.3 if expand_rate is None else expand_rate + expand_iters = random.randint( + 1, 10) if expand_iters is None else expand_iters + expand_lrtp = [ + random.random(), + random.random(), + random.random(), + random.random() + ] if expand_lrtp is None else expand_lrtp + # print('iters', expand_iters, 'expand_rate', expand_rate, 'expand_lrtp', expand_lrtp) + # mask = np.squeeze(mask) + left, top, right, bottom = bbox + # mask expansion + box_w = (right - left + 1) * expand_rate + box_h = (bottom - top + 1) * expand_rate + left_, right_ = int( + expand_lrtp[0] * min(box_w, left / 2) / expand_iters), int( + expand_lrtp[1] * min(box_w, (w - right) / 2) / expand_iters) + top_, bottom_ = int( + expand_lrtp[2] * min(box_h, top / 2) / expand_iters), int( + expand_lrtp[3] * min(box_h, (h - bottom) / 2) / expand_iters) + kernel_size = max(left_, right_, top_, bottom_) + if kernel_size > 0: + kernel = np.zeros((kernel_size * 2, kernel_size * 2), + dtype=np.uint8) + new_left, new_right = kernel_size - right_, kernel_size + left_ + new_top, new_bottom = kernel_size - bottom_, kernel_size + top_ + kernel[new_top:new_bottom + 1, new_left:new_right + 1] = 1 + mask = mask.astype(np.uint8) + mask = cv2.dilate(mask, kernel, + iterations=expand_iters).astype(np.uint8) + # mask = new_mask - (mask / 2).astype(np.uint8) + # mask = np.expand_dims(mask, axis=-1) + return mask + + @staticmethod + def _convexhull(image, clockwise): + # print('clockwise', clockwise) + contours, hierarchy = cv2.findContours(image, 2, 1) + cnt = np.concatenate(contours) # merge all regions + hull = cv2.convexHull(cnt, clockwise=clockwise) + hull = np.squeeze(hull, axis=1).astype(np.float32).tolist() + hull = [tuple(x) for x in hull] + return hull # b, 1, 2 + + def generate_hull_mask(self, + mask, + bbox, + h, + w, + clockwise=None, + expand_rate=None, + expand_iters=None, + expand_lrtp=None): + clockwise = random.choice([True, False + ]) if clockwise is None else clockwise + hull = self._convexhull(mask, clockwise) + mask_img = Image.new('L', (w, h), 0) + pt_list = hull + mask_img_draw = ImageDraw.Draw(mask_img) + mask_img_draw.polygon(pt_list, fill=255) + bin_mask = np.array(mask_img).astype(np.uint8) + if expand_rate: + bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate, + expand_iters, expand_lrtp) + return bin_mask + + def generate_bbox_mask(self, + mask, + bbox, + h, + w, + expand_rate=None, + expand_iters=None, + expand_lrtp=None): + left, top, right, bottom = bbox + bin_mask = np.zeros((h, w), dtype=np.uint8) + bin_mask[top:bottom + 1, left:right + 1] = 255 + if expand_rate: + bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate, + expand_iters, expand_lrtp) + return bin_mask + + +@ANNOTATORS.register_class() +class MaskLayoutAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + ram_tag_color = cfg.get('RAM_TAG_COLOR', None) + default_color = cfg.get('DEFAULT_COLOR', [0, 0, 0]) + self.use_aug = cfg.get('USE_AUG', False) + self.color_dict = {'default': tuple(default_color)} + if ram_tag_color is not None: + with FS.get_object(ram_tag_color) as object: + lines = object.decode('utf-8').strip().split('\n') + lines = [id_name_color.split('#;#') for id_name_color in lines] + self.color_dict.update({ + id_name_color[1]: tuple(eval(id_name_color[2])) + for id_name_color in lines + }) + if self.use_aug: + mask_aug_dict = {'NAME': 'MaskAugAnnotator'} + mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False) + self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg) + + def find_contours(self, mask): + # @mask: gray cv2 image + # contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) + contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL, + cv2.CHAIN_APPROX_SIMPLE) + return contours + + def draw_contours(self, canvas, contour, color): + canvas = np.ascontiguousarray(canvas, dtype=np.uint8) + canvas = cv2.drawContours(canvas, contour, -1, color, thickness=3) + return canvas + + def get_mask(self, mask): + if isinstance(mask, Image.Image): + mask = np.array(mask) + elif isinstance(mask, torch.Tensor): + mask = mask.detach().cpu().numpy() + elif isinstance(mask, np.ndarray): + mask = mask.copy() + else: + raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + return mask + + def forward(self, mask=None, color=None, label=None, mask_cfg=None): + if not isinstance(mask, list): + is_batch = False + mask = [mask] + else: + is_batch = True + + if label is not None and label in self.color_dict: + color = self.color_dict[label] + elif color is not None: + color = color + else: + color = self.color_dict['default'] + + ret_data = [] + for sub_mask in mask: + sub_mask = self.get_mask(sub_mask) + if self.use_aug: + sub_mask = self.mask_aug_anno(sub_mask, mask_cfg) + canvas = np.ones((sub_mask.shape[0], sub_mask.shape[1], 3)) * 255 + contour = self.find_contours(sub_mask) + frame = self.draw_contours(canvas, contour, color) + ret_data.append(frame) + + if is_batch: + return ret_data + else: + return ret_data[0] diff --git a/scepter/modules/annotator/outpainting.py b/scepter/modules/annotator/outpainting.py index 75f09c8..661444a 100644 --- a/scepter/modules/annotator/outpainting.py +++ b/scepter/modules/annotator/outpainting.py @@ -98,9 +98,15 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta): draw.rectangle( (left + (self.mask_blur * 2 if left > 0 else 0), up + (self.mask_blur * 2 if up > 0 else 0), mask.width - right - - (self.mask_blur * 2 if right > 0 else 0), mask.height - down - - (self.mask_blur * 2 if down > 0 else 0)), + (self.mask_blur * 2 if right > 0 else 0) - 1, mask.height - down - + (self.mask_blur * 2 if down > 0 else 0) - 1), fill='black') + # draw.rectangle( + # (left + (self.mask_blur * 2 if left > 0 else 0), up + + # (self.mask_blur * 2 if up > 0 else 0), left + src_width - + # (self.mask_blur * 2 if right > 0 else 0), up + src_height - + # (self.mask_blur * 2 if down > 0 else 0)), + # fill='black') else: bbox = self.get_box(np.array(mask)) if bbox is None: diff --git a/scepter/modules/annotator/raft.py b/scepter/modules/annotator/raft.py new file mode 100644 index 0000000..94078eb --- /dev/null +++ b/scepter/modules/annotator/raft.py @@ -0,0 +1,62 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +import random +import numpy as np +import argparse + +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.config import Config +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +try: + from raft import RAFT + from raft.utils.utils import InputPadder + from raft.utils import flow_viz +except: + import warnings + warnings.warn("ignore raft import, please pip install raft.") + + +@ANNOTATORS.register_class() +class RAFTAnnotator(BaseAnnotator): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + params = { + "small": False, + "mixed_precision": False, + "alternate_corr": False + } + params = argparse.Namespace(**params) + model = RAFT(params) + if cfg.PRETRAINED_MODEL is not None: + with FS.get_from(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + model.load_state_dict({k.replace('module.', ''): v for k, v in torch.load(local_path, map_location="cpu", weights_only=True).items()}) + self.model = model.to(we.device_id).eval() + + def forward(self, frames): + # frames / RGB + frames = [torch.from_numpy(frame.astype(np.uint8)).permute(2, 0, 1).float()[None].to(we.device_id) for frame in frames] + flow_up_list, flow_up_vis_list = [], [] + with torch.no_grad(): + for i, (image1, image2) in enumerate(zip(frames[:-1], frames[1:])): + padder = InputPadder(image1.shape) + image1, image2 = padder.pad(image1, image2) + flow_low, flow_up = self.model(image1, image2, iters=20, test_mode=True) + flow_up = flow_up[0].permute(1, 2, 0).cpu().numpy() + flow_up_vis = flow_viz.flow_to_image(flow_up) + flow_up_list.append(flow_up) + flow_up_vis_list.append(flow_up_vis) + return flow_up_list, flow_up_vis_list # RGB + + +@ANNOTATORS.register_class() +class RAFTVisAnnotator(RAFTAnnotator): + def forward(self, frames): + flow_up_list, flow_up_vis_list = super().forward(frames) + return flow_up_vis_list[:1] + flow_up_vis_list diff --git a/scepter/modules/annotator/region_canvas.py b/scepter/modules/annotator/region_canvas.py new file mode 100644 index 0000000..18a6a82 --- /dev/null +++ b/scepter/modules/annotator/region_canvas.py @@ -0,0 +1,95 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import random +from abc import ABCMeta + +import cv2 +import numpy as np +import torch +from PIL import Image + +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@ANNOTATORS.register_class() +class RegionCanvasAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.scale_range = cfg.get('SCALE_RANGE', [0.75, 1.0]) + self.canvas_value = cfg.get('CANVAS_VALUE', 255) + self.use_resize = cfg.get('USE_RESIZE', True) + self.use_canvas = cfg.get('USE_CANVAS', True) + self.use_aug = cfg.get('USE_AUG', False) + if self.use_aug: + mask_aug_dict = {'NAME': 'MaskAugAnnotator'} + mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False) + self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg) + + + def forward(self, + image, + mask, + mask_cfg=None): + if isinstance(image, Image.Image): + image = np.array(image) + elif isinstance(image, torch.Tensor): + image = image.detach().cpu().numpy() + elif isinstance(image, np.ndarray): + image = image + else: + raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + mask = np.array(mask).astype(np.uint8) + image_h, image_w = image.shape[:2] + + if self.use_aug: + mask = self.mask_aug_anno(mask, mask_cfg) + + # get region with white bg + image[np.array(mask) == 0] = self.canvas_value + x, y, w, h = cv2.boundingRect(mask) + region_crop = image[y:y + h, x:x + w] + + if self.use_resize: + # resize region + scale_min, scale_max = self.scale_range + scale_factor = random.uniform(scale_min, scale_max) + new_w, new_h = int(image_w * scale_factor), int(image_h * scale_factor) + obj_scale_factor = min(new_w/w, new_h/h) + + new_w = int(w * obj_scale_factor) + new_h = int(h * obj_scale_factor) + region_crop_resized = cv2.resize(region_crop, (new_w, new_h), interpolation=cv2.INTER_AREA) + else: + region_crop_resized = region_crop + + if self.use_canvas: + # plot region into canvas + new_canvas = np.ones_like(image) * self.canvas_value + max_x = max(0, image_w - new_w) + max_y = max(0, image_h - new_h) + new_x = random.randint(0, max_x) + new_y = random.randint(0, max_y) + + new_canvas[new_y:new_y + new_h, new_x:new_x + new_w] = region_crop_resized + else: + new_canvas = region_crop_resized + return new_canvas + + @staticmethod + def get_config_template(): + return dict_to_yaml('ANNOTATORS', + __class__.__name__, + RegionCanvasAnnotator.para_dict, + set_name=True) + + +@ANNOTATORS.register_class() +class RegionCanvasCropAnnotator(RegionCanvasAnnotator): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_resize, self.use_canvas = False, False diff --git a/scepter/modules/annotator/video_segmentation.py b/scepter/modules/annotator/video_segmentation.py new file mode 100644 index 0000000..76045a4 --- /dev/null +++ b/scepter/modules/annotator/video_segmentation.py @@ -0,0 +1,153 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from abc import ABCMeta + +import numpy as np +import torch +from PIL import Image +from scipy import ndimage +try: + from sklearn.cluster import KMeans +except: + import warnings + warnings.warn("ignore sklearn import, please pip install scikit-learn.") + +from scepter.modules.annotator.base_annotator import BaseAnnotator +from scepter.modules.annotator.registry import ANNOTATORS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.file_system import FS +import pycocotools.mask as mask_utils + + +def single_mask_to_rle(mask): + rle = mask_utils.encode(np.array(mask[:, :, None], order="F", dtype="uint8"))[0] + rle["counts"] = rle["counts"].decode("utf-8") + return rle + +def single_rle_to_mask(rle): + mask = np.array(mask_utils.decode(rle)).astype(np.uint8) + return mask + +def single_mask_to_xyxy(mask): + bbox = np.zeros((4), dtype=int) + rows, cols = np.where(np.array(mask)) + if len(rows) > 0 and len(cols) > 0: + x_min, x_max = np.min(cols), np.max(cols) + y_min, y_max = np.min(rows), np.max(rows) + bbox[:] = [x_min, y_min, x_max, y_max] + return bbox.tolist() + +@ANNOTATORS.register_class() +class SAM2DrawVideoAnnotator(BaseAnnotator, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.task_type = cfg.get('TASK_TYPE', 'input_box') + from sam2.build_sam import build_sam2_video_predictor + config_path = FS.get_from(cfg.CONFIG_PATH, local_path=cfg.CONFIG_LOCAL_PATH, wait_finish=True) + pretrained_model = FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) + self.video_predictor = build_sam2_video_predictor(config_path, pretrained_model, fill_hole_area=0) + + def forward(self, + video, + input_box=None, + mask=None, + task_type=None): + task_type = task_type if task_type is not None else self.task_type + + if mask is not None: + if isinstance(mask, Image.Image): + mask = np.array(mask) + elif isinstance(mask, torch.Tensor): + mask = mask.detach().cpu().numpy() + elif isinstance(mask, np.ndarray): + mask = mask.copy() + else: + raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.' + + if task_type == 'mask_point': + if len(mask.shape) == 3: + scribble = mask.transpose(2, 1, 0)[0] + else: + scribble = mask.transpose(1, 0) # (H, W) -> (W, H) + labeled_array, num_features = ndimage.label(scribble >= 255) + centers = ndimage.center_of_mass(scribble, labeled_array, + range(1, num_features + 1)) + point_coords = np.array(centers) + point_labels = np.array([1] * len(centers)) + sample = { + 'points': point_coords, + 'labels': point_labels + } + elif task_type == 'mask_box': + if len(mask.shape) == 3: + scribble = mask.transpose(2, 1, 0)[0] + else: + scribble = mask.transpose(1, 0) # (H, W) -> (W, H) + labeled_array, num_features = ndimage.label(scribble >= 255) + centers = ndimage.center_of_mass(scribble, labeled_array, + range(1, num_features + 1)) + centers = np.array(centers) + # (x1, y1, x2, y2) + x_min = centers[:, 0].min() + x_max = centers[:, 0].max() + y_min = centers[:, 1].min() + y_max = centers[:, 1].max() + bbox = np.array([x_min, y_min, x_max, y_max]) + sample = {'box': bbox} + elif task_type == 'input_box': + if isinstance(input_box, list): + input_box = np.array(input_box) + sample = {'box': input_box} + elif task_type == 'mask': + sample = {'mask': mask} + else: + raise NotImplementedError + + ann_frame_idx = 0 + object_id = 0 + with (torch.inference_mode(), torch.autocast("cuda", dtype=torch.bfloat16)): + + inference_state = self.video_predictor.init_state(video_path=video) + + if task_type in ['mask_point', 'mask_box', 'input_box']: + _, out_obj_ids, out_mask_logits = self.video_predictor.add_new_points_or_box( + inference_state=inference_state, + frame_idx=ann_frame_idx, + obj_id=object_id, + **sample + ) + elif task_type in ['mask']: + _, out_obj_ids, out_mask_logits = self.video_predictor.add_new_mask( + inference_state=inference_state, + frame_idx=ann_frame_idx, + obj_id=object_id, + **sample + ) + else: + raise NotImplementedError + + video_segments = {} # video_segments contains the per-frame segmentation results + for out_frame_idx, out_obj_ids, out_mask_logits in self.video_predictor.propagate_in_video(inference_state): + frame_segments = {} + for i, out_obj_id in enumerate(out_obj_ids): + mask = (out_mask_logits[i] > 0.0).cpu().numpy().squeeze(0) + frame_segments[out_obj_id] = { + "mask": single_mask_to_rle(mask), + "mask_area": int(mask.sum()), + "mask_box": single_mask_to_xyxy(mask), + } + video_segments[out_frame_idx] = frame_segments + + ret_data = { + "annotations": video_segments + } + return ret_data + + @staticmethod + def get_config_template(): + return dict_to_yaml('ANNOTATORS', + __class__.__name__, + SAM2DrawVideoAnnotator.para_dict, + set_name=True) diff --git a/scepter/modules/data/dataset/registry.py b/scepter/modules/data/dataset/registry.py index 33ddaaf..1a1820a 100644 --- a/scepter/modules/data/dataset/registry.py +++ b/scepter/modules/data/dataset/registry.py @@ -304,9 +304,10 @@ class DataObject(object): delimiter = sampler_config.get('DELIMITER', ',') path_prefix = sampler_config.get('PATH_PREFIX', '') prompt_prefix = sampler_config.get('PROMPT_PREFIX', '') + oss_prefix = sampler_config.get('OSS_PREFIX', '') return MultiLevelBatchSampler(batch_size, index_file, image_size, fields, delimiter, path_prefix, - prompt_prefix, rank, seed) + prompt_prefix, oss_prefix, rank, seed) def build_dataset_config(cfg, registry, logger=None, *args, **kwargs): diff --git a/scepter/modules/data/sampler/sampler.py b/scepter/modules/data/sampler/sampler.py index e6bb7e1..e48198c 100644 --- a/scepter/modules/data/sampler/sampler.py +++ b/scepter/modules/data/sampler/sampler.py @@ -79,6 +79,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler): self.num_fields = len(self.fields) self.delimiter = cfg.get('DELIMITER', ',') self.path_prefix = cfg.get('PATH_PREFIX', '') + oss_prefix = cfg.get('OSS_PREFIX', '') common_prob = cfg.get('PROB', 1) sub_data_weights = cfg.get('SUB_DATA_WEIGHTS', None) sub_data_weights = {} if sub_data_weights is None else sub_data_weights.get_dict( @@ -137,7 +138,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler): f"{p * common_prob} and samples'num: {sub_data['total']} in this cluster." ) self.rng = np.random.default_rng(self.seed + we.rank) - self.oss_prefix = '/'.join(index_file.split('/')[:3]) + self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix self.index_dir = os.path.dirname(index_file) def __iter__(self): @@ -434,6 +435,7 @@ class MultiLevelBatchSampler(BaseSampler): delimiter=',', path_prefix='', prompt_prefix='', + oss_prefix='', rank=0, seed=8888): self.batch_size = batch_size @@ -457,7 +459,7 @@ class MultiLevelBatchSampler(BaseSampler): 'index_level': 1, 'num_fields': self.num_fields } - self.oss_prefix = '/'.join(index_file.split('/')[:3]) + self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix self.index_dir = os.path.dirname(index_file) def __iter__(self): diff --git a/scepter/modules/model/base_model.py b/scepter/modules/model/base_model.py index f98a5e0..52d7bda 100644 --- a/scepter/modules/model/base_model.py +++ b/scepter/modules/model/base_model.py @@ -3,8 +3,10 @@ import copy import torch.nn as nn + from scepter.modules.utils.config import dict_to_yaml from scepter.modules.utils.distribute import gather_data, we +from scepter.modules.utils.model import get_parameter_dtype from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe, register_data) @@ -43,13 +45,15 @@ class BaseModel(nn.Module): self._dist_data[key][k] += v else: self._dist_data[key][k] = v + def collect_probe(self): probe_data_dict = self._probe_data for k, v in self._modules.items(): if isinstance(getattr(self, k), BaseModel): for kk, vv in getattr(self, k).collect_probe().items(): - probe_data_dict[f'{k}/{kk}'] = vv + probe_data_dict[f'{k}/{kk}'] = vv return probe_data_dict + def probe_data(self): gather_probe_data = gather_data(self._probe_data) _dist_data_list = gather_data([self._dist_data]) @@ -97,6 +101,13 @@ class BaseModel(nn.Module): self._probe_data = {} return ret_data + @property + def model_dtype(self): + """ + `torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype). + """ + return get_parameter_dtype(self) + def clear_probe(self): self._probe_data.clear() diff --git a/scepter/modules/model/registry.py b/scepter/modules/model/registry.py index 9a2e9ed..4e58c6c 100644 --- a/scepter/modules/model/registry.py +++ b/scepter/modules/model/registry.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +import torch + from scepter.modules.utils.config import Config from scepter.modules.utils.registry import Registry, build_from_config @@ -15,17 +17,25 @@ def build_model(cfg, registry, logger=None, *args, **kwargs): raise TypeError(f'Config must be type dict, got {type(cfg)}') if cfg.have('PRETRAINED_MODEL'): pretrain_cfg = cfg.PRETRAINED_MODEL - if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str, list)): + if pretrain_cfg is not None and not isinstance(pretrain_cfg, + (str, list)): raise TypeError('Pretrain parameter must be a string or list') else: pretrain_cfg = None - device = cfg.get("DEVICE", None) + if cfg.get('MODEL_DTYPE', None): + default_dtype = getattr(torch, cfg.MODEL_DTYPE) + ori_default_dtype = torch.get_default_dtype() + torch.set_default_dtype(default_dtype) + device = cfg.get('DEVICE', None) model = build_from_config(cfg, registry, logger=logger, *args, **kwargs) + if cfg.get('MODEL_DTYPE', None): + torch.set_default_dtype(ori_default_dtype) if pretrain_cfg is not None: if hasattr(model, 'load_pretrained_model'): model.load_pretrained_model(pretrain_cfg) return model + def build_diffusion(cfg, registry, logger=None, *args, **kwargs): """ After build model, load pretrained model if exists key `pretrain`. @@ -37,11 +47,13 @@ def build_diffusion(cfg, registry, logger=None, *args, **kwargs): raise TypeError(f'Config must be type dict, got {type(cfg)}') return build_from_config(cfg, registry, logger=logger, *args, **kwargs) + def build_scheduler(cfg, registry, logger=None, *args, **kwargs): if not isinstance(cfg, Config): raise TypeError(f'Config must be type dict, got {type(cfg)}') return build_from_config(cfg, registry, logger=logger, *args, **kwargs) + def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs): if not isinstance(cfg, Config): raise TypeError(f'Config must be type dict, got {type(cfg)}') @@ -60,7 +72,7 @@ LOSSES = Registry('LOSSES', build_func=build_model) TUNERS = Registry('TUNERS', build_func=build_model) # reigister cls for diffusion. - DIFFUSIONS = Registry('DIFFUSIONS', build_func=build_diffusion) NOISE_SCHEDULERS = Registry('NOISE_SCHEDULERS', build_func=build_diffusion) -DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS', build_func=build_diffusion_sampler) +DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS', + build_func=build_diffusion_sampler) diff --git a/scepter/modules/model/utils/basic_utils.py b/scepter/modules/model/utils/basic_utils.py index bc2a005..20fc19f 100644 --- a/scepter/modules/model/utils/basic_utils.py +++ b/scepter/modules/model/utils/basic_utils.py @@ -111,9 +111,11 @@ def pack_imagelist_into_tensor(image_list): image_tensor.append(img.view(c, h * w).transpose(1, 0)) # h*w, c shapes.append((h, w)) - image_tensor = pad_sequence(image_tensor, batch_first=True).permute(0, 2, 1) # b, c, l + image_tensor = pad_sequence(image_tensor, + batch_first=True).permute(0, 2, 1) # b, c, l return image_tensor, shapes + def limit_batch_data(batch_data_list, log_num): if log_num and log_num > 0: batch_data_list_limited = [] @@ -123,4 +125,4 @@ def limit_batch_data(batch_data_list, log_num): batch_data_list_limited.append(sub_data) return batch_data_list_limited else: - return batch_data_list \ No newline at end of file + return batch_data_list diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py index be2294c..11d57ec 100644 --- a/scepter/modules/solver/diffusion_solver.py +++ b/scepter/modules/solver/diffusion_solver.py @@ -196,6 +196,7 @@ class LatentDiffusionSolver(BaseSolver): self.logger.info('Use fsdp as the backend of ddp.') else: self.logger.info('Use default backend.') + self.find_unused_parameters = cfg.get('FIND_UNUSED_PARAMETERS', False) self.use_scaler = cfg.get('USE_SCALER', True) self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', False) self.use_orig_params = cfg.get('USE_ORIG_PARAMS', False) @@ -408,7 +409,7 @@ class LatentDiffusionSolver(BaseSolver): self.model, device_ids=[torch.cuda.current_device()], output_device=torch.cuda.current_device(), - find_unused_parameters=False) + find_unused_parameters=self.find_unused_parameters) self.optimizer = OPTIMIZERS.build( self.cfg.OPTIMIZER, logger=self.logger, diff --git a/scepter/modules/utils/ast_utils.py b/scepter/modules/utils/ast_utils.py index 9ab627c..bf76789 100644 --- a/scepter/modules/utils/ast_utils.py +++ b/scepter/modules/utils/ast_utils.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from __future__ import annotations + import ast import logging import os diff --git a/scepter/modules/utils/config.py b/scepter/modules/utils/config.py index 15cdd65..5acd95a 100644 --- a/scepter/modules/utils/config.py +++ b/scepter/modules/utils/config.py @@ -12,7 +12,6 @@ import yaml from scepter.modules.utils.logger import StdMsg - _SECURE_KEYWORDS = [ 'ENDPOINT', 'BUCKET', 'OSS_AK', 'OSS_SK', 'OSS', 'TOKEN', 'APPKEY' 'SECRET', 'ACCESS_ID', 'ACCESS_KEY', 'PASSWORD', 'TEMP_DIR' @@ -21,7 +20,11 @@ _SECURE_KEYWORDS = [ _SECURE_VALUEWORDS = ['oss://', 'oss-'] # -> "#####" -def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[]): +def dict_to_yaml(module_name, + name, + json_config, + set_name=False, + exclude_keys=[]): ''' { "ENV" : { "description" : "", @@ -227,6 +230,23 @@ yaml.SafeLoader.add_constructor('$', env_var_constructor) yaml.SafeLoader.add_implicit_resolver('$', pattern, None) +def check_surppor_type(v): + if isinstance(v, str) or isinstance(v, numbers.Number): + return True + elif isinstance(v, dict): + for k, v in v.items(): + if not check_surppor_type(v): + return False + return True + elif isinstance(v, list): + for v in v: + if not check_surppor_type(v): + return False + return True + else: + return False + + def _parse_args(parser): if parser is None: parser = argparse.ArgumentParser( @@ -353,13 +373,11 @@ class Config(object): if file_name.endswith('.json'): self.cfg_dict = self._load_json(file_name) self.logger.info( - f'Loading config from [{file_name}] as json file.' - ) + f'Loading config from [{file_name}] as json file.') elif file_name.endswith('.yaml'): self.cfg_dict = self._load_yaml(file_name) self.logger.info( - f'Loading config from [{file_name}] as yaml file.' - ) + f'Loading config from [{file_name}] as yaml file.') else: self.logger.info( f'No config file found! Because we do not find json or yaml in --cfg {file_name}' @@ -394,8 +412,43 @@ class Config(object): elem = float(elem) return key, elem + def recur_raw(key, elem): + if type(elem) is dict: + new_elem = {} + for k, v in elem.items(): + k, v = recur_raw(k, v) + new_elem[k] = v + return key, new_elem + elif type(elem) is list: + new_elem = [] + for idx, ele in enumerate(elem): + if type(ele) is str and ele[1:3] == 'e-': + ele = float(ele) + new_elem.append(ele) + elif type(ele) is str: + new_elem.append(ele) + elif type(ele) is dict: + new_ele = {} + for k, v in ele.items(): + k, v = recur_raw(k, v) + new_ele[k] = v + new_elem.append(new_ele) + elif type(ele) is list: + new_ele = [] + for ele_ in ele: + new_ele.append(recur_raw('', ele_)[1]) + new_elem.append(new_ele) + else: + new_elem.append(ele) + return key, new_elem + else: + if type(elem) is str and elem[1:3] == 'e-': + elem = float(elem) + return key, elem + dic = dict(recur(k, v) for k, v in cfg_dict.items()) self.__dict__.update(dic) + self.cfg_dict = dict(recur_raw(k, v) for k, v in cfg_dict.items()) def _load_json(self, cfg_file): ''' @@ -586,13 +639,12 @@ class Config(object): def __setattr__(self, key, value): super().__setattr__(key, value) - if hasattr(self, 'cfg_dict') and key in self.cfg_dict: - if isinstance(value, Config): - value = value.cfg_dict - self.cfg_dict[key] = value + if check_surppor_type(value) and key not in ['cfg_dict', 'logger']: + if hasattr(self, 'cfg_dict'): + self.cfg_dict[key] = value + self._update_dict(self.cfg_dict) def __setitem__(self, key, value): - self.__dict__[key] = value self.__setattr__(key, value) def __iter__(self): @@ -602,6 +654,7 @@ class Config(object): new_dict = {name: value} self.__dict__.update(new_dict) self.__setattr__(name, value) + self.cfg_dict.update(new_dict) def get_dict(self): return self.cfg_dict diff --git a/scepter/modules/utils/model.py b/scepter/modules/utils/model.py index a8b96d5..da580e3 100644 --- a/scepter/modules/utils/model.py +++ b/scepter/modules/utils/model.py @@ -3,9 +3,11 @@ import os import re from collections import OrderedDict +from typing import List, Tuple, Union import torch import torch.nn as nn +from torch import Tensor from torch.utils.model_zoo import load_url as load_state_dict_from_url @@ -147,3 +149,43 @@ def init_weights(module): module.weight.data.fill_(1.0) if isinstance(module, nn.Linear) and module.bias is not None: module.bias.data.zero_() + + +# copy from transformers.modeling_utils +def get_parameter_dtype(parameter: Union[nn.Module, 'ModuleUtilsMixin']): + """ + Returns the first found floating dtype in parameters if there is one, otherwise returns the last dtype it found. + """ + last_dtype = None + for t in parameter.parameters(): + last_dtype = t.dtype + if t.is_floating_point(): + return t.dtype + + if last_dtype is not None: + # if no floating dtype was found return whatever the first dtype is + return last_dtype + + # For nn.DataParallel compatibility in PyTorch > 1.5 + def find_tensor_attributes(module: nn.Module) -> List[Tuple[str, Tensor]]: + tuples = [(k, v) for k, v in module.__dict__.items() + if torch.is_tensor(v)] + return tuples + + gen = parameter._named_members(get_members_fn=find_tensor_attributes) + last_tuple = None + for tuple in gen: + last_tuple = tuple + if tuple[1].is_floating_point(): + return tuple[1].dtype + + if last_tuple is not None: + # fallback to the last dtype + return last_tuple[1].dtype + + # fallback to buffer dtype + for t in parameter.buffers(): + last_dtype = t.dtype + if t.is_floating_point(): + return t.dtype + return last_dtype diff --git a/scepter/version.py b/scepter/version.py index e56d294..e955c24 100644 --- a/scepter/version.py +++ b/scepter/version.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -__version__ = '1.4.0' +__version__ = '1.4.1' version_info = tuple(int(x) for x in __version__.split('.')[0:3]) diff --git a/scepter/workflow/calculator_node.py b/scepter/workflow/calculator_node.py index 3fcf468..294ffd2 100644 --- a/scepter/workflow/calculator_node.py +++ b/scepter/workflow/calculator_node.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. import math + from .constant import WORKFLOW_CONFIG @@ -14,16 +15,16 @@ class CalculatorNode: def INPUT_TYPES(s): return { 'required': { - 'parameter': ('INT',), - 'type': (list(s().cfg['CALCULATOR']['TYPE']),), - 'value': ('INT',), - 'round_method': (list(s().cfg['CALCULATOR']['ROUND']),) + 'parameter': ('INT', ), + 'type': (list(s().cfg['CALCULATOR']['TYPE']), ), + 'value': ('INT', ), + 'round_method': (list(s().cfg['CALCULATOR']['ROUND']), ) } } OUTPUT_NODE = True - RETURN_TYPES = ('INT',) - RETURN_NAMES = ('INT',) + RETURN_TYPES = ('INT', ) + RETURN_NAMES = ('INT', ) FUNCTION = 'execute' def execute(self, parameter, type, value, round_method): @@ -41,10 +42,11 @@ class CalculatorNode: } def _raise_zero_division(): - raise ValueError("Division by zero is not allowed.") + raise ValueError('Division by zero is not allowed.') - if not isinstance(parameter, (int, float)) or not isinstance(value, (int, float)): - raise TypeError("Parameters must be int or float.") + if not isinstance(parameter, + (int, float)) or not isinstance(value, (int, float)): + raise TypeError('Parameters must be int or float.') try: operation = _OPERATIONS[type] @@ -54,11 +56,12 @@ class CalculatorNode: try: res = operation(parameter, value) except ZeroDivisionError: - raise ValueError("Division by zero is not allowed") from None + raise ValueError('Division by zero is not allowed') from None try: round_func = _ROUND_METHODS[round_method] except KeyError: - raise ValueError(f"Invalid rounding method: {round_method}") from None + raise ValueError( + f"Invalid rounding method: {round_method}") from None - return (round_func(res),) + return (round_func(res), ) diff --git a/scepter/workflow/config/scepter_workflow.yaml b/scepter/workflow/config/scepter_workflow.yaml index 72be8e4..06a7f21 100644 --- a/scepter/workflow/config/scepter_workflow.yaml +++ b/scepter/workflow/config/scepter_workflow.yaml @@ -153,4 +153,4 @@ CALCULATOR: ROUND: - ceil - floor - - round + - round \ No newline at end of file