From 32e1787b470900fa3c504ebbf5979b7fdd433c6d Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 23 Jul 2025 10:42:39 -0400 Subject: [PATCH] add bria-3.2 model support Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 +- html/reference.json | 7 + models/Reference/briaai--BRIA-3.2.jpg | Bin 0 -> 31678 bytes modules/modeldata.py | 2 + modules/sd_detect.py | 2 + modules/sd_models.py | 4 + modules/shared_items.py | 11 +- pipelines/bria/__init__.py | 0 pipelines/bria/bria_pipeline.py | 651 ++++++++++++++++++++++++++ pipelines/bria/bria_utils.py | 443 ++++++++++++++++++ pipelines/bria/transformer_block.py | 549 ++++++++++++++++++++++ pipelines/bria/transformer_bria.py | 315 +++++++++++++ pipelines/model_bria.py | 90 ++++ wiki | 2 +- 14 files changed, 2075 insertions(+), 7 deletions(-) create mode 100644 models/Reference/briaai--BRIA-3.2.jpg create mode 100644 pipelines/bria/__init__.py create mode 100644 pipelines/bria/bria_pipeline.py create mode 100644 pipelines/bria/bria_utils.py create mode 100644 pipelines/bria/transformer_block.py create mode 100644 pipelines/bria/transformer_bria.py create mode 100644 pipelines/model_bria.py diff --git a/CHANGELOG.md b/CHANGELOG.md index e0c9e0fec..d130a0773 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,7 +6,7 @@ Feature highlights include: - **ModernUI** layout redesign which should make it more user friendly and easier to navigate -- New models [WanAI Wan 2.1](https://wan.video/) for text-to-image workflows and [FreePix F-Lite](https://huggingface.co/Freepik/F-Lite) +- New models [WanAI Wan 2.1](https://wan.video/) for text-to-image workflows, [FreePix F-Lite](https://huggingface.co/Freepik/F-Lite), [Bria 3.2](https://huggingface.co/briaai/BRIA-3.2) - Redesigned [LTXVideo](https://vladmandic.github.io/sdnext-docs/Video) interface with support for general video models plus optimized [FramePack](https://vladmandic.github.io/sdnext-docs/FramePack) and [LTXVideo](https://vladmandic.github.io/sdnext-docs/LTX) support - Fully integrated nudity detection and optional censorship with [NudeNet](https://vladmandic.github.io/sdnext-docs/NudeNet) - New background replacement and relightning methods using **Latent Bridge Matching** and new **PixelArt** processing filter @@ -44,6 +44,10 @@ For details, see [ChangeLog](https://github.com/vladmandic/automatic/blob/master - [FreePix F-Lite](https://huggingface.co/Freepik/F-Lite) F-Lite is a 10B model trained exclusively on copyright-safe and SFW content, trained on internal dataset comprising approximately 80 million copyright-safe images available via *networks -> models -> reference* + - [Bria 3.2](https://huggingface.co/briaai/BRIA-3.2) + Bria is a smaller 4B parameter model built entirely on licensed data and safe for commercial use + *note*: this is a gated model, you need to [accept terms](https://huggingface.co/briaai/BRIA-3.2) and set your [huggingface token](https://vladmandic.github.io/sdnext-docs/Gated/) + available via *networks -> models -> reference* - [LBM: Latent Bridge Matching](https://github.com/gojasper/LBM) very fast automatic image background replacement methods with relightning! *simple*: automatic background replacement using [BiRefNet](https://github.com/ZhengPeng7/BiRefNet) diff --git a/html/reference.json b/html/reference.json index c7da5a67b..4561a8eef 100644 --- a/html/reference.json +++ b/html/reference.json @@ -506,6 +506,13 @@ "skip": true }, + "Bria 3.2": { + "path": "briaai/BRIA-3.2", + "desc": "Bria 3.2 is the next-generation commercial-ready text-to-image model. With just 4 billion parameters, it provides exceptional aesthetics and text rendering, evaluated to provide on par results to leading open-source models, and outperforming other licensed models.", + "preview": "briaai--BRIA-3.2.jpg", + "skip": true + }, + "Meissonic": { "path": "MeissonFlow/Meissonic", "desc": "Meissonic is a non-autoregressive mask image modeling text-to-image synthesis model that can generate high-resolution images. It is designed to run on consumer graphics cards.", diff --git a/models/Reference/briaai--BRIA-3.2.jpg b/models/Reference/briaai--BRIA-3.2.jpg new file mode 100644 index 0000000000000000000000000000000000000000..dc7067ec55b3e74d6099abbf5e78d3791a87fb95 GIT binary patch literal 31678 zcmbTd1y~(h(k{Glhd^+5+qeXGcXtae8wlc{aCawcoP^-+4k5ThaF?HRX3or+ z`M&>o?!CRc7cHxsuIlQlx87R)vhcD7cq1<@Ck=ptf&#pM{Q+Lq0X9-zwpIXuk`jOs z001BYV4(;BaIb4nuZu7g(ZANEq38iHf9{8VJ;)AtZ2_3Ci}LI8$MpVn`Qw~F|2mo4 zI=ZnqTDp;Qv#_zSbH4tkWquLAtfs(3w)=dr7W!}DkCfP$3>t}(au~@qCL9(8yBPdiHB$siXKaq3_&h`^@wy*4>XWC`oMtYckdSfl z@CgWMXzAz~7`eFL@$mBTOGrvd%gD;fYiMd|>*(s~n_E~~S=-p!xw^S~czSvJ1b+?* z4GWI|CnP2%r=)&KOV2MTEGjN3Ei148+5l;6YHn%m>h9_7>mL{#nwp-Oots}+Tw34Q z+}hsR-P=DnJHNQRy1u!+yZ=Kjs8>1vn*Jf!f60aUDi<^?EDS8dA96uKd%i9(n6PjZ z?C@A(Y6zyzZz(x~5V6JM^Xh&eQE{rD;+VNiBI8nXt<#+SA=;mk{htZ;>HkTxe+l+) zxt0K^Fi@`#4+aw;0@xuRBY~bB$S1nFTwzpGEWq7PQ*aa1D%+q`7D4vF8Gfh6(Bc&< zhez%^o%ahh_s3pW1}4){N%T)Nig=9VZ|&-`_|7YYc}R?FJBfVyHZq{KDp4e* zww#1@=isOH7rs~%A^9$_p000ZBf4d z;TzN=K%+>j&$DzK=jcWA1m5muks3_}(r5KIKhZ>on_xAd6++KX+Ix;9G+oMUF7`!9 z{J>4=`n`mP6RVrHiUl@x!;|A7T;{E@_u|2WdDzT2n#b+hyuKuN?yZ2~%(G)D(o|qI6EPTfo%d!?_a15ds8S4c zM1U?(l|;6q9S5>uI9Su_DJ53ogATqWvR?1SUyY@}-wmm;Beo#A>Z*i*HmJ7%W<&t< zi{IgS?G1@dyQM*k$J3(ld!a)l>5(OtJ+RaOc_7jgLtZ!I3*Z-RIa35h0o*AO-&GOr z9UVG6P@D2InB$R<& z5yQz&C#v51cK)iwS@-yno69FA&9zA7Q*}|Ia9GWuVPkBmoTWQ6`Ei|NY)^#5jPFpbUd1m7G> zJn-PI!WwIawyC96i8O~Xcz*$How{-z z+7vM+`nz4?S$Yd{Y_1a^_V3N{5e)g1!w)I0%IjvSOdoHs7h>?s_H5Y}<{C4{%}(Dw zPh4}0&q-h}JXSFHwh{LJa-Kf%(s}`aImf3vbmO2XM7l12j{sfj^H~(gpV^b|rvVH} z&-5`m`+5gy28paAxiKAk1L*kvgYW^N4N=t!(Vr=TVU`LPjx+sNcPc|ei?vlQH8Vp*yQe#DNUJQoqtn1_QCh*`3Bf*jGYGrOfsLMo_n_5{Ewhq@tJ<09zI6KbnTl ze?NbAd45#@gY945=ikQk=NgKfD9V4F`d?T2*QQT#0n+(S6BRSTrf{;Sv3R14NAUKC zRiZ1Yv`+DBv z<}&fRsOA@-**uja~fiqaGY&VS<*S8aQwm|jaJz#k7$RtfV39&Ma?wN>H8`X|v7eR*x8!3w(Zd15?=qUIIMo~d!NEf8jJQS^@;DNeNw)j;BXrsihmE+?Ex;p3sp z+mn*4FCWlrIhvIhcgAjfb!LiAQ02Rl=k?U~?_}hiofmC(B(}}oi4E+FzL&-z)6cdL z)@Bg4@C^)iK@&gGxK(1SZtaK zKX@s*->fnvPZyjZ7cLHfjcv=>)|W1kwq*2I)9qrjg9+$~Gyib{l%lpA>RNe9zvj8* zCde}UBrOav?Am_kPPgoiUGb=;?TzcpMG*v4UzF1wp9k1t_WKyget;k}7b#sL6;Of~ z2L5da|C>KSNq=0JUlHsN_wk<|!y(33QFk`oh9Rn4jpSmb&M0G z*z=It*Bzd1i`ICX0uEgCYLn)P@yp(17RO)nQir(sJGxg9<=OF`F)wi z&CQ1I8v~L1@t6dVwNZ%9Dr-iZ+b}G-dRj(uwq(HG=ucc z?Re^07C+EfKA&uYl{dGUhbbv2wnPJ0D^qtxbo{8wqUpG7cRl%u zR1MsH{^?vaZvk5=qeFYh?arvZ(*;q#-;iMd`-DaO(1fWFC;H}QGgytfnGXjoc}SKU zffz9YW-h%4ca^=Z`6~ZeY)Itz9Ok()!0mSEy%SrYd`9&TE4u5Xj(mi(vLy0nOc{RFgEK<~1bM+13HMKF~V7RBsO#TL?6AnhDIMUMMSexKD+3ol1c^}yQ5xA;8I7*wt z^Bif;JV@%0r=pG^y&1x5>6ecl~U6yu2% z43!;VUsW+vH&7*WJT^MK#ZHs2J(6|SQn2x;c^9YfTEl=3s|MyaOkTc1MXc4?W*Y+m zp|oNN(j2-b8@wH7efHaE=76ChQ;T9?h$}hoOG6T%;a}e6PY1?K^JSstw}@cPbc$j*?A%i&Z5a435H2V7Op9sqIJO^$>Jul@N4gj43Z&lH zm3Y7t;8ao!J*J=|iSue%DC44r4Efq!k)$Q&geO+o*Xa>5e$X}+s3qF}qm*Z;cky@Ig-QlS#Q``2bTp?f{*O0q% z)?*X_d&I>ps^*QH`osm}C6ZB`>1xEW5e!#77tlG`m2+lb!>ou-DQl(Z3Ow7UqzR`+1P@KEOTp3F5k!lgwJ77^oLisV@029IpqeR_4zqG|~RzZwH^- zAWf9OR(Qpooqn0o>D8(H;4^vUQ*b|ma-vjkg#{Ccbw4X4m9*UJ?m+Rr=iDwC{)T3Q zYQj5X>|yN%5W)pidfycPNi72uSihRYsJ}oWIXg-uw=)p+lP-0ieAK*xjp;a3x+r%2 zSxTqg!dFwo_`5;b6#lHZL}v0cl^!2m;(Uj@M#(^o-gsIFPTrrP40=ixdh&0e`Jcf` z@xRb8lr24!P2vCS`B#+t!^r-c|Bk^4;XalgrtJEMxu|fZ6|&i5g3)Nz@fqw0JaJPY z6j7F={8^n3s2oLB5~pSyBeT{D6GVMS>?g3IQ@??6%W3j zg$`rAUvB8Qmh#2HX2|-YG|cQF|McEJPaox6OYA0wIJYIg_8+YT-$iNaV`pn+*S`QR z4-gJDsJTG5@-~~NW0h06&e&%y2>W{9NaLy_Yh- zg`)sw8gkr(kN=$icfk7hP=u0{>uTsJ1Oz?!PQw&H&PPH(>YV7$x zTwJH}c zkc9l8qiM*8>pZa+KrCU*$foj)|183Ro#lj0hKx&NdVNf^y~rhFtO9SaI%(st2e!N* zv>&6%HElzImxT7e&NgkIFTao4_B&Z%Yhhul~nG*P!M05PSf@wnTNwY5w;M9YhlPpaQFEgwP+R1mFGlFA(!W{dt%y*ofYVLp=dtiJ)``794I!;Y~t#RlY6p zHFQ#8L4~eHxv|!($IC~^CeI=MHS^U+X;c+DaS_*Svxc<`%+QG?so40f@=GkY?XAIqEk8|ur ztg82?2fGVvvN%7KrIozjA=NLSjfDF#W;-c)!7E8Xo88kPE%jD@UDFCubl%`=*hsqG zKLn9KUl`Mfl5gUO#P@4>0qA$@KoCr=y{vKZX(I`o&rq;uhh?w3?m2ytsk)iESATlR#l35tPkL?3Dg87PMF7l4o zBmaYRB!*hHr|}QlH6{3Y0thEhWu<<-cNR{)WhK-7*Zv=k{hBr@Jy#T6*AEO!kWIwZk`I=+?xwL-X0Yq07}v$+X7+FqE| zL;S6^k!GPf&>6fh!+D|*<|@!*j9Ys~^VK`N);Q>)Er~dI&tC>&7uM^h-m6u7X{BC7 z;3mijLWyS8R}kn~(9oeq&4iSFcn~a6L{NNh)CYCe5BQfdF56*VSXvNSxD7^ndWU3VdM$enj5H>blaUk$*|GtS( zx*1>5Q9-?sJq@chx@Vh!jr0u6Qf52yd8L6TnP_f8V~tsBcBH?>uk%%0wX1%zDkF(d z_@+Fy>GTASOz-%M^j({Q*L>8xmp|~@I^v4xJ=-U<{pUl!Uo){?-|91GWy@GfvS_T$ zDZzEGVW+(E2uF;mAAr6Cd1SX^aev=qA>Q_#H28iFAAAtSmj*#*sjQ5d?h+QX&+Y;=bFtoh!$r76yXF>{pfvfXAGQfSRSt*=J(z# z-z}M+YS1j_cE?4cD;f(iobwMY!#UsKRP$&;(g3;9cC0_JVlCpcOfU^$B=X0mb3VyH zY${yPJ}BU*XSC~mE~E8AFDYKxCycu0iKw?RK$r}(BJFWCdI9tr@#zCewD5LO9k$8c z5t7FzwgR{>@fG!}AuAi1e0n<#I<{#bmAX{bHeE{M?2-u`_cbMQCXhgE+Ozo+aJbL> zDRDb=*MXROUhtY<1>MW#mwco88v?}b@aiCGcsMVO3jUaLYvoAEO8SjYPK?}dS}Jem zm=ZoPR8y_;Qn`!x*4`2A*&T*{dp#3yY!bs;l zKr@1GWiYcwa2xr_H4Zb!l@X1hk=}KL!67Vx9#c|)QRTzaOTtVy#e5|NGu9Xga&47O(o>$}zWsQ9r)?7`LT@`+&k9%)1{kVWx za+4kzE2)MIQafol=Al4XvL?-fDpyNf7@Ywhh=Gbmox@D)g|}vbqxw4++a3!!E5?ko z44e#RSPRoRYLIEwUJG2kw(q8fPc~&z`10*v?kp3$_2mrPn$&F}<=X5H9X)?%yZ?ZR ze_%um@mp`orx;CQl3ni9HDC_QxsA%1X+`%aLt9yP3c6fx&lL_=d^Vg{>7n_BC_%1- z0@P}IhVGehb$Zqev}LqeavY!|Xrlw)5>kRgJy(p?qY?R4$bwf0&ATMj&%m4J;D^_G zw;Q>(RxNjFwRJgHi8c;b#QxB>vXpr2H@5EPYEml`hM`A^Y3&O+JaEBSBj0%WO;wdG z?2_e(A_w6+vLes5I;;{k4=ZOemvZ{;9Y`@8=(G{?{}le*-dda!4@M_&hCtVj*}rf zA6?%5kND28{Uc45u?KU{S!rOCF46cZzT3lx?Bi9$3C?k` z%d_5E;uym>s7Xf(ETge#{6r$HP;}t#1&u3Q_^+pMh9y}Q7*xnW$}uR zyT*DQ&E=YcNt>xIXTt;L^3XV=N*;M-_%~k*?P|Zl8lO?s2D;o~Y%vJbbx*ZlOCI+= z{}}q4E0H!1BM;1c$%(pJ%B~^-&n5e>;Mz=}V!{qPm7VI8Jo4+o)eok#2d(>wB4t87YWN;-If6y8~+G$dZ^kBUG zt(PcsdiNk;-`pC=K=jZt_*fBl8(Ae*wyL+RKR~VEZso%7D1@+`es%uH%6XkA!))8U z7b7TtNcW}MbC>4pV98SFybgO$Fiot8{59<+{x^phn~Nbi8Q1xCI;Zo6QWq5%f(l`h zh6ka9Epy6m@@>ivJ=cu(i7QL+L3&g$WPLU?r;g=io`eF*#*K__sw_^V<#lLrg?9SKzj-*b+n0j*wUsx&P z0#SnR=if1=LJB7#$eD?XHwqp<@r7p%6f#B|qz`-jGE;$dN1Rf^r^UgtpMn5QLIrWp6VLO| zWwKs83qJ;O5HCi|h|7(PU}|`q+2x+fs4YEXBc1VGHQGZOp-2`q!jJj^)37Q#-kp;a zm-BXp6gYX}p2a!Pkfs(4&)6LX)99;Xj3%)}r?!Zx20<8t2qt^g{^9vlMb4%*Ru>Iz z&$#f^@VDNgV+8Pn@000AS6O!BHa059HrIy(^W?uJ5QeRHTLI(u=`L~Bt%iRoUOIia zvA+}(-O@?0m+`a)l^PKFE_>~U$9c$daH_uz(jge)?_Eg6zTsRuy^gnAuDU$5S)MhA zt_efv(5T{RGgno}WrPFJ1JJut>K!*Rzfh)_f8?WWt_e?vT!pm1Yh3F`8=xrG7$@t# z)3G8#Vv?Aav{WxIhVZcxf~fT7h$R=Bn&95Ve7@lhTE@9dVmr+v)g^V;PbfVZH$AO5 znI;!G!r0b&B%(d0{q@`uyKTZ68C@ZJo&Q^&>0mB|d&FW)WWxpb^r;}wWiOz#tN1pq ztsR(ivLYR5s8g3bd=u?W?3G+JX5Q)1mcHR4e8Keo0X?wF;g)iIc)IGGP0ur@Y^Q{y zW?WdRp1(HfDNFbd;HJq4b*!w`>usLcMm>l)QoscDI!r)}sz}oW*n$;3&UktqS-otR zsNh*LilvfbWpD1`dUyhwl^gdtu8DRp7}^MtW{Ii}XNF-yB=%bgUP`##-W@nm-jPWE} zAU2^UjO6cekB5hBSIKsEMk@Sj+a7%KK$kz60JKxJ0U*LTeu%lhH3mWC*5NI3r z7alXHXn}R=X-kW?zReLMrJUGqkYRAjY5%FK0z;B9QZT&J{?-(VA>Yqo<-Oqh_L`oN3!DjUq9Y=;CQArTYmCP{a^Z z5JfJC1KdLyrGEBfyiu>3oideSEN-sYEK|F%))Evfs+h#{BUlf`Z)-7R--21t$rD2s z!HS0dyeJ|T}KOFiq(pYCCSHo@;i(vBu_gt$E2J~gAkS}*Z&kKA6LoODZq zH^w9gEIR=)s_(!Y;H1vIIA z6>8u${W0ofjCsa1ul|(1UB4D!)?A0r-@u*HcX=S5f1|fD;#6Fa_7xWtH%k?JMpLQ( z`2$hV4!T+g8jnt*WuO1yM7O-^9FuwpL518KKYn<+_ULBI#Q0m}uj|y92lMtK+r@KF z>`y@ZzJMwA_z+~r94_T6AJq=v5Eno;E zK=Mf~`6JeFG*(_HRMN-9V6BN{OQG;J@ZmQKo9||`sj4B)aKBdjt;se-lKq9r@CeUZ zS1uSQodRUU);z}&5idmP`tBPHk9R=G|r&r&T9eLLPaK_VA=f_FD zI!s&JbhPf89+P4GxLeJO`1Fb)E&!wi*N7bVVuun0uASj6j+mG z8F)ZxtwStxLMq2`)=zc;?i*ybd@lgQUpI?Yfv@CT$1Tcof8B|kD^Vt}?4IyE5FLbN z@T)v6E$Ny$eYfUg7R^Z3B!oe_dWmyBXeYPnivm`&kd%|5<2ybTFkxU$zW%S$yJO3R z{-@NS)j7_+EjLc zMoQk3J+lUKF`D`Jqz0RTh+^wkpDt89-zAEq4D8ar>yr@eApQ@_GB`4x_$F8&0~+BS~6FAo}^jL&|LyxB8ZaI~DHHXBx^H6zElbXk$o(gfE+ z)ghBFtv-Rpo_?pNe@~hPGLbMsC!$R6uzM4LcQ9 z$^N^_knkV1L8v5+*UI2undP4`>OUI8e^BapT*!dKM;O=bUrI=VY-oC<_{|T(El6W9 zevs`QP4E2+4*|Bs?Ns1_FGpr+Y1Oiy9cX7oFt{A9ZaiGRV^Jw#?=_mTVi zgA4S>Q*>wJumUgbaj9VnTN)vwim~3!{h2Uot+5edQq@$hXl;?bP2r^FP~^{6?UE3I z7y%=DB{fi`VE3qxmrbzccNCFFMY27{u+Zuh;I0*~D!g2I@3dcePZ_O`OC~t)L)o%C zKu?T&$3>wIEuR$!^;Sod?DFIVFjO*YhIdbLNiTLmVS2A4O?X;?`;ay^#!QOK2mSF# z1_gXyZ9by8IP1l!_0BUuqG_O?vZcbi=#yK#3L_{I?=>g7%kuJWbsRVUnIJlVPYM~Z z4Pb)c+R+AnamTrs>eh^e#2l8)?71Hcxkr=wnn!OrM#ox_0JGZ8lI#$JkZ!6<)oYL# zq~6l^$&Zw*d3xo06{*Km|46zfzpMzLr>Og~bxd;D^j6w^DY>#MxvW3wiaXhP zu6^g1)(mUnhb2VV%@{Q2cA}UgNZHKcDtSYwNPWu6NN0jb7y@_SsKY}Sn{M!+dFI^A z>E&h{15hA4vhbkM6ERmgH+{o?dcfWH{V?M>q7b7{G}mCJCDrscEsq`_ImaNZ0oaf%`)TtFTM;{f=ioziE`AHf`j;N{&Is=NNzh922O*czU&o5 z1g)%fiQrj65ZQXcvSf}a)w@-zrWQH8TV;@@t-be>2)}v-eC1FyL5Z6WStw-SgjtKA z1?Ced`s?)0XDbE*l?U2(aI=rnDoejrm0he}%su{?B3Oj*V%V1`LhPc(B3gzz8$`c8 zdl{dZ+nPG(>x*W$PGP$Jl$Yha*Y9~jTiat`NjDaBx7KW*r)(Gl77{LQh~{fW!F%QIR{4tc2S{Pbcfoy3 z>d{$|STZClSO$jD_1%ogZl-Z!#R_bn@Lnlo#LYtkahI5~v}RWu3!8L%ZMCFuQ`KaX z+$t38Wyh8SZ*Hu_j5eN(dNpP!(ZVR1sJG+tiqi%WnKF@&L{*^sL>0IdX2;++U5iBi zfrKY8geOH&Su{2U3>rHtXTGseUf>n7ODGnR-K38%06`rH5@K~oy5>RDPguVP@4^E@ zuhHYAQ&tZ*0h!IIZ=aIQTRzE_NO}E~wKR>B2Ksns{Y>X{V=lXh7IxQcPJ3gjEv-e` znWrI<~vO$UqN?j6K7dD@>-MSmo}y3aR2@RxF0W>e9Y9G%Z5m z`K|GR-ftaVhs?MO-PvlHH%1)QgJ|}~Zc8SBxGt=DAXYc7h+O)^m5@YdjHXxik!&`sJ@`c(4929lY zVwZ2&ZXJPGx}C7#!BUM2O&fJrkSD}Q^F7gVt~=HzLtur8a;d@-z7ST{OSNPvbU|0l zChQ}l`_R=unTvl{kF8mU`0T(ki8m>P3X07pU^AUiOKqNRsrWb~{0ag66Ar#{ zah`ccy>0P?OMYpCej}h{oMNGL4(tPNi{Zoel|%1+i?T_y7*^mSqTkVA6>HO-jBPHL z|611E7FnSQOg36nSz>p|$8`09M;ne4;6)YD&m zDOad;SH+FEG+>!oESx|OgN9mgcj2b|e7Q0>8{(YRpxm|fj#=qv3>L$)*ds^mplmAY#9jYgiu=Q7U9gfQ0aB>wUu=oNM+!WZ@F0sO}(@&!Z~Qf7L*;H4pxVSBFXVEl@#;y z8q$&?6W<78AJ*Km>nTe{b}A!&M0iX8EstYGGu=(T>r&UmpXrQpZn%0$Xm#|^nJJ6g z&7{eKgM{?L!N(TycyLsXe;)S1pzekC6636(7v{tR1|0htp%g|Z9H343pwByc_toY>wX ze0&$wNt$eUy)fGM5Z?f;=eL~d8}T)3$^~J<4Tr6>S7~{ru#&PE!G8K-6GHxmTu8eR zW_gbuC3&j>EOya{bV*tHQD8#gE2RS|E_&M(gPd?uW0E}+YC^MksKtk)>i8+ z$0Rei8XD{0O&G9DILV|9JTkPp;Bnvzq{e=e;*RWLXJRI6Ar5fQ2tEjNGl;paYDvl2_YOpVfh=e}wiSAj&6m%~ zB?uHPQ${`KL$nBQ<5?2;!;53wFi6=LGk@l5UuilE33Wfl)ZC>7XaD;ziT`fb?GzB zmQGcH)I0NL)ghk*pjK@9lx&x-?4_X0kPvXroeW0Tj$L+6-P==l=;FsU_pAy5(w?hw zaKzTCS`{b*sq%!nsaZUs&uIu$ki{u10UYAa+j<>K8XA}cl9YHIiA}I7S4Bn{<8*E1 znE;X*D98*LnOqWLRDuk-#UGO<(=M^V`bx<)Mh1~i?yEbLI8ps^u%Q~(SElH7Z%Ftx zIMD$Q1rrBgRmjjBb(uHdLH;v4Qv0*lLtWw}j^JiJ<68VQj_phgGLG;@jPrTY08EK~ zTCxsJEgUeK@_U>?l~56LkH3(uVYT+aqPj4Dj=XBFZ;pn-mm;Wa1CiA`ncONM5JLcK zWpK;DT-em`!HWi9jh+~g+qLvGOs z-UkbfNaZF7tYK4*^YpO9Io-^Sa95G0rYer8@oPokNPvTWT~iP>r$|wUOijaiL7bsG zHy0LZejbN1CaW)yk%hg@#lNh#mjCKl7FL8|usqR_o#xHf4HsC=%*YV%ZWne|1BHFi z(SNU%SrRiwdh#|soeP;g+uzPGQr5{l%AGc|NaN&6=zu`nVJW~-_IRlfjF7D*sXp=( z@LQ_qlEh|$dnOQcA((lwy~%B{{~KqcZ#YyOCF&AK&LA%{>zaEjZUmbHHbP{!i=VaS zO87GE@*q&&p;27Up7o5>RpFECu{xXKMUKAK55WnT{vtks-6{a9^Btja!HP#A0NqwR zg!jqz<4HG(ZOO!`uPRLcaQ_|G=;j#-0~8&sk2(tqA;73DtlIVoaagwGvL}w2bW!+N%gBzunFg)^Oj< zdt|)bg?%$P4QBb0n4}yam9lc&Y=}O_Pj*_h8*NOx$^kDns!)P!Q~s?)@U-BuoBSR; z(LA}tVLVyO|7CFFtJ_wi6}V6@#cv^0W`xmWCAMCZ{d?a2-Hct6)QP4uPX8|=vWKp! zj6k->$!D&aBU&(9$4)Y9hTRK*;i0>{X-4aN-AHDS*YkcA&Ib;|x7=m36+8u1nX0es~IIZQ=0<-z*VOEd7zRrnQ}I1$hTHKm26yDj@yhqkI(%POTU=I^G@q&lsZ0_dAjm(|T) z{lrot_QzNUIWXYQ#=r|U0pO#t9>}Qb>si?2pkL`+s$PI2eZJOjGZL1vrpYGvRYbE* zyTu`MGYAX@`aL|%vDlKZho&c`5lW=)E<->pNn@OC1lL8!@%getNCu}i34LN)WrCZ( z$?RV-4K1E*Ny(BF(OqJXb(~A>F3Oep{;_`Xn7e>kBV%xG$yh7WfrOLshAH>(?|5lU z(Jt~!7>F{{i5$Yf?$2`X1?qTRw3W?$?L_6aOG(nlom4j4l}bFA5(#Y`*{tQ=B(unn z&6Us6ZC_(t@3k|1sG&r>PWlg}ewn(;w4|4Y33ZfxsLTsh3OAUqlE`XeroN_V&K7XK zXuc*H6>s&kNbojajRj%NAiMtL%3*X8$iB)?_cLEjZH(|M9W>m;X<(u{Vl?uKm7v;z z2})lw67U277<_F)CS!85Ars5{h&0l1&`^V<4UBbWF0w>1jCdWZXXUxLyt@Y*4j?9B=|ZU& z*L7jL0C zTvscf5;kG)m_PZbjc(x?jTm&$&<*?d^@-us+MmcSycf~;Ta{v1F6);i{n_3U22-FT zJLgbcJXz{Gv>^1VC*d4d)DT?S{tiMy28sbCexU7}=Tn?VAfF!k9;!D1 zjw3I2Y)ATdc2|pxb(LKt+h!H%BmxJANA^`M*Ke0VQRWtK@a^$i^ONgyMb6ux>OeH64H6$00X7{70Z&(ny}qdzp4TxD?H(KrYEZSU*#`63T^)YJj6ZlL zGu`j358MK{&J{+6NGFX{Z`fLJn$~H|n87)_LRXoCpL}W}jY8-k zFGc0EnKypalo>lThe{F0wvn&Ul4FT7@ zmgQ6K3KqRXV)Dm~rYzP)~rza{$1E_|)}?l+t3iqaN>Mn#Ocf$z!YDB0e2scIf%OsN_QM~(Hw z-XKl&-E6{CE;6x)r{V95*M&iKvN(v%B(BP*M|e3dQ~QmUb=%+w$FM5(YySwBvX#bk z|NK)*{Se1cpV?~Vo>l+h3I|67%$jY@(K)Ce6+2<3zE6XWOvjcYT%hz4=6N9y!w^Yi zEdPP-r?EMz9&^evly=Wj_2Nb;Vb7T*G{UHlK%(R9tCGeD($NePTdhjE)|;Z^BhsE< z4`%bcAJ`u?zcut9p~pMMWjh;=FNgn@nG0SZQEgVS+tGdjP$g>CnO&nEpq3pQPJG|} zt>tfvvcOqjPE?JK=DMp6bGvsx;*)#Lv1ey!&@-Be>JOy3q5%$T%6AC&mb*%Pcppm# zFe_Uzkq!hmAB-{E#DJaC$-Mx4#tg+*2zHycS@KVE`5qECJjgXq`N0sNEXgkkZD5G8 z$L7ME$#p`YM}vC$+vu?%ZW?utc7GtkgPe0$`m7maJ_-HTvRLhp?%IzPUkKBVlJont zr9Ys-)EJp9lMp@MRZBd`J<#a)YoAlBle=piilU04t5gsp@2=yaW)!bAi3%W>MD!?B zeJ)j(%1EPW_n*g#;b%OBW7if9(1KZ+7tG_&A!?~KG(An%*Ffa?E=2_+bddjp6SGw! z_*O@XXzEtcu1Lefp)oi3HpEiK2A;>&;6XoBcnq%P%3W2KdUL95#9>`xM>fxYR3S@w zQ#IDCqxW&1y7W%0{FiNv9&~33%x(Cg;kmNG*L=1naSMoEEve3KEpNrU;ot-Y=}cd= zMhlR>(zLT(Ae3sXA(aoY4b1FJ`5y;DF;XqX(}*29_&|XI54_wFlqC7)5G>$R@s@&!@9^lZZE!!05I!ItI{CDR4d@7 z4r+>2h;}}Y9jaJl_tkVBHs!q3Fj~@+7C^R(5T4mIpR@{?U8$-^@2XhH_pjhXY!}>m zoL#B}?W-(O_-ebHc{vA0fJ)P4k8{0ELc|hP7JH1R6t~=sgr~YDyQA8o7-L$hO|;KY z=hHJ2lAacL#rA*%%}3Ge1df9PvX*j&PKah0rIkcg;s!DQUrlcv*Yx|nkAny(D2O!D zAu&Q?^gy~BHe!HugQNmdBHbo*jtyBwkMmN#1eL+>H0|)f9t>A%h`QF_| z@Xx^%s3*yXKu0KAL|Ma0nWmVgrZKOlbFR;I}KNh2+dDPa`Xz#L`2{FOdVm^O1B3HEh6h+)?Zb+iSJBnM|R#G zt(0OhG5>=;ys^n7%QN;?Y{Ae;lGu|3l`q4PCo{zhyKd@n%FU&R|*tlHJxEDin$P%5g+_Z+r^xvZK z)MKKEsap1kd|RoHxnx;ijstK#cl$0yH+}zM1!tLq^DR?1T<@!bPQXu2L3t~=y~|`> zwY)zj;{IX5HKn6r3w<>H*qN|=ZAUZufQU&_=&I&svLm6|sPFITk=U+*Mz5t+^065A z>w6Ck)93@bHYy>1A@}6O53hz;HW?_NCU$B`cUh&*$(XZ8^ngNF+zak{zn#E zR?AvwT$kV<76WihZR31*v?VFwJK00E)Q_?Fgk39W<*Ch&fzfwTwfYT>hf-E$$s#yx z3A@nR$;cC478Gru5$G&O=K_31lOnRDPu$#ghmP>dgYdC5@t&@;yj?Ha=^vd6pgh=*F<@SH)J^ottY1zx4Q zM!baTIR*i7lJDeS+P9~F$#Vvn6?n1~=rBprbkmd-v*1QczI30-Xh>cyKrHUpVs|R< zcs;BB{><-^JoI%K+tIS-raKk~t6llE&^i^H&xK&Jn_6yP2H0nKN6YgiuU|U(=<;88 z2yO;D++mV_z--vtHs^t-F^1FS@xn9dNZzR+Qh8%8UMZb^J_i(WGNGsB-vs^4;pjE0 zg$730L)y1?%lKg1_z#Og9g2i2j;*0DCq~Mv280t%>(aPEm1gQTreP|ZprP+71;F0B!e%-PmZ%l9C7eN;cd#PC;6pr?GFZF7VqIawqeAK~0 z?xK!B4X70I+0><0RBg#jOvs!I(_z;&7hREcW-&C|QQW9kPeI;I6&|Uf3U3wNpIdQx z8y_xBziOzwwDFhC0VmZvnhUpc$MO!(v(kmhvXxsnn!Hi~p9Sna`ek7t-*t}d1u}Sf zJ4^EZzsnTv)We0o+R57*D6zp8%%E0Z)0-Jn{Y_u2nSL2|Wa*Rj)Tba_aQC(9_ z$+z97<6aON!|mcaWd^W>aG7v#tR7{z28iDIY~6IC+uPveptlmtgX;bL7v}caE1!pkf9lQ}t1g z8O46NoYqc5tV=S)ISg8=pME5lFNR=@y~p&MWIwxyYUK|>!?N98kYl@AiZS%%uNx@T zJ$x?w1DZ0*R;fpAvJKsY8Uf00^^Lz8WzBrdLQ-BiutqFeHsMCnRopLnz*_z8w}+X@ zUzG1S>R81iRo|O1ZrO4tsY$q|p_Y;HPELnkk03`!4gP|phjufqVP*chm+ObqVV8pY zrdUnVtHsL=m-gMalyEPHcUq3-hE|vMy9x<#uU<{5)UF2lVcN;X&g#mC3G|>)SHs|8 z(r66nzNzv4+FZapP?nha;U9kpwB7DjsH?FTRh;-R?j~oNwU0Pbn*VyAqY3o7COb_G zsBSR@D*c}Mm@aJ_FT=wN31JQP=dYV`sFDsZ-JrTA8H|&H+#ViXI4{@=O-YBBqXiha zL@#FI1CM+v%R}#kx?oIy`nD)lrP7w_a-=)dmdCi`L0(PwOGkQ|W?y$o=M&FGPs@_- zfT$olqEWS_p$|&5LjSNF4}NM=vGVO>0uIF)6f83>+s*cxs;AlqLrB(IiT+ug5&*7^;>9iUTba0dT>2q_E#kST#E{GBgpywlE3m5n!;_+kzdyc$}`5)b2gRWcg|P{ z{Foz{T695XLeNYSx$b13L5*l6d2(J2a2-71zvDu_z9F65k=d?5_nWAB>*ZMABCEJF zw5urD37*o!{jq$iF6mHGQ>C!)?T*Xhp`p6_Y;$5LBP-7-E41GA5WbFvETFGs- zqiewm=em^DzSfG&$xB=VIzaZk_w^C)r^wW2SSmmaQ6W5hs&}rUE^farog`$`4~!dL zNcZRdB*z=Cr5v1dlQ~x`>KdcjztTjnAsT%E=4UQHDtWA!2M%Ty`or-^awe_7mN?Tu z@~2;s+GnVNw8&a8-_xpuERa3*CpG;qYvy!vp z4O-9A)zA%l6n?#q@ltt=(5;vKFsVqvPGcI!xoxz1xtm`?@E{%{CfC6xw8pQW z{f=y7Zqs2*XHC@2!x8VDJ_@uJmRg+UETcxg_dNIFA-K^$(9IH4VSKN-U6*X1@|c|V zO!vSes~^LQFo(M~^M_t3@HvCS4Z^;*=Twg+7I{o`=YfI`ocxw7T*IuWRu{o)F+BXvU*v3(}dVrg+ zZf6>a?MGnt%F29<4604LxF-sU*CoTGZ-X!V3{P5;!p8|yHTCd+Uqh*k4E<<2asOUa zHB%vWVlha|D6-vQU|U9aG$_YE0-j-owOGL_WX<4i%CN-jo?{+3*!Z3%}G*Xj^SQ7j`h2&f%iJ6`Kp$e?1fj zW1(M`2{t?IhNC7784^TNY0XVcPqTw0N{CJgUDh@Y6XgTmjWY7i(G#ivrz_ zn6paN&1P)$N>h*rKbI&$3eq;pW;2CJ?u9i4pO1_a1oAF;)0mNt5yCgD?!q(U#u>Am zX7FL-9Prp$u495HkJqA3xQuk1ChPOHHj}x%ofIKw7Z*`Hvem`B3mNMrnH5D4peW7z zIdaTSW3+OXtCwKclJQ~ykez-_c>=RKcB)1Xf8u|0)0~aWTDY24d8I&u&mNo9UGsHU z>AKV(Q)!QziEdnc}F+kqDz&y5g@Y*I65s zNd)9%n9ESVdk-@F1%b?4O;h_n6)Y0qiih`2ZoJoA%#ofgm4`Vv#~6YqhU=!>v$l>d zsCR#iB=j`hp#4~-t4%yh{ntA1wDyhOKV<_X`yjc_5<63M997s-W7`vKFffD^%R%ZR zm9mROfOF&wH*BRZeg*04?WCp)R!J456C7t7iQY~@(KHm8rG>n1o2mSRhka~su-{i- zt)FDx_lCtxu=3tUIz+ zSpiEHT++`dio>HLsP||$|HNp=5zpTD;Ua%`3NHcUu6biQJiSMYOkQC&9O+J5-`)9O zpu%9MO+Yc-9!VW92+e*Zgx$#Nz6Oh1=)`hdakN|$SiL3Kt+#g2PV^%f@aBP)Sb2&V zE%3UQ@<(vx1<%W8o@IZ4PPGWZgNb%iGz%;g82KD-B2Wn#d^b%u1~@AHZ^TnG4PB!` z@DLWCDdc~~U+)x7Gl{)U`f_yzKhb^z_3i&p+>P*9(IF(O{iZbOF_E)6UTi)n8c|7bosspvf z{E=ou2Wo%i+{FdRUm0L;SrwVp>xAkdP>JRWns*huPU6=7@wLk8dFh2Oesk13ov~R@ zTc&qRIF%xErG@pdx1~=nBxCuk?kSLZDp~WSvg~*A9{mz~L7b|G_ySSZ+(_HXL11%N zHPfwOZndfUZU(TwrPKS4=Q`TlAkp7Gw`!E>K@h5aio3_crU|OJDG6PcKPl`;mx@(B zuM=PHe`=iI zEvYF`L*yuC9Z@WvhSdLuMeebEPZOgIRx$*#O7}yfGBXN`*N6;r6E33@r7Zg0p9e7K z=r#`jGG5~&facdY_~C^(uUhA@F>>FBmnL~0&%Zb@z{YN@U&v>HO+DE7djP`V;U zvcy|Y#o&i{p-jye9eCyLA%7SSqu_6YP+vwCZ*Nmw2^|U3!o)A^&CtGd_FYN~*=dntjP zuj*{R$-Vm)wF!8*;Iq;j}k2g*G4_*ezfY8gE%b` zaA|7s5k1&^)M{wkPBMX-PF!TC!m@jf0IUX{X*Z^?T9DLK!?Ptf+sbKZ4k(hWOirCC zl7uL5Cfo6G|6zql)9>y>5B)3#SL?nG#EL^Ro;ET2FQYTZnzUHD{PXaqeCncv-Uo%Z z-1NPKM;bbqUGX2mv-h41trahXnhMcy(c)n-l-h93<(&fKxMGf8EdFgV{40E{1*`QB zc*)qiO-=W)&Q~=-Ijcggn%nU2%%vhBUHseQ$>>(x)c5Dw(gfqk=)XE<1FXph?RkYw zo-x)8y$oVEDSEgDEFng!)iwjE;Guf*FODA?xObDshbz}MM8gdq$e*{^r!e#otA=y9 z!LdGl4Tf&A=gjE1y>kWl=znh;Uc>?-%KvJVxGaoF1<%6@MW!YUv$R*U%t*uVa$dPo zUDRevDA&9rk$+PFDGV=IuSo!qn(C)bW#|SAbMAH;1VZ!mOi$zBqRX$;qG8YbCCISG z^k8qFG3Q2O{lQ_nsOsN;mO2#R2=n8`O%S?4>YhUhr5rA*O843*J}uQI+hZ;VwW`0g zc4@+MyR*lDRAk{Q-?Hq6wd}?$LTjp{zoMWO)Uju(=r)hWXZqi5+a@H799BON zfn~vEZ!|kCgiJ9u1eMj?lEF9gQmWaLcYS0;0bu)6H8JjrFeH!cL#21|%QWG!$fGQ| zKc5581-Gv@w+ZlP3O7hp6FhkV7B5fdE`U`x&>qxlI|5dZ3W*l+x%MiP#x$W8fbsHu z)y}K;-M;aZtiOo;@(9(?Mt6(Y5EfS>Z1kanrxIpXgHe{S*<+UG*<+1vg;y(Ht9zWG zG95`cxVs3UV$pK{2$lm+H&?ktal2?#wsJ>QykZc+D%f=B@-=&M06=?bq8N`hD{x^9 z6qT{XpgK))mVa%bwd5PoDEKXzSapUXL#f2PlpZslk(9lauZp}SAExaN;xa>iAhy}E z@nxTBUYdXB8}^+GRN>FZ!Z#w7R%Br@|N5@-R>z-I1mhMr(P>W~;VfpZ>73W{leruQ z>|GWf8l_u`?Hl*3KJ6Crx2FG*3DW>|;E*=n+C?L8iQ(*EKKEPd7e8vD1^0VD_ms#* z7Q_9(uKGV^(>Uu&v}#xK+k*o%G%iV22>7%b{Kc-gq-uTc(`y5g@ZGht^50$R9qZBL zF7848R0r86fPz%T3Ks)M0!DiEztJuHss218Z9bDBBteYQLg+g#pwAoaG@RD=Inju|&E$rvet(ynA(1wG5@ zP?7B(r-Ok`H3>cDd?&ot>@7rG_i>P^2n-Z7=fH?7!CcRj=C7>e-I;RlYGr^~y+wd+ zf+rR$eUy7+dJw;PYm8;Z+8lg_|saJ+xD&>2eAMA8JDW+eiT;!}3@z=EuQYaTV zX+P!}P4FBo?#8x~ubJ%sHmM*Pz~ii;ZBITuHDQTMcsu9%-h9jK8_fw87I+k&zNUZk zz)K0}unu@H_+E{{%9QtQtyOqgtD-=^L541&6Nd{kNl%*&S~nBl)424>;k#{4kdkGi zH=U77kp&4cyT!iI#5bl6p)l%k6!TFZI)x@byVG4bD<^-6`U6|eS=N>zru{QNh5krU z?4N5vc4!N9Grn+`(~FzlRTcJmrDps@)!6b?v@QpBTe&VyZ?;g0er)V4Om*{>W&z$t z?%U`6;=#{!Nu#ko_RG%V77)G)CYL#Wy00*3xv5UkWGw3d_2DD@Z0Lf{LoaUN%MJ_@$YsYYyDl89+9bmOlRPYKeXxoGb;!8J-dbA+>7*q%TRtqj+% z*=D)CfU$_*@R)yCql^H(0lKddiOqmsOU_nO$H-YbDKb`tELf7y;mVDFp0 z`RZL>{{FH_%;MJH|Yd=4AP29_;JP$eRUacCzzhEqH~M+^>5O{>`_0 zr}ye0`(}kBN9xM;?8k%0kUzC7Z?JNv?eRgH+a&9wkNCIhx4{o{_@_FzFwePTNbB#T zw2o~{^O0J@f}rSx)zDm}Ip3KhT8lz;ul=_OdboMVp_b5XQtqSS@W(0%rCOH_64Q=W zyOk?N50o&A1Xj)q@}Q&Q`II=Osx<#hr}bBFe7R)~1dTr*(m+$g(FVtQ3BBT++U9+RP+FZUDD4|;!*DLv8&pkqg7-5|E@TmWWj+`F5G{+|0;K^EYG zk3>yTE4(ATrt+%%*t!S`>0+z%If2jwR70U}^mV$A2OAz$T3X-W!^A1R68!PcCv~Dp z%kPdLw#iAI1xr~ zYIsC3uh#K<%Ni%^usJ|LLe2YKo26J?m8g-Hk+b^-f0RFrNBm2elHcelQNes1nR~$I zFh;ZC$U<(_?rR}e{NDHIUFHQ=Z4n;b{@DP9Lsl=tnjv6C{>ck~ff=!_Tsf;ll8%xm zZSl6_w2FYQ5k^sto)KK;RmoAKpbAbLa#vj;cr0|PJ#@jd=tI7ZQwousJ0l+fs1AI! z-?3`=l5Z@_3H5m;Qy@ztr~dpoXT;ah6{Lj2x2VuSZyqQ~tNW9r=PJ))xIPofRQOTE zJMO%W*Gr22+0pcJKs6SwZ@1zj@D3Txp#mwVCR>=^J>=|Cq_Ek8E1MjTHs8<#xHN>l zXbXi9oUq%0zwI6Z4AOC)d8U1maeMfkTCfB~DF-?dMB-K5N7D%m5j3!<{9xO&sgM*V z?bo1izS$s@M;kbui=Ekx)~|7uldtv2RJ43Aem>wNrn9|LDr{~vy5(7sn&hfTMCxGQ^K$VC9#f0K-&w=H2AYCusIYPP{G>NC?4^Ncu%}WZrvHIm%lB!XE81*~ zH*UIi@Wjv#XDOqJu=vBRw|gv2&oc>3rRF#;rS`KQO+R0V8Y&_4l~Yapbq!mhn5dAy zFr|__u=JvakrQJ^bpiV?Y97g6lpZe5MorbVcZD}ly!ErpV2wD=#a=MHM*`!cO-gAK z+WgCk(WD|&gPTTOdL#>-^y6V0W)HI|l37I--M=2Z{Fzp1096eST$3cvF3JzXyd0i8 zL{!#f(u=kyX)QLi17De0EG{rrI_a2H8V33Xv=g+gm-2ks#YG@nBTc3ja^r0G^1tF3 z2WZL~yt|)_CM7I(?4$7jSl(Lz8m)TNLijp`ofv6>x=GpARsYGo@=KY z6JCDtEXS(W{g0#ZXqiF}_CpT=T$66>K4)T45plqeHuJgq@W*bVz0LUoYxCORRy1aD z60!dh^zCgX9b_d3l0E3egW3>Dyg|H<7;5Fz4lr6Cu65868TWm`66;#$G22Ud=8@=f zGq*;Ef0(?Qz1w$V%#gJD4{KG{v4);-j=I~jDx#9?{sqFpJF)3+?Xt&`Q6hiu(BeHV zG^M*0>r*lw-GY#m>ko$NE3xd3>Q1l?saTf!x{e~b)aOnd+24IW_+AwvL|2h-GU%scqA>`;_3f|E!7#U#P2I zf%=<)-P{&DKx&(yyRwt^e%;i;AkhOL_G^>d083*JP!F?0v*LjtSutQ}{DYBzdx? zi!gT9tUp|HyJ-`9dRyRUdZ$Jh_~E0HN4fnLry8u#YMQrj_i7Fd9YcK}b=hVks;P?h-By%M> z=YRVYW4|&OgJFhg^-F|BMwTaOZoo*$HaJ(n)UnzM5@9j(t$3pSdbT!LJ|#0&;sB>q zudOra-3M0s5}JrdDdK$`YiX0cfqP0~H7iC+YFkXZk1CcAR2Zp6fkb=OXZzD5B`&nc z)LdF?61B|2oVy2x0rdcM5pN+OX?IzbiZ^x7ibE@jI?I#3ysvd_1GWKUctL5>m|wRu zL3mdxQRoe~^EygN8}MnVzMv-W_3`1fE^It32q;&k$HMAqDsJc-`%?5u;yt<2>#V-I z9!!Rt+xE6+|Md1Vk!Zgc#MpAA(v|G9SPIMcQqcv-94Hyq7o%^H*wdbQ$4$E6xX#KuU`Hkzz>jXFY^~0nL+N&fE z?OWK&eEuEt)!=c89tM~eI5w1}hVab5AlS2ejg!wG&l*^)KAnBD=`_UPq3n5LQn{Y1 zaZ`s2X-V+nv((ct+t+*kf(PB+{C%CYlN~pqR;js6+V=KqlOyGy-5mlDRy9N%)L&Do zB1W!SB5oMi-oa$a<=GMWN`d zsn)3z*GMMe&?C%DI`A6<#%^7}2O<))?}0ONzn04wcgbH?tDNI%E)Nyswd+5!^g|1s z;1+R+FbL92W`RgF6uiYU=3QSumT`XXpJEp0v_5K2X8A^`jML1NSqmuS+rD~GQ~rVz zA8po<<=xKv*eRiuGQ*mFCf{nPv~6z1B1hrj;Ou6hMx(@78O2|D$LTrjrNEVXOE1~3 zf!s{VECWqAp!ML=u4;KFR7{u{* z;$Mo>0d=nL#ur-rN|=v2BV0DI&$I7(^liYW4lq0K-B0m-iLy1rYGnDxhw0>JZ#QQ; zO~zy6-#dtY-{G{hS^OjNyY{yv2GP2~vZyo6x5F_1ErGl7_u=&Jq-|;t~(!^GgRhM~}2p<1eN) z$qw9f2L$>IOyswN*6EeE#2)h%lCRW`Jwu$4G8pgOC0a>LSdbbYg)%y9eJbaohk!d6quAN*UDKI7X;k5zPos&9i(Y9@~VEwqu zcykzKQFP5h_lJ6=rP-+AQDJi|f-#ZJ`GlHT*-Ks9JO%4qd7=pRrG(m923S(XAFkMx z=v}v+ZyjF-T}xiN{{?#t@HN&!IZSWI;J!7_cSHl-^bEldHq}RIPb}t0pn~@H;9U~KNUOn4lGG2uQsg%! z@>DiUrZPz^3CJrq7oNH7tx1qd&30tU%y!iaTg*BbdD2zmSb+MJMJL^4IV9Bt@hFDW>3%|XSJTK7$Fu^-r@e6f@z+5Q8;cTj4J zOB4iv>Y&+uZ$GU<%|i@#ph~+Yrzsi>Pr#^pcg-QfZOE#RT;k_E!(wOByq0ebM(9ht;DJme^$k{H!0{(ktj$0$*@VOa14js)KoQmjabJ}zqhQ}wjr%m z$+=4fS{xkk+%g-!8k#0`(X%&npuQSMp7zr(F8EV(ec)PO~1Y#Mv5nJJrB0OoQ;qZRH#F}N=XPKe?3!Tz9-Ct*_VWBX7^g5el&>~G`z zHyGq2c>u0<^n(cPo4ft9B448lmRO?CO>5UWo>`yvz&3*ESScH9drbHzd@|DaIm&?j=$6q{5otSmgmpBrIvrM%*zUHd(aW!F`@{ ziGcQ3F&(<31gD}OG}sN+{@}U`{F!oD=#~OSCNZ^4bYmUJ9+OM5htF^hmcz4|%P~;W%=KTnQ0Sinp5-i12eGy5 zeKAXlq2gUuzVqN=eDFw!%(&;q>2iGYeG{>*3#bR@M~%{%kqButbS3~Nji%NI->=8} z>+h$N#@WMZBXNpzpD!Hy(x1LJNkMC0PZz~LFRbFL^|bp)Oh5ptek8GuV{|~E{Hl3|+)#<4Cglis0^$EM!qsq5pB=wgAZO(XX_($FO@ug$QJzaAup%(Iem5vy`4uBOYKTjbdYOnajd(aK(ck3 zBVKKpfdC2-2}?=5ZbdBN{fl>7v#h7jYpk}wo=5tXIwml)Vib+eqI|G>R(Ez^QybvX z))q)e0B!8+>J60)E}70J^*{i-93G-lTZnjM=Be?y2yaxK@Q32n;z5E`o#?~ zozJ^*$YihwqCbqNd)9W=zkeLWhr9K<;)~Xpf=AjpOdO&^S#4Vk zKviPdbp=c8ru9Z}lW;*k}Jq{f`sR$=7d-nR;Rt5&+j!zt$s<11u)UYfGs#Rp(D- z@<~+}JGEa?Dk})`qD*LHfd%!t=ZP^k3rlEp;8dvw8Te(JGv403E0MCI+Z#`o}Me|TqWe@%BM-#%0H|g8R}!F zXrnsPI7G>Bblb_6&(G!V-@z$RKeLJ4mxb$hgF!5fG+$;#o;O1T>FWiEFASu*>#VA@ z&68Dt6VH9#i^@t9fPxX3bFA;Y>&$Nh-%G7yit{fzB&d2qXWC9T07=tNiWg(_+Vo?) zMEl*X`kIKZZBfN2sUNlY%p|-&$+5P{p$}F$Z68t_*>xQ-z>B7)FK~2foDftQt)2* z(N?bK^~|LLt9PRo!S0S3AiJdn@f9zr0R#|J1XMO#qN=W*rhF!zNd}NM+hzS2d4f{x z;V#V&>8^&MetFfbDrb1qfBm^yy|J)Vz9j!a$m?EpW@Zt8wkTa@;_?UZnQ$}&lbVYj z+Hkq2bkKR%nVQN^E0Mim+9b%C0k+Jqe=Io_^oI@PurzWmk9s-_Z>35mWwl!kb}F@6 zIe_{U{X=dM0s*f=wUiY-Fi-HVJFb9&e7g)$t5D+?yRHmLCrY5p*Oy#u>3%u$Rf_zH zT)tls(e6vw#@L`L*F{8h*IPPV6Ovsxh4tABB}zEdFriep{Iy!<%fTb?7rU93xTgYi z^757OaDG`_H(EiaIp$rs7J0{rKmqoFQyG=w4Rw1(h~;qYlOrSIapui^6Q!rC)b0Ji z2)b%{1uAs9)J0WKiCUGOa%Mca<5goFz#zg@*_-;M30FeS)?5eGWo?>h-51X{LpWI2 znlE5L2ZGI~m0>WDL*(_f?Vw*tW9Za8zD!9gMjA=M`l{%c(_j$ux#O4-N1NduGjbdH z2^r3MRs0oN4k2ns$HF}ynfrbEG2%1zHxyYNZ*;$Z7L$6iu`ric)TZBVF0r9h;|h z?nHm~-c{=WKTf)9|JLTQd~G2e0S$VT-_x!zXX(08`U~RG1!qYY$0342jCt2xH;NCW z*8pfY*kF)lqHpAaJv{|F+FU2&?78Jnt9VlmfG4?+1S2SnuVS?Aq1nTBIC9_(2&>;a z-!`n1%M34|7p=S%o(BTd2C@)v#4y5j z2Wd)28y@T$46EdrKmb=@wSM3dJKOzN>2^TP!vuY5fN_PteNNI$$ODO~-DRMJ&V^sd zQIQ9`GJV=YK8MM^X!?Tc^WC_;k=%+BD?Rd{6S{&VraNX(m{FAw^ixjMIoId+$f z(!HhyF#Ohiwd9J<5>3I|=UTrP>@g@?k9mgXD2xD<1@ZBM(25;fOY^10NnW5M&v02VA1V zAQ631%1s(DbA98$K9<^*tU&iu}JJ8~OXeaHPJ+^wG;pM*XUT=U`eJkF?j@DKZW|~xAFiv;p`%O@q zZg8wNAiu-XVs*!wU+Ntw>!S@&NE@wY+GJVmj-PH2A?|6%c17&K;|&+|I4$d8sp*G) zJkF7GKDFs4Muf_=l3 zDejTAK1Cn0^FSiH-s&;055z>9NDF8+p#!e$CmkDzjN?S;7&sAZM1mXzKjN!mWe;Oy z3)X&V;}%kp%hY%7O>x_Y=@{_@N_Q$u9gYVKfdEf0$R93vF^6TjWKKauZ>oqlFx(T@MAP)-lliLYk%+NzGQ?x=nw2EAG~yhTQ3#zAla0-ZTH>dj#f-JaBW3CYZvM& zCvRy8r8{+B`cGHTsCfi92B)mWG&BOiM-hrO)9js;7S_&eX9*iMg-penVM*;-De*h7 z$7T6~lw9IMNZJoI*!-9ssbYWD(BjTu{OP!q4{Ze+VuW!*kL;2okTL zP7J}G1+L-r^P*17b@_Jgx;OEd)QBG5A2WbS$YTVXDYD-)kZA}NUa7!YHseKS24nv~ntnXuG`L=xB14nb>< zaOad}+Afd*`)7>XcDB4$7vw{A)hGYGGA%c_E%u6h7~?8Tc{J6YSxKdh@>V_|DP_0q z^tuH1p;Jf@Mk1E?^Voi?(h+nKbCXrv%$ZOIj1+{yRcV1yF1^)d-ABhz%&CbO=LP>Q zMEbrA6l83Sn>n>3?YKMepCd#lSokatIT@$!_t0!lVJFv)8Np=qgY{^|5L>TTDN7B~ z|2bMRM%lNs|DW@JBmcjnfB{bbZ{+@e)A@P!67dh~h#!43?L}L&CplLabgS<7e_shd z2R<3(NNg;Y=+RVo*6$gikMkX^B*4HTuGP%rsboO^y_R-QjCv-Z1R~(AgI3GMXAyU6 zj^3WA0@?4%ZP{tMhIy2OH&1jhf4f|*rF(RWI1B^Wv5z=JA+vwnA8eU`F|=L#P*hX3 zNbxOlNVsgb^Z!QYv$r2E+y`|MuY~+{ZdqzwWRy-J^b0feJLO9ejp74ysUPazt|E@@ zFN-gg^p}d_FNUcx$rDNGk>P!^xDPTya)D=ubSjp<+c8=wc)m&;@I`esj*7nU~ zey|$Oi^$FGY=$_%n|4tGD8y)hjpxlPwYSepi#A@S2fs|6!lDZOA@sH2>ElC*tOPQl z&xehz?DN$aA@u7^Uso-RJnDzip@@lbAmXZKYD;G%V#29*=%S{OM~m@#ybyy~Vs`U` zhVBY~*`#w~)a>n~#n|$_^Nk9_7o{UoUd~ZUwKg$LXth+aq&hZwcslQ7Ep6wvdE|HfHVnPzWh(-pqge zZ}G#5_!w@H`8)nI9fU0MX419VXP8Iqu%yDv)`tu2mpVQaiV;gcZ;^6iMTb8Bne!FC zdFQ6#m}UhWbGI{2)3BtVBdr6+jLZ{-MasarpKCz1@vEXk!k-RScLvI;AbX4=S>V(f zTx;ULXr8AZQd-n+`u!AQH{ZB?m)IB!v4b%#L(s}LlW%U249yiMZz;cx6u6OnC?LE= z3$2AHe6_c~PD^*IXd=w&K-2!booyl?4-$rHTw|cb5dy-W*3pKF*A?3qV{rp*I%^Hj zXW>#{+MyK@Rfz?u*}S8KfqX{9blXA<8oHG+JwUXtw5!BIr%)TwgfTm)CNrdy&=Q25}(9VhwMR-NaY+ZHE6){|^tU>~{bF literal 0 HcmV?d00001 diff --git a/modules/modeldata.py b/modules/modeldata.py index 11b3b7c61..48d21028d 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -56,6 +56,8 @@ def get_model_type(pipe): model_type = 'pixartsigma' elif "PixArtAlpha" in name: model_type = 'pixartalpha' + elif "Bria" in name: + model_type = 'bria' # video models elif "CogVideo" in name: model_type = 'cogvideo' diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 93dcd7e73..b16478580 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -99,6 +99,8 @@ def guess_by_name(fn, current_guess): return 'FLite' elif 'wan' in fn.lower(): return 'WanAI' + elif 'bria' in fn.lower(): + return 'Bria' return current_guess diff --git a/modules/sd_models.py b/modules/sd_models.py index 9dba245ed..d8ef24a3e 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -369,6 +369,10 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op=' from pipelines.model_wanai import load_wan sd_model = load_wan(checkpoint_info, diffusers_load_config) allow_post_quant = False + elif model_type in ['Bria']: + from pipelines.model_bria import load_bria + sd_model = load_bria(checkpoint_info, diffusers_load_config) + allow_post_quant = False except Exception as e: shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}') if debug_load: diff --git a/modules/shared_items.py b/modules/shared_items.py index e735c68ab..483186905 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -47,12 +47,13 @@ pipelines = { 'WanAI': getattr(diffusers, 'WanPipeline', None), # dynamically imported and redefined later - 'Meissonic': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser - 'Monetico': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser - 'OmniGen2': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser - 'InstaFlow': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser - 'SegMoE': getattr(diffusers, 'DiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser + 'Meissonic': getattr(diffusers, 'DiffusionPipeline', None), + 'Monetico': getattr(diffusers, 'DiffusionPipeline', None), + 'OmniGen2': getattr(diffusers, 'DiffusionPipeline', None), + 'InstaFlow': getattr(diffusers, 'DiffusionPipeline', None), + 'SegMoE': getattr(diffusers, 'DiffusionPipeline', None), 'FLite': getattr(diffusers, 'DiffusionPipeline', None), + 'Bria': getattr(diffusers, 'DiffusionPipeline', None), } initialize_onnx() diff --git a/pipelines/bria/__init__.py b/pipelines/bria/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/pipelines/bria/bria_pipeline.py b/pipelines/bria/bria_pipeline.py new file mode 100644 index 000000000..c07f6e4de --- /dev/null +++ b/pipelines/bria/bria_pipeline.py @@ -0,0 +1,651 @@ +from diffusers.pipelines.flux.pipeline_flux import FluxPipeline, retrieve_timesteps, calculate_shift +from typing import Any, Callable, Dict, List, Optional, Union + +import torch + +from transformers import ( + T5EncoderModel, + T5TokenizerFast, +) + +from diffusers.image_processor import VaeImageProcessor +from diffusers import AutoencoderKL , DDIMScheduler, EulerAncestralDiscreteScheduler +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from diffusers.schedulers import KarrasDiffusionSchedulers +from diffusers.loaders import FluxLoraLoaderMixin +from diffusers.utils import ( + USE_PEFT_BACKEND, + is_torch_xla_available, + logging, + replace_example_docstring, + scale_lora_layers, + unscale_lora_layers, +) +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput +from pipelines.bria.transformer_bria import BriaTransformer2DModel +from pipelines.bria.bria_utils import get_t5_prompt_embeds, get_original_sigmas, is_ng_none +from diffusers.utils.torch_utils import randn_tensor +import diffusers +import numpy as np +if is_torch_xla_available(): + import torch_xla.core.xla_model as xm + + XLA_AVAILABLE = True +else: + XLA_AVAILABLE = False + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + +EXAMPLE_DOC_STRING = """ + Examples: + ```py + >>> import torch + >>> from diffusers import StableDiffusion3Pipeline + + >>> pipe = StableDiffusion3Pipeline.from_pretrained( + ... "stabilityai/stable-diffusion-3-medium-diffusers", torch_dtype=torch.float16 + ... ) + >>> pipe.to("cuda") + >>> prompt = "A cat holding a sign that says hello world" + >>> image = pipe(prompt).images[0] + >>> image.save("sd3.png") + ``` +""" + +""" +Based on FluxPipeline with several changes: +- no pooled embeddings +- We use zero padding for prompts +- No guidance embedding since this is not a distilled version +""" +class BriaPipeline(FluxPipeline): + r""" + Args: + transformer ([`SD3Transformer2DModel`]): + Conditional Transformer (MMDiT) architecture to denoise the encoded image latents. + scheduler ([`FlowMatchEulerDiscreteScheduler`]): + A scheduler to be used in combination with `transformer` to denoise the encoded image latents. + vae ([`AutoencoderKL`]): + Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations. + text_encoder ([`T5EncoderModel`]): + Frozen text-encoder. Stable Diffusion 3 uses + [T5](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5EncoderModel), specifically the + [t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant. + tokenizer (`T5TokenizerFast`): + Tokenizer of class + [T5Tokenizer](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5Tokenizer). + """ + + def __init__( + self, + transformer: BriaTransformer2DModel, + scheduler: Union[FlowMatchEulerDiscreteScheduler,KarrasDiffusionSchedulers], + vae: AutoencoderKL, + text_encoder: T5EncoderModel, + tokenizer: T5TokenizerFast + ): + self.register_modules( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer, + scheduler=scheduler, + ) + + # TODO - why different than offical flux (-1) + self.vae_scale_factor = ( + 2 ** (len(self.vae.config.block_out_channels)) if hasattr(self, "vae") and self.vae is not None else 16 + ) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) + self.default_sample_size = 64 # due to patchify=> 128,128 => res of 1k,1k + + # T5 is senstive to precision so we use the precision used for precompute and cast as needed + for block in self.text_encoder.encoder.block: + block.layer[-1].DenseReluDense.wo.to(dtype=torch.float32) + + if self.vae.config.shift_factor is None: + self.vae.config.shift_factor=0 + self.vae.to(dtype=torch.float32) + + + def encode_prompt( + self, + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + num_images_per_prompt: int = 1, + do_classifier_free_guidance: bool = True, + negative_prompt: Optional[Union[str, List[str]]] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + max_sequence_length: int = 128, + lora_scale: Optional[float] = None, + ): + r""" + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + device: (`torch.device`): + torch device + num_images_per_prompt (`int`): + number of images that should be generated per prompt + do_classifier_free_guidance (`bool`): + whether to use classifier free guidance or not + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + """ + device = device or self._execution_device + + # set lora scale so that monkey patched LoRA + # function of text encoder can correctly access it + if lora_scale is not None and isinstance(self, FluxLoraLoaderMixin): + self._lora_scale = lora_scale + + # dynamically adjust the LoRA scale + if self.text_encoder is not None and USE_PEFT_BACKEND: + scale_lora_layers(self.text_encoder, lora_scale) + + prompt = [prompt] if isinstance(prompt, str) else prompt + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds = get_t5_prompt_embeds( + self.tokenizer, + self.text_encoder, + prompt=prompt, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + ).to(dtype=self.transformer.dtype) + + if do_classifier_free_guidance and negative_prompt_embeds is None: + if not is_ng_none(negative_prompt): + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + + negative_prompt_embeds = get_t5_prompt_embeds( + self.tokenizer, + self.text_encoder, + prompt=negative_prompt, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + ).to(dtype=self.transformer.dtype) + else: + negative_prompt_embeds = torch.zeros_like(prompt_embeds) + + if self.text_encoder is not None: + if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND: + # Retrieve the original scale by scaling back the LoRA layers + unscale_lora_layers(self.text_encoder, lora_scale) + + dtype = self.text_encoder.dtype if self.text_encoder is not None else self.transformer.dtype + text_ids = torch.zeros(batch_size, prompt_embeds.shape[1], 3).to(device=device, dtype=dtype) + text_ids = text_ids.repeat(num_images_per_prompt, 1, 1) + + return prompt_embeds, negative_prompt_embeds, text_ids + + @property + def guidance_scale(self): + return self._guidance_scale + + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + @property + def do_classifier_free_guidance(self): + return self._guidance_scale > 1 + + @property + def joint_attention_kwargs(self): + return self._joint_attention_kwargs + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def interrupt(self): + return self._interrupt + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Union[str, List[str]] = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 30, + timesteps: List[int] = None, + guidance_scale: float = 5, + negative_prompt: Optional[Union[str, List[str]]] = None, + num_images_per_prompt: Optional[int] = 1, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 128, + clip_value:Union[None,float] = None, + normalize:bool = False + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. + instead. + height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The height in pixels of the generated image. This is set to 1024 by default for the best results. + width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The width in pixels of the generated image. This is set to 1024 by default for the best results. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + timesteps (`List[int]`, *optional*): + Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument + in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is + passed will be used. Must be in descending order. + guidance_scale (`float`, *optional*, defaults to 5.0): + Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). + `guidance_scale` is defined as `w` of equation 2. of [Imagen + Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale > + 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, + usually at the expense of lower image quality. + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + num_images_per_prompt (`int`, *optional*, defaults to 1): + The number of images to generate per prompt. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + latents (`torch.FloatTensor`, *optional*): + Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor will ge generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead + of a plain tuple. + joint_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + max_sequence_length (`int` defaults to 256): Maximum sequence length to use with the `prompt`. + + Examples: + + Returns: + [`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict` + is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated + images. + """ + + height = height or self.default_sample_size * self.vae_scale_factor + width = width or self.default_sample_size * self.vae_scale_factor + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt=prompt, + height=height, + width=width, + prompt_embeds=prompt_embeds, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + max_sequence_length=max_sequence_length, + ) + + self._guidance_scale = guidance_scale + self._joint_attention_kwargs = joint_attention_kwargs + self._interrupt = False + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + + lora_scale = ( + self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None + ) + + ( + prompt_embeds, + negative_prompt_embeds, + text_ids + ) = self.encode_prompt( + prompt=prompt, + negative_prompt=negative_prompt, + do_classifier_free_guidance=self.do_classifier_free_guidance, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + + if self.do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) + + + + # 5. Prepare latent variables + num_channels_latents = self.transformer.config.in_channels // 4 # due to patch=2, we devide by 4 + latents, latent_image_ids = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + if isinstance(self.scheduler,FlowMatchEulerDiscreteScheduler) and self.scheduler.config['use_dynamic_shifting']: + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + image_seq_len = latents.shape[1] # Shift by height - Why just height? + + mu = calculate_shift( + image_seq_len, + self.scheduler.config.base_image_seq_len, + self.scheduler.config.max_image_seq_len, + self.scheduler.config.base_shift, + self.scheduler.config.max_shift, + ) + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, + num_inference_steps, + device, + timesteps, + sigmas, + mu=mu, + ) + else: + # 4. Prepare timesteps + # Sample from training sigmas + if isinstance(self.scheduler,DDIMScheduler) or isinstance(self.scheduler,EulerAncestralDiscreteScheduler): + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None, None) + else: + sigmas = get_original_sigmas(num_train_timesteps=self.scheduler.config.num_train_timesteps,num_inference_steps=num_inference_steps) + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps,sigmas=sigmas) + + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self._num_timesteps = len(timesteps) + + # Supprot different diffusers versions + if diffusers.__version__>='0.32.0': + latent_image_ids=latent_image_ids[0] + text_ids=text_ids[0] + + # 6. Denoising loop + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + if self.interrupt: + continue + + # expand the latents if we are doing classifier free guidance + latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents + if type(self.scheduler)!=FlowMatchEulerDiscreteScheduler: + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latent_model_input.shape[0]) + + # This is predicts "v" from flow-matching or eps from diffusion + noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + txt_ids=text_ids, + img_ids=latent_image_ids, + )[0] + + # perform guidance + if self.do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + cfg_noise_pred_text = noise_pred_text.std() + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond) + + if normalize: + noise_pred = noise_pred * (0.7 *(cfg_noise_pred_text/noise_pred.std())) + 0.3 * noise_pred + + if clip_value: + assert clip_value>0 + noise_pred = noise_pred.clip(-clip_value,clip_value) + + # compute the previous noisy sample x_t -> x_t-1 + latents_dtype = latents.dtype + latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] + + if latents.dtype != latents_dtype: + if torch.backends.mps.is_available(): + # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 + latents = latents.to(latents_dtype) + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + + if XLA_AVAILABLE: + xm.mark_step() + + if output_type == "latent": + image = latents + + else: + latents = self._unpack_latents(latents, height, width, self.vae_scale_factor) + latents = (latents.to(dtype=torch.float32) / self.vae.config.scaling_factor) + self.vae.config.shift_factor + image = self.vae.decode(latents.to(dtype=self.vae.dtype), return_dict=False)[0] + image = self.image_processor.postprocess(image, output_type=output_type) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (image,) + + return FluxPipelineOutput(images=image) + + def check_inputs( + self, + prompt, + height, + width, + negative_prompt=None, + prompt_embeds=None, + negative_prompt_embeds=None, + callback_on_step_end_tensor_inputs=None, + max_sequence_length=None, + ): + if height % 8 != 0 or width % 8 != 0: + raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.") + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." + ) + elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): + raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") + + if negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + + if prompt_embeds is not None and negative_prompt_embeds is not None: + if prompt_embeds.shape != negative_prompt_embeds.shape: + raise ValueError( + "`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but" + f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`" + f" {negative_prompt_embeds.shape}." + ) + + if max_sequence_length is not None and max_sequence_length > 512: + raise ValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}") + + def to(self, *args, **kwargs): + DiffusionPipeline.to(self, *args, **kwargs) + # T5 is senstive to precision so we use the precision used for precompute and cast as needed + for block in self.text_encoder.encoder.block: + block.layer[-1].DenseReluDense.wo.to(dtype=torch.float32) + + if self.vae.config.shift_factor == 0 and self.vae.dtype!=torch.float32: + self.vae.to(dtype=torch.float32) + + + return self + + + def prepare_latents( + self, + batch_size, + num_channels_latents, + height, + width, + dtype, + device, + generator, + latents=None, + ): + # VAE applies 8x compression on images but we must also account for packing which requires + # latent height and width to be divisible by 2. + height = 2 * (int(height) // self.vae_scale_factor) + width = 2 * (int(width) // self.vae_scale_factor ) + + shape = (batch_size, num_channels_latents, height, width) + + if latents is not None: + latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype) + return latents.to(device=device, dtype=dtype), latent_image_ids + + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width) + + latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype) + + return latents, latent_image_ids + + @staticmethod + def _pack_latents(latents, batch_size, num_channels_latents, height, width): + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 2, 4, 1, 3, 5) + latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4) + + return latents + + @staticmethod + def _unpack_latents(latents, height, width, vae_scale_factor): + batch_size, num_patches, channels = latents.shape + + height = height // vae_scale_factor + width = width // vae_scale_factor + + latents = latents.view(batch_size, height, width, channels // 4, 2, 2) + latents = latents.permute(0, 3, 1, 4, 2, 5) + + latents = latents.reshape(batch_size, channels // (2 * 2), height * 2, width * 2) + + return latents + + @staticmethod + def _prepare_latent_image_ids(batch_size, height, width, device, dtype): + latent_image_ids = torch.zeros(height, width, 3) + latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None] + latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :] + + latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape + + latent_image_ids = latent_image_ids.repeat(batch_size, 1, 1, 1) + latent_image_ids = latent_image_ids.reshape( + batch_size, latent_image_id_height * latent_image_id_width, latent_image_id_channels + ) + + return latent_image_ids.to(device=device, dtype=dtype) diff --git a/pipelines/bria/bria_utils.py b/pipelines/bria/bria_utils.py new file mode 100644 index 000000000..3cddeafa1 --- /dev/null +++ b/pipelines/bria/bria_utils.py @@ -0,0 +1,443 @@ +from typing import Union, Optional, List +import torch +from diffusers.utils import logging +from transformers import ( + T5EncoderModel, + T5TokenizerFast, + AutoTokenizer +) +from transformers import ( + CLIPTextModel, + CLIPTextModelWithProjection, + CLIPTokenizer +) + +import numpy as np +import torch.distributed as dist +import math +import os + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def get_text(caption): + + existing_text_list = set() + + if caption[0]=='\"' and caption[-1]=='\"': + caption=caption[1:-2] + + if caption[0]=='\'' and caption[-1]=='\'': + caption=caption[1:-2] + + text_list=[] + current_text='' + text_present = False + for c in caption: + if c=='\"' and not text_present: + text_present=True + continue + + if c=='\"' and text_present: + if current_text not in existing_text_list: + text_list+=[current_text] + existing_text_list.add(current_text) + + text_present=False + current_text='' + continue + + if text_present: + current_text+=c + + return text_list + +def get_by_t5_prompt_embeds( + tokenizer: AutoTokenizer , + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + max_sequence_length: int = 128, + device: Optional[torch.device] = None, +): + device = device or text_encoder.device + + if isinstance(prompt, list): + assert len(prompt)==1 + prompt=prompt[0] + + assert type(prompt)==str + + captions_list = get_text(prompt) + embeddings_list=[] + for inner_prompt in captions_list: + text_inputs = tokenizer( + [inner_prompt], + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + embeddings_list+=[prompt_embeds[0]] + + # No Text Found + if len(embeddings_list)==0: + return None + + prompt_embeds = torch.concatenate(embeddings_list,axis=0) + + # Concat zeros to max_sequence + seq_len, dim = prompt_embeds.shape + if seq_len= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because `max_sequence_length` is set to " + f" {max_sequence_length} tokens: {removed_text}" + ) + + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + + # Concat zeros to max_sequence + b, seq_len, dim = prompt_embeds.shape + if seq_len torch.Tensor: + n_axes = ids.shape[-1] + cos_out = [] + sin_out = [] + pos = ids.float() + is_mps = ids.device.type == "mps" + freqs_dtype = torch.float32 if is_mps else torch.float64 + for i in range(n_axes): + cos, sin = get_1d_rotary_pos_embed( + self.axes_dim[i], + pos[:, i], + theta=self.theta, + repeat_interleave_real=True, + use_real=True, + freqs_dtype=freqs_dtype, + ) + cos_out.append(cos) + sin_out.append(sin) + freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device) + freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device) + return freqs_cos, freqs_sin + +from diffusers.optimization import get_scheduler +from torch.optim import Optimizer +from torch.optim.lr_scheduler import LambdaLR + +# Not really cosine but with decay +def get_cosine_schedule_with_warmup_and_decay( + optimizer: Optimizer, num_warmup_steps: int, num_training_steps: int, num_cycles: float = 0.5, last_epoch: int = -1, constant_steps=-1,eps=1e-5 +) -> LambdaLR: + + """ + Create a schedule with a learning rate that decreases following the values of the cosine function between the + initial lr set in the optimizer to 0, after a warmup period during which it increases linearly between 0 and the + initial lr set in the optimizer. + + Args: + optimizer ([`~torch.optim.Optimizer`]): + The optimizer for which to schedule the learning rate. + num_warmup_steps (`int`): + The number of steps for the warmup phase. + num_training_steps (`int`): + The total number of training steps. + num_periods (`float`, *optional*, defaults to 0.5): + The number of periods of the cosine function in a schedule (the default is to just decrease from the max + value to 0 following a half-cosine). + last_epoch (`int`, *optional*, defaults to -1): + The index of the last epoch when resuming training. + constant_steps (`int`): + The total number of constant lr steps following a warmup + + Return: + `torch.optim.lr_scheduler.LambdaLR` with the appropriate schedule. + """ + if constant_steps <=0: + constant_steps = num_training_steps-num_warmup_steps + + def lr_lambda(current_step): + # Accelerate sends current_step*num_processes + if current_step < num_warmup_steps: + return float(current_step) / float(max(1, num_warmup_steps)) + elif current_step torch.Tensor: + residual = hidden_states + norm_hidden_states, gate = self.norm(hidden_states, emb=temb) + mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) + joint_attention_kwargs = joint_attention_kwargs or {} + attn_output = self.attn( + hidden_states=norm_hidden_states, + image_rotary_emb=image_rotary_emb, + **joint_attention_kwargs, + ) + + hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) + gate = gate.unsqueeze(1) + hidden_states = gate * self.proj_out(hidden_states) + hidden_states = residual + hidden_states + if hidden_states.dtype == torch.float16: + hidden_states = hidden_states.clip(-65504, 65504) + + return hidden_states + + +@maybe_allow_in_graph +class FluxTransformerBlock(nn.Module): + def __init__( + self, dim: int, num_attention_heads: int, attention_head_dim: int, qk_norm: str = "rms_norm", eps: float = 1e-6 + ): + super().__init__() + + self.norm1 = AdaLayerNormZero(dim) + self.norm1_context = AdaLayerNormZero(dim) + + self.attn = Attention( + query_dim=dim, + cross_attention_dim=None, + added_kv_proj_dim=dim, + dim_head=attention_head_dim, + heads=num_attention_heads, + out_dim=dim, + context_pre_only=False, + bias=True, + processor=FluxAttnProcessor2_0(), + qk_norm=qk_norm, + eps=eps, + ) + + self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) + self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + + self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) + self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) + + norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( + encoder_hidden_states, emb=temb + ) + joint_attention_kwargs = joint_attention_kwargs or {} + # Attention. + attention_outputs = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + image_rotary_emb=image_rotary_emb, + **joint_attention_kwargs, + ) + + if len(attention_outputs) == 2: + attn_output, context_attn_output = attention_outputs + elif len(attention_outputs) == 3: + attn_output, context_attn_output, ip_attn_output = attention_outputs + + # Process attention outputs for the `hidden_states`. + attn_output = gate_msa.unsqueeze(1) * attn_output + hidden_states = hidden_states + attn_output + + norm_hidden_states = self.norm2(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + + ff_output = self.ff(norm_hidden_states) + ff_output = gate_mlp.unsqueeze(1) * ff_output + + hidden_states = hidden_states + ff_output + if len(attention_outputs) == 3: + hidden_states = hidden_states + ip_attn_output + + # Process attention outputs for the `encoder_hidden_states`. + + context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output + encoder_hidden_states = encoder_hidden_states + context_attn_output + + norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) + norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] + + context_ff_output = self.ff_context(norm_encoder_hidden_states) + encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output + if encoder_hidden_states.dtype == torch.float16: + encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) + + return encoder_hidden_states, hidden_states + + +class FluxTransformer2DModel( + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, FluxTransformer2DLoadersMixin, CacheMixin +): + """ + The Transformer model introduced in Flux. + + Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ + + Args: + patch_size (`int`, defaults to `1`): + Patch size to turn the input data into small patches. + in_channels (`int`, defaults to `64`): + The number of channels in the input. + out_channels (`int`, *optional*, defaults to `None`): + The number of channels in the output. If not specified, it defaults to `in_channels`. + num_layers (`int`, defaults to `19`): + The number of layers of dual stream DiT blocks to use. + num_single_layers (`int`, defaults to `38`): + The number of layers of single stream DiT blocks to use. + attention_head_dim (`int`, defaults to `128`): + The number of dimensions to use for each attention head. + num_attention_heads (`int`, defaults to `24`): + The number of attention heads to use. + joint_attention_dim (`int`, defaults to `4096`): + The number of dimensions to use for the joint attention (embedding/channel dimension of + `encoder_hidden_states`). + pooled_projection_dim (`int`, defaults to `768`): + The number of dimensions to use for the pooled projection. + guidance_embeds (`bool`, defaults to `False`): + Whether to use guidance embeddings for guidance-distilled variant of the model. + axes_dims_rope (`Tuple[int]`, defaults to `(16, 56, 56)`): + The dimensions to use for the rotary positional embeddings. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["FluxTransformerBlock", "FluxSingleTransformerBlock"] + _skip_layerwise_casting_patterns = ["pos_embed", "norm"] + + @register_to_config + def __init__( + self, + patch_size: int = 1, + in_channels: int = 64, + out_channels: Optional[int] = None, + num_layers: int = 19, + num_single_layers: int = 38, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + joint_attention_dim: int = 4096, + pooled_projection_dim: int = 768, + guidance_embeds: bool = False, + axes_dims_rope: Tuple[int, int, int] = (16, 56, 56), + ): + super().__init__() + self.out_channels = out_channels or in_channels + self.inner_dim = num_attention_heads * attention_head_dim + + self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope) + + text_time_guidance_cls = ( + CombinedTimestepGuidanceTextProjEmbeddings if guidance_embeds else CombinedTimestepTextProjEmbeddings + ) + self.time_text_embed = text_time_guidance_cls( + embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim + ) + + self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim) + self.x_embedder = nn.Linear(in_channels, self.inner_dim) + + self.transformer_blocks = nn.ModuleList( + [ + FluxTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + ) + for _ in range(num_layers) + ] + ) + + self.single_transformer_blocks = nn.ModuleList( + [ + FluxSingleTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + ) + for _ in range(num_single_layers) + ] + ) + + self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) + self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) + + self.gradient_checkpointing = False + + @property + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors + def attn_processors(self) -> Dict[str, AttentionProcessor]: + r""" + Returns: + `dict` of attention processors: A dictionary containing all attention processors used in the model with + indexed by its weight name. + """ + # set recursively + processors = {} + + def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]): + if hasattr(module, "get_processor"): + processors[f"{name}.processor"] = module.get_processor() + + for sub_name, child in module.named_children(): + fn_recursive_add_processors(f"{name}.{sub_name}", child, processors) + + return processors + + for name, module in self.named_children(): + fn_recursive_add_processors(name, module, processors) + + return processors + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor + def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]): + r""" + Sets the attention processor to use to compute attention. + + Parameters: + processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`): + The instantiated processor class or a dictionary of processor classes that will be set as the processor + for **all** `Attention` layers. + + If `processor` is a dict, the key needs to define the path to the corresponding cross attention + processor. This is strongly recommended when setting trainable attention processors. + + """ + count = len(self.attn_processors.keys()) + + if isinstance(processor, dict) and len(processor) != count: + raise ValueError( + f"A dict of processors was passed, but the number of processors {len(processor)} does not match the" + f" number of attention layers: {count}. Please make sure to pass {count} processor classes." + ) + + def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor): + if hasattr(module, "set_processor"): + if not isinstance(processor, dict): + module.set_processor(processor) + else: + module.set_processor(processor.pop(f"{name}.processor")) + + for sub_name, child in module.named_children(): + fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor) + + for name, module in self.named_children(): + fn_recursive_attn_processor(name, module, processor) + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedFluxAttnProcessor2_0 + def fuse_qkv_projections(self): + """ + Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value) + are fused. For cross-attention modules, key and value projection matrices are fused. + + + + This API is 🧪 experimental. + + + """ + self.original_attn_processors = None + + for _, attn_processor in self.attn_processors.items(): + if "Added" in str(attn_processor.__class__.__name__): + raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.") + + self.original_attn_processors = self.attn_processors + + for module in self.modules(): + if isinstance(module, Attention): + module.fuse_projections(fuse=True) + + self.set_attn_processor(FusedFluxAttnProcessor2_0()) + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections + def unfuse_qkv_projections(self): + """Disables the fused QKV projection if enabled. + + + + This API is 🧪 experimental. + + + + """ + if self.original_attn_processors is not None: + self.set_attn_processor(self.original_attn_processors) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor = None, + pooled_projections: torch.Tensor = None, + timestep: torch.LongTensor = None, + img_ids: torch.Tensor = None, + txt_ids: torch.Tensor = None, + guidance: torch.Tensor = None, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + controlnet_block_samples=None, + controlnet_single_block_samples=None, + return_dict: bool = True, + controlnet_blocks_repeat: bool = False, + ) -> Union[torch.Tensor, Transformer2DModelOutput]: + """ + The [`FluxTransformer2DModel`] forward method. + + Args: + hidden_states (`torch.Tensor` of shape `(batch_size, image_sequence_length, in_channels)`): + Input `hidden_states`. + encoder_hidden_states (`torch.Tensor` of shape `(batch_size, text_sequence_length, joint_attention_dim)`): + Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. + pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): Embeddings projected + from the embeddings of input conditions. + timestep ( `torch.LongTensor`): + Used to indicate denoising step. + block_controlnet_hidden_states: (`list` of `torch.Tensor`): + A list of tensors that if specified are added to the residuals of transformer blocks. + joint_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain + tuple. + + Returns: + If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. + """ + if joint_attention_kwargs is not None: + joint_attention_kwargs = joint_attention_kwargs.copy() + lora_scale = joint_attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + else: + if joint_attention_kwargs is not None and joint_attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` via `joint_attention_kwargs` when not using the PEFT backend is ineffective." + ) + + hidden_states = self.x_embedder(hidden_states) + + timestep = timestep.to(hidden_states.dtype) * 1000 + if guidance is not None: + guidance = guidance.to(hidden_states.dtype) * 1000 + + temb = ( + self.time_text_embed(timestep, pooled_projections) + if guidance is None + else self.time_text_embed(timestep, guidance, pooled_projections) + ) + encoder_hidden_states = self.context_embedder(encoder_hidden_states) + + if txt_ids.ndim == 3: + logger.warning( + "Passing `txt_ids` 3d torch.Tensor is deprecated." + "Please remove the batch dimension and pass it as a 2d torch Tensor" + ) + txt_ids = txt_ids[0] + if img_ids.ndim == 3: + logger.warning( + "Passing `img_ids` 3d torch.Tensor is deprecated." + "Please remove the batch dimension and pass it as a 2d torch Tensor" + ) + img_ids = img_ids[0] + + ids = torch.cat((txt_ids, img_ids), dim=0) + image_rotary_emb = self.pos_embed(ids) + + if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs: + ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds") + ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) + joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) + + for index_block, block in enumerate(self.transformer_blocks): + if torch.is_grad_enabled() and self.gradient_checkpointing: + encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( + block, + hidden_states, + encoder_hidden_states, + temb, + image_rotary_emb, + ) + + else: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=joint_attention_kwargs, + ) + + # controlnet residual + if controlnet_block_samples is not None: + interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) + interval_control = int(np.ceil(interval_control)) + # For Xlabs ControlNet. + if controlnet_blocks_repeat: + hidden_states = ( + hidden_states + controlnet_block_samples[index_block % len(controlnet_block_samples)] + ) + else: + hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] + hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) + + for index_block, block in enumerate(self.single_transformer_blocks): + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = self._gradient_checkpointing_func( + block, + hidden_states, + temb, + image_rotary_emb, + ) + + else: + hidden_states = block( + hidden_states=hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=joint_attention_kwargs, + ) + + # controlnet residual + if controlnet_single_block_samples is not None: + interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) + interval_control = int(np.ceil(interval_control)) + hidden_states[:, encoder_hidden_states.shape[1] :, ...] = ( + hidden_states[:, encoder_hidden_states.shape[1] :, ...] + + controlnet_single_block_samples[index_block // interval_control] + ) + + hidden_states = hidden_states[:, encoder_hidden_states.shape[1] :, ...] + + hidden_states = self.norm_out(hidden_states, temb) + output = self.proj_out(hidden_states) + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return (output,) + + return Transformer2DModelOutput(sample=output) diff --git a/pipelines/bria/transformer_bria.py b/pipelines/bria/transformer_bria.py new file mode 100644 index 000000000..f75cde912 --- /dev/null +++ b/pipelines/bria/transformer_bria.py @@ -0,0 +1,315 @@ +from typing import Any, Dict, List, Optional, Union +import numpy as np +import torch +import torch.nn as nn +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import PeftAdapterMixin, FromOriginalModelMixin +from diffusers.models.modeling_utils import ModelMixin +from diffusers.models.normalization import AdaLayerNormContinuous +from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers +from diffusers.models.modeling_outputs import Transformer2DModelOutput +from diffusers.models.embeddings import TimestepEmbedding, get_timestep_embedding +from pipelines.bria.transformer_block import FluxSingleTransformerBlock, FluxTransformerBlock +from pipelines.bria.bria_utils import FluxPosEmbed as EmbedND + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + +class Timesteps(nn.Module): + def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1,time_theta=10000): + super().__init__() + self.num_channels = num_channels + self.flip_sin_to_cos = flip_sin_to_cos + self.downscale_freq_shift = downscale_freq_shift + self.scale = scale + self.time_theta=time_theta + + def forward(self, timesteps): + t_emb = get_timestep_embedding( + timesteps, + self.num_channels, + flip_sin_to_cos=self.flip_sin_to_cos, + downscale_freq_shift=self.downscale_freq_shift, + scale=self.scale, + max_period=self.time_theta + ) + return t_emb + +class TimestepProjEmbeddings(nn.Module): + def __init__(self, embedding_dim, time_theta): + super().__init__() + + self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0,time_theta=time_theta) + self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) + + def forward(self, timestep, dtype): + timesteps_proj = self.time_proj(timestep) + timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=dtype)) # (N, D) + return timesteps_emb + +""" +Based on FluxPipeline with several changes: +- no pooled embeddings +- We use zero padding for prompts +- No guidance embedding since this is not a distilled version +""" +class BriaTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): + """ + The Transformer model introduced in Flux. + + Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ + + Parameters: + patch_size (`int`): Patch size to turn the input data into small patches. + in_channels (`int`, *optional*, defaults to 16): The number of channels in the input. + num_layers (`int`, *optional*, defaults to 18): The number of layers of MMDiT blocks to use. + num_single_layers (`int`, *optional*, defaults to 18): The number of layers of single DiT blocks to use. + attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. + num_attention_heads (`int`, *optional*, defaults to 18): The number of heads to use for multi-head attention. + joint_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. + pooled_projection_dim (`int`): Number of dimensions to use when projecting the `pooled_projections`. + guidance_embeds (`bool`, defaults to False): Whether to use guidance embeddings. + """ + + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + patch_size: int = 1, + in_channels: int = 64, + num_layers: int = 19, + num_single_layers: int = 38, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + joint_attention_dim: int = 4096, + pooled_projection_dim: int = None, + guidance_embeds: bool = False, + axes_dims_rope: List[int] = [16, 56, 56], + rope_theta = 10000, + time_theta = 10000 + ): + super().__init__() + self.out_channels = in_channels + self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim + + self.pos_embed = EmbedND(theta=rope_theta, axes_dim=axes_dims_rope) + + + self.time_embed = TimestepProjEmbeddings( + embedding_dim=self.inner_dim,time_theta=time_theta + ) + + # if pooled_projection_dim: + # self.pooled_text_embed = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim=self.inner_dim, act_fn="silu") + + if guidance_embeds: + self.guidance_embed = TimestepProjEmbeddings(embedding_dim=self.inner_dim) + + self.context_embedder = nn.Linear(self.config.joint_attention_dim, self.inner_dim) + self.x_embedder = torch.nn.Linear(self.config.in_channels, self.inner_dim) + + self.transformer_blocks = nn.ModuleList( + [ + FluxTransformerBlock( + dim=self.inner_dim, + num_attention_heads=self.config.num_attention_heads, + attention_head_dim=self.config.attention_head_dim, + ) + for i in range(self.config.num_layers) + ] + ) + + self.single_transformer_blocks = nn.ModuleList( + [ + FluxSingleTransformerBlock( + dim=self.inner_dim, + num_attention_heads=self.config.num_attention_heads, + attention_head_dim=self.config.attention_head_dim, + ) + for i in range(self.config.num_single_layers) + ] + ) + + self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) + self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) + + self.gradient_checkpointing = False + + def _set_gradient_checkpointing(self, module, value=False): + if hasattr(module, "gradient_checkpointing"): + module.gradient_checkpointing = value + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor = None, + pooled_projections: torch.Tensor = None, + timestep: torch.LongTensor = None, + img_ids: torch.Tensor = None, + txt_ids: torch.Tensor = None, + guidance: torch.Tensor = None, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + return_dict: bool = True, + controlnet_block_samples = None, + controlnet_single_block_samples=None, + + ) -> Union[torch.FloatTensor, Transformer2DModelOutput]: + """ + The [`FluxTransformer2DModel`] forward method. + + Args: + hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`): + Input `hidden_states`. + encoder_hidden_states (`torch.FloatTensor` of shape `(batch size, sequence_len, embed_dims)`): + Conditional embeddings (embeddings computed from the input conditions such as prompts) to use. + pooled_projections (`torch.FloatTensor` of shape `(batch_size, projection_dim)`): Embeddings projected + from the embeddings of input conditions. + timestep ( `torch.LongTensor`): + Used to indicate denoising step. + block_controlnet_hidden_states: (`list` of `torch.Tensor`): + A list of tensors that if specified are added to the residuals of transformer blocks. + joint_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain + tuple. + + Returns: + If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a + `tuple` where the first element is the sample tensor. + """ + if joint_attention_kwargs is not None: + joint_attention_kwargs = joint_attention_kwargs.copy() + lora_scale = joint_attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + else: + if joint_attention_kwargs is not None and joint_attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` via `joint_attention_kwargs` when not using the PEFT backend is ineffective." + ) + hidden_states = self.x_embedder(hidden_states) + + timestep = timestep.to(hidden_states.dtype) + if guidance is not None: + guidance = guidance.to(hidden_states.dtype) + else: + guidance = None + + # temb = ( + # self.time_text_embed(timestep, pooled_projections) + # if guidance is None + # else self.time_text_embed(timestep, guidance, pooled_projections) + # ) + + temb = self.time_embed(timestep,dtype=hidden_states.dtype) + + # if pooled_projections: + # temb+=self.pooled_text_embed(pooled_projections) + + if guidance: + temb+=self.guidance_embed(guidance,dtype=hidden_states.dtype) + + encoder_hidden_states = self.context_embedder(encoder_hidden_states) + + if len(txt_ids.shape)==2: + ids = torch.cat((txt_ids, img_ids), dim=0) + else: + ids = torch.cat((txt_ids, img_ids), dim=1) + image_rotary_emb = self.pos_embed(ids) + + for index_block, block in enumerate(self.transformer_blocks): + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module, return_dict=None): + def custom_forward(*inputs): + if return_dict is not None: + return module(*inputs, return_dict=return_dict) + else: + return module(*inputs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), + hidden_states, + encoder_hidden_states, + temb, + image_rotary_emb, + **ckpt_kwargs, + ) + + else: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + ) + + # controlnet residual + if controlnet_block_samples is not None: + interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) + interval_control = int(np.ceil(interval_control)) + hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control] + + + hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) + + for index_block, block in enumerate(self.single_transformer_blocks): + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module, return_dict=None): + def custom_forward(*inputs): + if return_dict is not None: + return module(*inputs, return_dict=return_dict) + else: + return module(*inputs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), + hidden_states, + temb, + image_rotary_emb, + **ckpt_kwargs, + ) + + else: + hidden_states = block( + hidden_states=hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + ) + + # controlnet residual + if controlnet_single_block_samples is not None: + interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) + interval_control = int(np.ceil(interval_control)) + hidden_states[:, encoder_hidden_states.shape[1] :, ...] = ( + hidden_states[:, encoder_hidden_states.shape[1] :, ...] + + controlnet_single_block_samples[index_block // interval_control] + ) + + hidden_states = hidden_states[:, encoder_hidden_states.shape[1] :, ...] + + hidden_states = self.norm_out(hidden_states, temb) + output = self.proj_out(hidden_states) + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return (output,) + + return Transformer2DModelOutput(sample=output) diff --git a/pipelines/model_bria.py b/pipelines/model_bria.py new file mode 100644 index 000000000..bcd458dc4 --- /dev/null +++ b/pipelines/model_bria.py @@ -0,0 +1,90 @@ +import os +import sys +import transformers +from modules import shared, devices, sd_models, model_quant, sd_hijack_te + + +def load_transformer(repo_id, diffusers_load_config={}): + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True) + fn = None + + if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default': + from modules import sd_unet + if shared.opts.sd_unet not in list(sd_unet.unet_dict): + shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}') + return None + fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None + + from pipelines.bria.transformer_bria import BriaTransformer2DModel + + if fn is not None and 'gguf' in fn.lower(): + shared.log.error('Load model: type=Bria format="gguf" unsupported') + transformer = None + elif fn is not None and 'safetensors' in fn.lower(): + shared.log.debug(f'Load model: type=Bria transformer="{fn}" quant="{model_quant.get_quant(repo_id)}" args={load_args}') + transformer = BriaTransformer2DModel.from_single_file( + fn, + cache_dir=shared.opts.hfcache_dir, + **load_args, + ) + else: + shared.log.debug(f'Load model: type=Bria transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}') + transformer = BriaTransformer2DModel.from_pretrained( + repo_id, + subfolder="transformer", + cache_dir=shared.opts.hfcache_dir, + **load_args, + **quant_args, + ) + if shared.opts.diffusers_offload_mode != 'none' and transformer is not None: + sd_models.move_model(transformer, devices.cpu) + return transformer + + +def load_text_encoder(repo_id, diffusers_load_config={}): + load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True) + shared.log.debug(f'Load model: type=Bria te="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}') + text_encoder = transformers.T5EncoderModel.from_pretrained( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.hfcache_dir, + **load_args, + **quant_args, + ) + if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None: + sd_models.move_model(text_encoder, devices.cpu) + return text_encoder + + +def load_bria(checkpoint_info, diffusers_load_config={}): + repo_id = sd_models.path_to_repo(checkpoint_info) + sd_models.hf_auth_check(checkpoint_info) + + transformer = load_transformer(repo_id, diffusers_load_config) + text_encoder = load_text_encoder(repo_id, diffusers_load_config) + + load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model') + shared.log.debug(f'Load model: type=Bria model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') + + from pipelines.bria.bria_pipeline import BriaPipeline + sys.path.append(os.path.join(os.path.dirname(__file__), 'bria')) + + pipe = BriaPipeline.from_pretrained( + repo_id, + transformer=transformer, + text_encoder=text_encoder, + cache_dir=shared.opts.diffusers_dir, + trust_remote_code=True, + **load_args, + ) + + del text_encoder + del transformer + + sd_hijack_te.init_hijack(pipe) + from modules.video_models import video_vae + pipe.vae.orig_decode = pipe.vae.decode + pipe.vae.decode = video_vae.hijack_vae_decode + + devices.torch_gc() + return pipe diff --git a/wiki b/wiki index 0004abde0..c134fe9e9 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 0004abde02d2d32b1a1c1518f6658ab80a2bde14 +Subproject commit c134fe9e92dc01caa3967ac5e1b1e018b3eb473d