From a5d3a681077bf9ddf5f935c915b26f5c2dbf31ff Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 13:20:17 -0400 Subject: [PATCH 001/122] add cogview4 Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 33 ++++-- html/reference.json | 8 +- .../Alpha-VLLM--Lumina-Image-2.0.jpg | Bin models/Reference/THUDM--CogView4-6B.jpg | Bin 0 -> 40384 bytes modules/model_cogview.py | 95 ++++++++++++++++++ modules/model_flux.py | 33 ------ modules/model_quant.py | 8 +- modules/modeldata.py | 4 + modules/sd_detect.py | 6 +- modules/sd_models.py | 8 +- modules/sd_samplers_common.py | 2 +- modules/shared_items.py | 3 +- wiki | 2 +- 13 files changed, 148 insertions(+), 54 deletions(-) mode change 100755 => 100644 models/Reference/Alpha-VLLM--Lumina-Image-2.0.jpg create mode 100644 models/Reference/THUDM--CogView4-6B.jpg create mode 100644 modules/model_cogview.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 32d906a7d..f0b8b9c52 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,15 +1,30 @@ # Change Log for SD.Next -## Update for 2025-03-14 +## Update for 2025-03-15 -- fix installer not starting when older version of rich is installed -- fix circular imports when debug flags are enabled -- fix cuda errors with directml -- fix memory stats not displaying the ram usage -- fix runpod memory limit reporting -- fix remote vae not being stored in metadata, thanks @iDeNoh -- add --upgrade to torch_command when using --use-nightly for ipex and rocm -- **ipex** +- **Models** + - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) + new foundation model for image generation based o T5-XXL text encoder and a flow-based diffusion transformer + fully supports offloading and on-the-fly quantization + simply select from *networks -> models -> reference* + - New [zer0int CLiP-L](https://huggingface.co/zer0int/CLIP-Registers-Gated_MLP-ViT-L-14) models: + download text encoders into folder set in settings -> system paths -> text encoders (default is `models/Text-encoder`) + load using *settings -> text encoder* + *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui +- **Wiki/Docs** + - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info + - Updated SD3 +- **Other** + - add remote vae info to metadata, thanks @iDeNoh + - add quantization support to **CogView-3Plus** +- **Fixes** + - fix installer not starting when older version of `rich` is installed + - fix circular imports when debug flags are enabled + - fix cuda errors with *directml* + - fix memory stats not displaying the ram usage + - fix **RunPod** memory limit reporting +- **IPEX** + - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler - fix untyped_storage, torch.eye and torch.cuda.device ops - fix torch 2.7 compatibility diff --git a/html/reference.json b/html/reference.json index 4cc2edb28..7b6589215 100644 --- a/html/reference.json +++ b/html/reference.json @@ -376,9 +376,15 @@ "extras": "sampler: DPM++ 2M EDM" }, + "CogView 4": { + "path": "THUDM/CogView4-6B", + "desc": "An innovative cascaded framework that enhances the performance of text-to-image diffusion. CogView is the first model implementing relay diffusion in the realm of text-to-image generation, executing the task by first creating low-resolution images and subsequently applying relay-based super-resolution.", + "preview": "THUDM--CogView4-6B.jpg", + "skip": true + }, "CogView 3 Plus": { "path": "THUDM/CogView3-Plus-3B", - "desc": "This model is the DiT version of CogView3, a text-to-image generation model, supporting image generation from 512 to 2048px. Resolution: Width and height must meet the range from 512px to 2048px and must be divisible by 32.", + "desc": "An innovative cascaded framework that enhances the performance of text-to-image diffusion. CogView is the first model implementing relay diffusion in the realm of text-to-image generation, executing the task by first creating low-resolution images and subsequently applying relay-based super-resolution.", "preview": "THUDM--CogView3-Plus-3B.jpg", "skip": true }, diff --git a/models/Reference/Alpha-VLLM--Lumina-Image-2.0.jpg b/models/Reference/Alpha-VLLM--Lumina-Image-2.0.jpg old mode 100755 new mode 100644 diff --git a/models/Reference/THUDM--CogView4-6B.jpg b/models/Reference/THUDM--CogView4-6B.jpg new file mode 100644 index 0000000000000000000000000000000000000000..5876b935ba4fff7da75e4df471b58ef5b4aeb2b0 GIT binary patch literal 40384 zcmbTdbyOWq5HEOfcL`2_pce@465QP{lHd-(EqJhtySux)BoN%)-Q6!P%lF>f_s-eB zcDLu$Om%hlbpLv)r+cb;-WT6D0hqFqGLir&C@8>>j|1?&1+b8Cw=@F)6chjq0000H z01HI~fcxN}K1c|P_y8{~?F|NMr?koB>!Lr1(Mqg&!aEUzz`L{b^`v=ge$p z>in6LnT45^{iCK8fEWN70Ra&K9vKl45d|3;1)TsB9Sseg6b~PZfQp>@3ne)v1r3n< zD-AsdJq0C;2rCCKzmSj+HM6*kn1D35pb-CmH-SP%K|x1DC&9!d;isdd)4{AS9xu zp{1i|;Naxq=HcZN6PJ*bl9rKGQ`gYc($)bPo0yuJTUc5-IlH*JxqEm9{tgNb2@M0s z#U~^tC8wmO<>eO?78RG2metiaG&VK2w6^v9>+S0w7#td&o|&DSUszmP-rC;T-P=Dn zJUYI*zPW|m-9J1&{f7(c1Lyx8{wJ{i2QI7+T+pzvFt7;!;evv8`yd!BSU3t+cx(|B z1VaZLN;ZE)T+!Iv+HNE&cGXKfBgZLZ{4X3^)K~vO`ya^u-vJBw|Ap*-1N%R?mI0_R zP#+Hu1`8kzxE`X3HT^*hAd*RDYo1R*Iv9Dqc~m#Sl1CDF^zHIx0Q-Hh4hy-;9?r8% zpkVvnAQ$~1jBcce^UnKz@s>ymkozd{4w%j_U>$Q{5#W6ZV)Og*EWWWJoG3|;8h?i4 z`|C}87`i7sHOklBmqOm)c9Bx4LE#-Bh_TBcpxScEvpVrVBW!dn6|#1PFc;awIRj&g zJ&a7f0}m5>@*>JhGT6kaBK`l>VtP=?St5w!GevuUT>Aw^aa2qEYE5=r@}tdM`HYrI z;ymT`qTEii`W%jEnx&UR1cb7gEwR$3xv(k26*A(k9Y7zrLJcN_uu9IYvj%2`f#oS(y_)Y3RH zoUKM*`>jM8J(&b0a!U3>MCsP0#K4~MIpO|Fzbt5ld6|igW^S{j+3!KbF z@2MdrT=qn#yh*4r?i7Sv<5^>I%d41ul7Me*CGYp8&|F?6OR!=>Hag#V^CTZskJH6k z4v+dLG$;IWD;!fquw;RPyx<1&2Kq=d@tpiU`3`8(^PNz8L=FWTQ9qjwyg-CwN)%YA zpKbpqE1iSXwc_yc)c|0P-VcAhHpxXvVZp5$Fp)M-K)n!7l);N=hW+VrO7B$I z4vg)Lg-vEh?8XQ@(Z_a|8=YR>EVqZQe_F$1)=jEO1hr03uhhN1Sy?)VMCW{-&zisf_!%R;V{-bf(4d&J;XQ$;>eFyw{(}2Q-qHj-sKDfq1r^Dh~lOKiuXk8R#o=Chr ziUR-ZE4=Ki=r{1SVo_4+&gG4Jw23&Ha#KL?#cnG*#yF!fwIqs*_dzntGqK`P?hH1vs! z*OrA2N+MQ*ik3Os@nWrfo9ZIRHvIAD1F|75_|QKgh*r$_j0h$E#iRJN8!GVmDS8Fz z<5uI*7HLuCa1FmJXuqUVGToW>iJ2GfNpl5EVA_N6woKH8|x^`=ki&c@fcijK7d??^yL|9zL8mE>eg1^TR(rNBz{XvqPt5( zEmmIE-ylYBrAFO43jKsH{@Xl>u;x>t2-hbXig=FF%3kr{EHi$}U;alA8uokwS9Gry z_w5MntBqZ`nyc)yo>3P2hQndO;wQ5j*WE!d?1<$P*xXjkKwLD%yIt=Y#Vw;w-q_@l zB8r9_Wr~JYlGgmOoxs%F)ZLKy8zj%d!W#?&dPV=SF&{ zbVu4mN?^t0pI-nq=Wreb8C;ZdtAZeRFwoV@*?^>mh-NyQYxKODz4f7~XJ#mP+8||! zihR5S641-D5^r$N=AtF|TJ+n1zP7oheT9rwFSvLc7zt86Exhbz3G;og8Krrc=gYV^tXEOWs(0*lvHkzwE!umN7oqZ zL^x1!5qUEiK4%uXVcBx-VoD8Nm*~OCnpeX`J&YmBzP~lHq2B?6A-TQ_zlEO=i)e92 z$HcZmLcZ!`0nHj>Zrb4KCnmLw!ET=z1j^n4TTG`O)@WLzq)G1Guadwy3|u|EF7u#u zcbwj))Ug6?Wd8mK;aI%0i@;fxT;`L~o6PlSZ+M@%LOmM3?5}Ru)j5R$gz#jX!a@9$ z9?L=&Y11n_)OkV?h;J^pq#Gg5?H!LvgHCT_B)MnbjSqMk@2=y3$Y(E7TVB>on%CoW z(2s$mVNY-n@`L&Ku)MO=zbV8soQIm^Xhq2~U`unC)PCtruwXidEzWp^!Q9H~1B!}Y zW)O*20NrlPYXv&)fIi0?-s}!7SJ}pSBeM?P8O7?}9hNS;MRKk8Gq2R7Sz>{Z#XW0{ zgkId#jxZmeZ%2XZY$dV|CvW&HYlcfEw}5aD9nW8$9$2oUw-UzOkcqmMN|QN4C22uq@f}|u4yncjO)x{3)bdr zZP7h+c&NI6sv8ctVdIObMGIKIkdI`V%gtrGtOttPV6uUK1Irro$Q3IoO>M+DbfKNx z@EtH;nvf6Yfmgcli+Ku1t@BS?2eFL-d@a*zM+*tLvT#aO&oRFeT6dv^E8+DS@kDI2 zR==du^MJz>@(mTJiot%d+EwtWt8b`3PEkp1NI;^{;K5h7eWu#)CfB|L9;Vui9CSMo z>e*1Uu`=Ff`#ZiraHu1RZF*qUSb)TZ)B!ZvM_;aMtfOkb&SQ`7b9^_v>e4vFlxHfa z^4o*K``Y7a-g#@!EJ@ScjXE`DH_kI&_h#*w7CzsYBBgb8;LCkstSj$#*;CD?5UNfH zp#&?yOfAu70cAO)sO`dsnX6B&2DP-!a_HzQgsj5fc2LsiMXHmfgw99%ft2?+^N$93y9I zRE&5%ZVGTg-cD{*j}&Z%HGbSnrXMZvtYLfX5WX(%`jA?^C6|MgD1qEddxtP&7*lQL zC*WbVO0+3U)uhpXYeS^n-K2eJ^OxW+etiitvT=)rN46Mph^3Su)xs0y!NuSH7Y-A_0Dh032^mzWka&^DsA)I>GKVafDI}UA2WUH?za1;zpD9RZ8T9gy|5F(u?seo-YwYGn`8r(SRwCre z!$JMSZb$t_^$S^1x&^jjT7qaJw}eiFSTG)2zaLnO^mlg1L?xfY^MsA^Z^@10q>P<= zYeH+$&?mYtt(P*%fk18{hN_xx$RO*MIwHxVh)J@%UmtX!dg0VwYiioarkTnj{bB zMzN~V73rBRsKwV^QYVfT9X$Wt`Mlk>VgXm-g;y^caEe~go}oVxpgDkn!g@!`xmp)u z;IX+NXyw3HqLJ*>$64-MZ-qwQsCV@fMlYT$F|xwNU|*uMqS^6LxoBUo+`4=5ZIKsn z36xb)^-P&h|Lx4aJ6WCkJU{W^WYGAwFHy$K;@frt-Eb-X!kc-VF(wD)Y`yO}#tlcwXxcG7z)NWwEV_$5QZA|%Ai{sCl zb}J)5Y_hWl@J1u+7k%o*;9Hb%V|yUPcxpC$m^v-cYWDBhT+rbTVAeVp^2sY+8mW1$ z;Lq3#_>_$|sr%r1plFu#1NoJK8o!xd@vE7kfzS{#@?yJwvS72dGXmPRN2!?uPxtNqY}U54mzjRSxkbN=;~i+@fi-B z$p8rY679H<+MV>e4KbGPu82C=+F)Qsr3zI>kK~*yCe&~qcl{D;K6sPf&vciP#^1@Z zl~YT)-UQraoRMNWBMSvW7tHLeFkgtI-O|G~G0pUlgZPD&rwVWZFY`!xPjgYB$t2yDEUav@QR;5Df-X;qKtsKAptIb{f1 zW=V+a>qO!$SGR$fMRm@HRs&N_&8di^m1M0+v_$JPrTE+3aBxeLW0hTu{S~)sY|hrQPVLW~^@%UZW~w2u{SW)xWGP20Bu`r8 zP1PdT6IXA4Hi?d?o!yi6sHid)jCZP0B=QZ43X8?w0oW_^JQCf{0-3T7ixBeULQnN_ z(Xz>;3I|8h5kH3QKH%$!xqVM_{i&RFs{-36g7WQgNY`2Z8^;-Q&wjLH^thL1=rn}p z$Tqq(KS@7m0)5|S3BzheA4ko1Q$;DxJLqHD9=%?LaMpAY&3-S3dQiJlJN;^_)Ed2; zwci}y^U40c(EhL59Vb7^%FOwXP(NYM#*LolT78D5?cb36#fcg(Pw!zU9~I5>Gqj8@ zl~W*+MBh>-W?Wnwi~Yj=h}v3@iKCR&Jc=0oxfHGg+zE(?CiE#FnRmH3RAU|$fUk0ntxgLRDZ5CzhFE9gLAK( zom2fo32~)loAprzV!HZW_gdh$7(Yn`w(kZ_XV5Qp zjkR9tNPya8yg@m7R*iVAxF6kJxjNNsNLaGi=s=JQ`f2_p=P|hzWC|dpAlBnwrGhmvgvNK z1Fe@IFZqh@;)dVlyYZ`<$)R%#z4I&D!JiLs6*VB!zes3!RH z-&`Q`5hWM@Us3Hu+R{n8BC~Mhop5!NDMKf8CY7XOLp`aYc=;sp{Wvl8-&!$dv5_j& zmbkN(f$|cx1f|pF&ynspMz^Q_tZ|NM zpZ?2|*iPQ4Xf1gW$G9txUvG_XU7ZUT!gvS7{Mn?t^m_+{{ds}ay!D|<$GXf)=nvE{ z&S05l%@Na{CIVP;1gBuA$s^!uqae>AS}}fNJjVlAiP5_Ue`~YOa1~@Y{mC@;+7jy~ zfyFK@C>rq&U_sc?(;fyV3R#kBTeQj!2JjR5jF{)Z?B&s=J@#`wrmW&V1(A9mV-P6N zByQ=1G2L|n#6xLbQl^q(O##T>&5TVvtb+HIgts_;vF*-Nd0y3XimX1*gt<>yx#mmmDe10 zWBx0)8}%=NKhF*PCR;x5X6#IHlgB}(Mc-0=8OMqg`so5}1#3skNv%wuXFL+406+5b zk=a_Ho*Xd{YA%y9+6uxmB;@D0L!#CX(v-X6qmW=9K5Hm{KL63=SS0Zq|4}5i>lpQ3 zDDl|7>}ycPCfKF9MbqLPFquYLHd^^-oY7+y{ihTH%zXXR?(PQVAIQB=;WMtses=zl z>9%@93ggX6<`{N>qv!$PDGJy!Wna;-foIjlIIE|-*un&GrcTo&@8O5gURBa>>lS5n z(N1d6gFrSO7kBdah*gDP1Xg&I*ckSSg|wq&(52XFsy$O<`aDznXxukkR$f|*&lov) z#SJBQ;>&cq@ac}LuZiU1oAy?3u~wVD6oC>Gj5X_Ol?(sP+4>!Jsh1Y`~JX37DBE)#v-O zaG@QF^vz(Q!{TsbGu!Ydyjh*4l`3KwPG(p7Z)0s!2`j;iNPW&Y0zL-oe418vd)N6r zzf_=_=48ob+jlnMTiSiwti!Rcva&vH@|q~G+$cr+%@uzAt>Ht}@ZeqaTGTJr{CU@> zvGMXRQ~^#@!8?HOIIuwxxoz;s=CrTDQ2Gc23vX4Ps$E@U{#j$Mfw2|3I2CVr_!yo~ zeIL40;u8uUIO4Wk&Q9TjlZ7W`cv|saPG)(OPi4H%N~}>)ch1ky3npI_h2Z(znFP}~ zO5w>Wg!K{bFG*SbofXis$Cl3fMre8Sg;~>j#I9@ZApYTatC^Bkt$R$V2@G74Nqmt5 zehuHWNU@8QsX!S?Wp4b=Cw7CSH(l;=Lsi9bhQj8p`6d?$;*Z$#8)#HXe|0v?9cy-+ zRgH}Hxzye1=A(Ue@B43U9nwAa4!|`43m*-CeTCHwvoPeoV9&#zQ++5lVg|8ku+!&t zI}Ul4n0~xs>rC+Ehb|Pm9X-f7wTBdYS}LsAW_y7|hBlMXWR`4jEHH>1gmloK{9KFc zCn@v$z=?}QZdfuicvLf3dX>CwGf<18s9iaex+bXWehZwoQKPK8UBC#b3btppUn*ZS zE7-lXD5Y}vdu-rq-vLl~pMqudO??ul#I|{Bg^`XT<_HC}JocchBGLe^mI;4dxH$j~ z%1{!}pEwdw>uudgGn^)Xa0KE*q);~5ukbaT-CwqtsA4E$qNE20k6uoOM(O0REU0F> zWpOX)39~UFEngcWxex)E%gb>9m9vZPn?I&?_X)XiE0ay`W6h%@!dWo!DMe>tn3h7s=EtQHDUf$~%4V8X zOAs9c$6D>5y=jrf-LghxP;<6De zlGuk?zJpNvN(xgzh{{dkx%FTrHz@B8s5E+BDC-cpJ8rE%eY^Og0d2td&#Hi4MtK!%6#OLl*Q&lFkRlv0HPkc3skO;h5Q0twmUjgfmRZNHgcHF2%ImhSq{ z0hAdM92M8*6PVv{DgKbG+u}=~+Rv+`a`yMl`AjcTBb;;Yv|hJHJAd|9$^2KX6>_ww zfGXQoC+=4t^3npP#eI%a*n9SFw%=_M7V|qKLOhQlqSkVnGxlFksp5K2$vH5c8bir@ zy?OVeTm;Fo7mUtEs`hQBCQupf8S4_g<5C_6PuF|Zqlr)+LDMRCiw*~UQzoQis-EB9 z5LVHg=`EK9h?*M7(wKdQZ+MA!RTp`VqCH(^xRzU$6hqHqctzfR_ZG%3o_ZY1iNx=hz2W*BJB;@N88%F^~~ZJ^rLB%WZ2M zn>I8nXRIO9IcBTdzn~>CO@?_(Q>!JNxIr>+SCDEIC+S>21dKbE5%j2}%4g%|PQql zq~(EYTI*61{EvW31?$6`0iM2OJNQ3lj+xZdHp}*CFP-GF`Srh^;9sZ75+sjcrSr6n z>*VQde%i&yjQd#n!{9b7qNrJ?nmZwLF2e+0sH&LW;zqKmI}F2pCk50TliUwM!4H3> zWucxsGIM5?fy=Zbg$DC7uwZdaXW+TxE_ZHK5B;-<7rv9%_sIoGCS#Szyl9KB zCgY}oQ8^1@)ENpdl-Uj9Z8rPqZiRN80iGI8c^-rlEcUX?aoyhZH@S}NQ2w`-<~4W4 zqd~)?n(&R0ez*Cl2td|dRof+RkxG;iL%ez}exm;Rd&LaCh6?xQ(5-8Tq}q|%Lk-4R zlyJMAOjeqNFWT1^MVt?rf|zw};Hu^=d54oxa*HRDn17!wF)jnonL!AfK$fPVhe=jq z;~gGCA;}=)XhD*vStmQJXX?Y6>|MqDe5OWI>p2HNMcWIZ19;?9ah-UfOOukK)G@ zp{zt<^Sq>Y`dkk*7lL>*%^}G9vdg8oin}69tlI7+0j{}W2M;~Yv~q)u1Bgey7XcYc z)dRzitrf0EQYoF_&K9vjpvJkrKdF&hl{<_XutP%=R2WgBS^X83W;~5>36^f3a6S|< zdWjYqkK~T2o9rJg=iYd_noN6E04sqX-TD_KhFI$eZ4|9I^M0jd{}IA6_Vvu=dRwLW4iOy!Vxv zbc++25KiOEVLj}6F+k+aet*b;@nW2!`;*2VyYUF?01Z4rzSGIz4*C7Wj!v5`-e?3; zL#CHVSCCDX$0@gS%&(HpZsGZg zW93XzcZE1*v94|ex#n~2n|fWJi!D>;TAGAUy69tiNXoEq+fMc5!I93jTymnE-QEpw zEgu`G8_7IYaqb_bQ`BpGVe0jUC)8dt;E?`sk5-r$`9+Uq0)cAQsFGghc|^q`DV_9t zh!6#6tX^@BN>@v1ZGh!DGAi}9nk$PnB^)hJ<6?2a1~^2b_D|$2_kwxTPXv}sg!E** z17q8~Dbra&Qqj=18}xn_9DvtJ5%f(qAXWjOt<2EaHyA+KJKKp2N|qnO93jyn_bTtGP#1zTI_@Lx zk4Rl8JZ0MjqHed)+0zDy9_|J%rpcBatK~frzJ1ky>gL^KGCyG-adFg-=IL|L%rIr2 z{@0`8AKUK4-3Y~nUrE2bjH7x#3n33tibV5C{Pb@p63@C0WlbtoL!Jx%0lqCVss$kVR()Rg+f7W)IkLct#9a7x-^y@VnE~bs09D4tIh9^ zSR7xUE<@e{B6!K0XRgyL-JTrG4;H6+s)9Y>`Oy%bVl!8l&|h*U!cpO@UefnmC#Vl; z6VJCgAq~+^jy1kpbPaD1eDA>Jf!6l4ZXx=Q7)f5Y=uuz;lf1)tdK2)9aJRz^kwcr4M;%pQC-buCQLW@|Y)5_cR=+WARSMYL&qkxq2A4_e#( zvn0LCQLW2Hv6L7cK76s)W!2-WS&NVg8P3JYQ?>X>-ynr5{b@deH>YI7Mk%T~@KDQa zK)q0D0bRp8U;CRVa z!2-9~03;f%B6;1ra`bmZj82`$eYI%ZibBT%^`!Z9VNJ`>HLZkg9ES1aY}>sSm^%V{ zD#MwrL(^;jM|?j!6*Lqbd;El+3)R&C$Sw4RymZORl7ri3s$UW}s`$@7PiYjTutl>{ z?WbiJM}6_yotSU1WKkhCcIirlu&5+BqG)Vh@c4f?!0BnYK9R(ZQ#U7z?*K2nymI{0 zRH1QkGY?MS5Mv`hD}DIBa{ku{@7;-Yf&5>%`KbuYJ;m%dszc% zYT2Us=oK~!ZgBr(+CJ5o_T&XN$$Pm9mz11CF!4=3H+0Knkp+L_tGP9NdUFAXC-+z^ zN9_@hlvLSspKlOdx{+tGA9}Ieaz--Mhs9ZYpw`MSY7Ow7SMU=Z{?uzSlT}LjQ(APQ zf6S(Sz-g}D#A+#3^7ml9^+?8PiCEp>yhywtB`?E8;HSq#r2;DT?Sr;d2f|QyKhUdW zC?Dt~!DcJm3+noCvj<#D3d&0NCH~yJ_^e$74h{1kP2rlJX) zW<2*aj2<3XTu6j!F%8z5Z_?OpW=SH#u*LAiS>b_Wru^YepcTY0L=+j%#SiFJdm28K zn%!_(2Oz|Vko9+nu7GAM3DT{fgQ)V7CvM8MgXS8Co=75 zlKa8UzT`2Yf2?-b`NO(iBCH9UsjW#dBaAK_3BMBif|O48jYo(^mg?;-!Zy1IcW1?9 zsWQ)GW{mG*DO-D^wAHeN8P5b3Qnl7nPgLst6S>hhWSNEM8(madGn2J#-TCBnrt zjs3xYFW>TvrUVL);elQ0_eu;E-glnR$1JS3;kV3oxqKy>5)-Y9!Uh@wMTkcm6bA1NKuS8>+!B6#*RtV6=7AWFs33EgB&vg zfB0jw8pr5lZB8QcVHKMXG1kkc+#N>@%@f0(W}o4;nnDlFC7YSD9CY09<3Y*N)Wfdu zxl?iU>fu0{lI+H_3Brma#&dzS0m3n1vw`ubE_|(KSfanpt!-h@84k)v(!Z|GZLAPX z<7;j!!o!Ew5wxCLN)m`biG(z_q+4|VysOO46;!9}X&U`|&N}P~bL^5oM?DJimJ|v+h5S*tgKy_yD(lFntKIcVFRK!ecf+OM0k>HCSe{f- zdG{Yf>2uj$5GaW@w(Wmf+y8jyZoaI4H+_8is!*In3(ItZD>xJUPsoJhc95 zh(9gx21TFMw|I}f=|y(y(i4#Uv}{ZFaOL;3WekkWR zI*T~=4{*xd$nFNgX^~k9GDMu~g=5htKeVEp9s3qaim||C0086%>q$HNwM=>Noe{Nv z)mq|L@QC7m$F3R$H^2e;lLYQ!3tNnXe<&*196;v|Ju?exn?v3KM4Iz`PW*aNB1vEF z_hYZW2`!dU)OoW_8XPw@ZQRUjsNGI*+V|YMrrCAkL7rio{YjTr@*L({Yt-i0#`ZZIn|Cprsn)LF;F?1jC@)V9CTb%;MU zDT-}m%Z8~{a$i`0GN>*S1+?6Q=7$+c$Q-m9vQcEnR$i6gF0V@lis#GNisj5*<Ers~Aa7u`Zb8;etd!7v2ZB7+ zeQ*YQh25vF2TM+5d=spGkkW+-F)o<#9-s*__hEJ&O|)aOmkb3WHnZh2O6 zwp(jV7?&%bg%g>>3Q?+2kLr9s{J2MoTy?DFKCV{2cHfR7AKn3eb388`Y#zd?Ws}3# zw4eM}abs@I)N~Jx7#)52IMAeo$5RE~5#;o@@J>|%Faa1FCkdLs3O_D3^tYQ zS0Nfv?d6!cbheVSi(C=pbx$zxb; z`5!kKu5dAbY;j}YA;u}aHPn@op)nQqWFiY8sj%Jod#2Rb){O6v?+nVs&h{^waO=Bn^)@~!AOvs40t*VkQP#qDTB5re_e*>);i`j}~QhKtfVW~LwZmB~2r zyLB9!HMfC$%hr!(>U6E>Vr7D-m zz;pljSTpnRaNejE(sJ#-DoLekZ(ZElW4?S{NP3pNvFB(tVg4LZixBAdfI$A;3WDUK zOOCXtt!>R6NgoqxoTf44h(I{txug|;Z`NA5W_=G{`vWufX?wpeW`wNJ0^f3jhrZS* zdJ<+xsQAdBwBpE<3DOq+$_GqbWkEHJvNjDXNGycAPrc^4H4k0-{*WY$x@P5&3RjZs zxl8stBleWtAi3nErh7+@997qH$vW%68;ALPKT4ev)8G!e*Tltd_l5={cxk|(U z_2#O~@h1}F>4raFT8&$}s9l}d6z}<>W3Bos&m>9T0bSnooc+zy&WoG1!>{@K8>`Wk zws2zwT4Qk=%yn#N4WUzENNTJ`J2Rcaqj8e;ne`2lZq3t;?en`j7V5u-Xdb#@v(`$N zo(0F1R(YvNyDK%92|FP)+V~zU1m(1htn*1+CT*%YG4Fuk-1TMLN^glzXioe60Q5jQ zhl%OX6f|Xq2&X_Wku0|~j$6Hn#oiB5b)T#8O$j9jx(qqu&pH$~Rs3MflQNtYQEc03 zBaQD-{|afKGoORW$_s;R?WEX3dmiOI@E9s2NvE6E%Nh(uxoM_3c{wU2%y?O9S3ciC z%CM>SLiD^h`Jdw|pfa}F4oG4F;R932$^uXy{bZpKEH9dAayMtepxp+9}uUF2)e6xHJS^971 zd2N>b;`?RQRKEWfhS2`R({)2Cv@a(8uQjiF75TIevhXt?Hs_ia+76$qSna*JG5V`d z9N*NwmWVw@jh+fp;|-!*3r7y@NG%AreJ%OA&o{53tUZW`HynO%^5cusfTyqS9}xj$ zeFcACj%(qiNnAywD#ztJ3z#5BD1q5$^Z7k*edpR)(6Ke}(Gft@dY|%|KB4)T+^Nor zCMc6d^wqc|HeCRD|MPqmp`@N>G!Cv=jjFd|A`H$TX8>nv$Q2Vg$nULC=s;Wx+#Rix z7{&8~CS3U~IWXl+q+<+3@D6AnQ}Nih-e@p^7PKa}{3K*%dzw}}9O*5(TbRGQX^a=- zTh=(F-R9@fry<)hJea0VqQ6vZ_*}{Eu*uD-aWVSjL|OGWA2(CF#4-1op7%!jT!`dE zi8<5B*oiH>ZV(@jHYSBc6E|N`A4M){@-C? zZwnT9T-#>6{-k%x>-@d6&{OX})aL;i=pBjg)%@!Xina)@xjuJu`>8}LN!c;%8%k8s zx{M{vF1P%Maa;)zBqC-dXOo#|t)@@zOY2>;`%0{Lp&gBoahEFAA_)b z=+|ggD@SL=0K`hFaBO>xfz|JLL(zJS5DbA&Zb+34qMp2Ern+U;H3*mKI{T=Nx@x9neAbOk z%TP+wyjs8b6in6RfwzzZt-g8W?P;4QUntS0F%OtHKfz@$o&8!zvIiy0B>L7bz?wj< zCoZKk$yf`7x1z76Wyk=(xney|rqZ&y6tf$!bELOis156w$M8FPA73n@E^R`=xR%nu zY9|SGsI`D65qDBpT}1h?*}(z8$at8Dj2SWWjw4kAYXd--a~|AY1=eel`6fsBoYGz( zq~apjvw|C>wPaU_`5L7rV5;rC`s3iHnjK4~r-r}6T&f0A`KHk`md^ks+I~$VXl&vX7h(>qq$kDMOr%tCCL#1y52bz zBKHT^s@jVXJU%DxTDtTA+bEk?Naj3z;NWVWX_bAM@(?^dDQ7c3is<%6T7m9M^@g{6 z!Z+5H5@A;?9#4p}(fr}>Y%}JIw8XndSyOU~yCM&-{>AEGkDJtGRp+?I=FP5)>1+NF zJZDC1_c;NcV&VMR+sc?4jTSp?e<_s+?yNaH_gNp3m$GFf9b;k#!J7;T6a%lXlL~2q z-TEIAY~`IN3dIXYFYru2nwo1FYF$x;6n`x0ub`(kVLbG%}>y167^WbVD}k3Hz#x3{zWg%TAmCA0wGK+c$sxCgdHcD zqatEVLZLDN36TjO(fF7?xh+<122pvEc5vw9w-~mEv=)VdErzdSRvEv#9fe`u0So1y zf?<<*a#aPo5bvy}**cdgWtU;yDq^V0BBBrrqcFHo#s6Y@6g=9OFk4$FvG+6;J#nm1 zMVL<)1f5k)Rp`d$S652)Ko&da(-WrYtL3>_qtnkjd8>%lpMYL!*Yw?FjYUm%{8)vj z_n);F?z5~hW3$`W3zK%vKH-Q^mIY4v?HXmS#hWbIfbZED*S{wfpr>2rc8w0ppgYCh zmMN~lXsyHoV%iaNPKHuEmz;lr0NAh$!sT*ZGM=T zfS)ZM%DFq3)^i%QzidW?=Holo%e_lvt_<;fuc^_?l+@Zb%8MnlmrF>QmuWuNl-@Xb zFOu{d7+00989pJBQX7+*Jzlk#vd*r*j6SG?pQOtZIZ=eA#Q>=&-9+Jw`j0a#GpA%8 zTOUqywCkkMj-A-81@>WKu+#rNKpAmOkD8zgBb}$~6z=DJ%%gkXlaLe`ejh1!HzW%_ zyjamulWZh};>GpQEJ0iS8jNbJiMG%A_>dllo<0m%x}Zya%n! zxbZU~E}LvNLmb0hw$q*@l(K}J_ZW8DWROkKZDolIvElu6v`a`w=u|393IOoa)ox zCWf9N4L3)Dh0N+gW*RlkSoCPDFej>w>lu8iHI4WhvFf2!+UD!IwxZph$c%8#GrMU< z)8mU57fRlsMT|rRD5J(x{{&r0^tp3-a`o^BqoN9YQ;XYsTHb;ALU7d8_#P*;k2g?(WbFKBt{}}wp$YeMcB8k`w@|jWgt~Xk4pk2>%yFd_6b?8&+}jvuRM?Mk`4y^ zQ!bHGz41~+k1?32Hi}$N;9H~}u==x%VXbHUa;cM)^eM8;MQi+K=ZCwM-q3vYVBh%A zb3p^u2q->4oc!6jB$4Lz2xIwHD*540A760dHZO&at2ic|9(Dy9d*2(vEQD0&d@GZP zOW5V+`jx(?keXP!OQuuVUGZ6YUTF0z3a@<&KRUnWKtO}n8QbbVCpOvi&^-^l3C#+< zWQvL(v)(&F;R>xYE`;Y0x@B0^Z;YM`Qdgv5K`ONJ)f@7tu_yW}`<3KgV$7xpxkCy? zA}-F%+X;_~ZkU!oHfSP#Ra_wXgTsf}cDsV-!}aFr^#+<%Vs_zTg~TqDmrGLx>V6vA zRQamtt|aQDeotM)Sxo(P1nsAO6YZo78Xs~ZTW?+dH=i+Mk1)0#}4|vSb>Ut#PR;Ef75}WNQ855u5Une5 zCE~Uk+NtgeFg6&1do2LD&VQ7b;-PKSIIrll=FiWy=qBMb8StPW_6_r#*Q<2q+E*&& zc+yp%IfIa|ta%2iXC#Y)!jAo&JtN`V{x6QMGOCTP>(Ul?C{Wzpy|@&22oNAhad$720>y(%aEAcFEjSc+DNZ3! zTw0`fvA$oP@5ijkotZmpR@S+4&e?k(L<>=13uHz)tN(W7ohqfP0*M3n142iH$*e*G z`}0Kl_`Ld^0`J=1#|*;=KYY&#gIp#^#vV{_yr$ID{tT>Ho=Zi?mEKQsN2tor+z*`u zAEr{Lzx$AiSA$qg;AqF*EvRdvI+uH0jG~|Wqnls2Q{%9$pP6Fd*eMYQ>@4~gGGZ`r z5T>_PYWvwTLVI`<;vTxUt5at|$qNmZ6edd=A7e;DVXU1l=}kfptZ&f%>rwx){(mT< z3#1pFQ*`6lLEuXjRyDTVi}W3lA51gLJG3Q?Y>x{q_<{MD4XQekf5M+;fcpPVE%hvq z;*Lm99Coz6zMKK~WR?Dhf-P?BQ#KR25uDr$#-uK z{Jlk}h!!zsvU|pg6hOh+J}l=KB`#%O>`K5iV8G3%kb2njT0SRQCQ+_mju?Y4I+>EcPgkgTb2I zGkE0*gif6>e-<3md}^%@{mxFqrL7^IA3h@M7?xrT5Cn8bsitj~>#`VCHtq#A2RQ#PE-@#~Q z0{+oez3DI=4qTEMH#f~}vP_Vk{8h3TX_Um8;}a}*5~B7rR$NhL72+F5#X&&(983$@ zU-!)&jjA@Kg1APcbN2!*r@p7XAcO(YQ#P3PjqWN}LZisVUPQ`Hv zK<-`w*7f1ynQL`%00!Gqh)?cWYtd_@$2_w}JUyIm6ZvQ`MWHRB z0x0{S1M~(tB~j6_s^NMUf}K5xPXrXIg`|It6vhsrhH%t_`fewCO3w5D64NN;zayBS zwr`M*AMvzLxB%qypS!pRcKo2Y65&dLg?s!Mvgss0c7`{N^rv1&iaWPcDX+5SUUk;r zuY6Y-pS{d~HBr4#?h_H<@w9Sm_$|6wFByO+lD7yvaWhZisJ2rl$V%& zZ)yy)fS@8~*KX@F^p)|o1p9kF2Gyq?G%0}5ILM-qw5&idjPvZXZfW+zIv);HcFbT! zYJNQGj3Ur1;@KMZ$)HneY;HfOKF-R^s~j}m1c{#j*t(P8`qjD#jW(NXD7pE))^f6X zUCKUb7xP9sfI3tL;>5(dA7lZ+SgZTm(_gs574|WXD%wb)mHz(8H1twL*`4P+$SKwI zP=9WrXQGeIobR!haVGSvC-DW1{=lK4${Nc7Tl)8u-QThr!CfjzvAOzDk3&Z~CTYO~ zTR%OL`I0&2iqq9}jmvk#fH$3IpMIfYAW1$@uP4!<$mv*}oHR)*xd!G?YM12Yy0NFW zy~l-4b&L|k6oa+l*Q6oFp4pFNF3JBx;fs>pOvJb;2gLpC=}w+TrU0%TF7cbUqnGM2 zWg-doIgHEZDX?f>^aqT`Z%izJ4vS&`rlF^gooA^v0S=Rbd~2tKqVj!ic{vTE?`cS(5mQ_91c7Tt*Wr zCo0H>UA(}{I&y>BD^B9Wf1vExN(-r|{x&S2B7}U&%HldajaeGAEC)QEH<(WIe%Sv# z6Hc3Md13#b&;rJ>rsXj-Rr;s-;2J?YszL}|RK0JnzF}YHGeRMRJ(bBaG+<0K-OMAd zk7Y|?=VsA{1!+O(#Y@y2 z%40$1w#4V45|3gHZWEDb%ywzb878H|&(hT+Z$omz07l$L&o`tKo6L>tX5(rUu@TAY z{zqoDTBZj6_kAqibGPO$1v7${?hHX4O;!_ zE)U!|fU#N;N|`t0uT-efS-knW{!Qba%yPHa&{(Dfm}5D>&MxmsRx5rQFuRHF+}Hkz z=HQSKO;4y%A^4y94Fp@Co4sqxraYH$A_b9XeLI z%mHt?U!w(HZeBhe9%9lMaO@Wo-c4*O#l!>UVr6mi3y`QpuoC8$4AOe^14(Tpi>#CC zjNIL5UDTEdyD!5yso|>StxRM@7RKeClU2n>Y`sR!3-;+iQ;x`P0-#wKCuMt2x%tI; zt4lmz)B1$R-I{kl=JJ!FxsQ5`R=(!QIqCSI15}f-^5Uc3r2oP?vD3Qo{DMY-{5%dW zFI58I@7FtyRu?~k?9G18x@9LQXIg9I`qwT&-rD{4yxDfz&)gdjBs%X$>)pEM)j5L> z;_-2%-RbIFIkS^=e}O71uCWXSoLDb)i*C2cwuzADOly~a?;Xg@oZ;~o&rsRHviyt* z`krJ!t@K3+uiqYv@hT8Ki~eCk?L<)vwL7EJ1Vb)(Vc#xTk>aps!jl9tidg|`4yk5^ z{4fF+UpGlRFN>AZs7c4;1R6GOa%ST(?I^32j4~xa5v!>DsJ_C(R73`Mq|>W5P?@~~ z?~bwnu@W2QTsCfU*w@ciy!e)e2{8)cl+(_PFu0IxJ zUy8jUS{p3JwQ=l!j`Fg;GnhqSn=N+@%;x<4Qp8*;W>JcZ_fl{_%SO1k$&S61Z$Md2P zlB29vk*>vHrt8#+BoBo;I|>`~g4)upe(Q3KiSU>?tsQgoPHU6XRo&Lse2N_<5Q&5p z5LNjLsJA>%tajAaP}QbJuMd^f7R)!dto@9ajuOO4StG6=Jq@Nb!XaBKQuhT^S=P%^ z*4#5eZfN_vd`WGv-_hIZ_C&7wh%7~b`&Opt5vBhHmPXLRL3P4!@;#j@twekzTtaG5 zYD}p`y*hwBt00Q2AHU?dBM8)cYRfH?2cd{M+)m*uWfjKOA1F_C;V0^+4zetHqu%H) zZl#D|J2o$~mqxx1sPc{aB?o+ZveFdUz|fFd=M`IJffS8S5`kc*)i+6^7m zADF2!z`Sm1zbe0z&+@yy=<8RLn>Tk)3KUHK^k=AHI&{!aOPmIGi3SBlupa=fRPVCr zOY~H}m{bv(`G5D!@!}LT~Fh zBOlYX+%^Bbg{|~a1X{Ve8EgoKK=^0t4LZ(^?`uMd(EgR}XvCY7>GAYGP&0Y=Q}bgq z{kD){^}hQhP{eb;bX?lI0ryE>Wm?$(ZohaFSF9UWB^N0B5_GNFI-^Y^PoK83i$QRW z4)_K-j;7H6GOxY;t24Ns5fA@p+C%D8H|UnNsQ#l_{8&3e&t?y0Ph=Sb=zOm_tDa9& zt6{&xr&O!ftB0?pRwhMAJiP0mnG%$%J&$J+@bqDeLr-8!NUn^vh)nB(LQve zienaJYEHcZAC5wq%0^>#_j0=leKWlvod}sd+_!3_W_bky)Z5K}oLNwB>{CP3=Qk0@ zSuILIpg?A{y%gNkGf?xSzvW18aDvZYpf*8ur;1pY`7D&br)5d834LyES-MGqwaoQe zq+r;ZueXC&_BCh2ynwOaA+&OHwk-?uqg8);n=`!D$$w5{SWmcS+*VpktBei+w+oI> zY+o=9`y8c7B8^g^^@%+LKKzQwT_8!zF*C9qW^N98N2Jo~Z|yu%&g17fzD86vF29b1 zk`CA3C=uqoy7%l=t20_ec`65g9vA;;Bh}hmv9D=UQsZqhO;QykjrL0qc$GJuF4;0B z!qD{jdCssSTGaiJMOoU}ao1RWU_#jCJ8u^R#K^q%O~j1azp~2AG%!+Qjz=54#i@CB zfthDUSw}<)dQE_Fh~yP60tUp)n-5K`Rti!sV zlz)tK7o&Jed?6{XodlIJhkLE@ zpH`ptI0>d)`FBv*J4cdr6g={ZH{rYKBw=aq2P*hLT6jX*6gs&0nANGB)GLBAlQIimN^CBZ{&MY!bDl!JO9czewElEBE(ct$$W&2O>(wY_X07nN849P^?T)Zcdk z);8KQyouWRpdZaenA;?hvIlqHOs6(%=?ubK-TV{WFUpvrJ<}5R_GO{sK=BkoS`P1@ z_V3z`n?%|Uu9+TUs;U^T&k)%~2hs*N&u8T$!DK{uii>euY+6Sm zcH@011!R`$z

1E=2F^yH{27jmuNn8N?wNruEZz`;ijP{Hn`fttMkXYRzu|7905Q z4?-Hr1B1np>~tCSa7RQKltykG8)9SbRyqz|8%}l+mp?-OOa_Y8kfVCDAK+Wl{M@2T z)M(V3q#ibXC9bdI*emvk6uDERtlOV=Oa+0e-Z|(-n>t9BW6-U1@nL`R+BDbUN#20A zwL29jps4vBqX(}_J8^n>JHKFb10~!nTR|QQ31Oz>-rZ@ zJ-iNvXW*r%+XM$Af%w#L(#?wvQ`$(nuIXeVC68S^Udbn{fvcbNu$>rB=hiv9osIOJ zzYKX(kbv@2Z)zkzIboG&r)dPh`39Sa9cKNSW8Ps2dI(sCGaADbe<=BSoWgR81CIk`vBb|>uD3z*7O41sJOtuxDWh2cGbs*>=TnS)GZX&;XMbm z0d`E}^+o!H;uSe7gOb@9U(O%{+%Dh(|69l70xoJ)Z-kl<_jlL(5!YXLPF{@WChCNI zkK@-V1kfPDD(Jq~q$t;4|5N!89+)h5e)zF7Btv|afr!?;Ii$u9Rybxc-FX65W3<#V zVm~)5oJCJPPc@y{h&DLhLj)}yH#3KKq4*T3&z8T9L&`yL@buiZ3o7zFm$0d+kUxUk z>ozycG#5&RzIt8bevw)gWb;8+Y2b4FxLnjC5}avq%kCJ47Mxl^*^qhWVfVv;DoG+8Jx95_svvOI3=+@zOQ*}gh2_)>_;?Enux@4-DNBkGDLH^0dnv7%P-s!d zFS}@!NGCaVc;|dhDs5=Rcq(4#qDdH&TJEIRUskR3WzA6*y3TF6Go;`gjj_i*o(sLL zs0dcsQl($6hz%eoAlZA{U?`#pZk!*_1OiR7!{wVIKE;l`aAg4$VLE2c~P_G4Kafz!y+FWHLxy z1Ib?QStPaUrAC;%MJjWgv3NFqR89YY+NaNCsV)Ile~Qqs_5U;3@|s_PV~B&CGX%4C zb(yIn6?k#~KI|3ArH$)Q<;*{ka{Fn7k-qVc>csZV2-BF4>j3ZG`m|s@zTsGvh!Y}_ zFkLzQb1lD$GKnY^QnTZ>It~2;^A)Da7)?3dO5u-DszZ`fH6~O~3QX}A=hEnQw|U(9 zPYN&W2O+OtCdMa5UK>|D)kPp^OCjrsAn91!_KC?qAxfGG_&rIZ1tZq+!l{?^+$`rVI*nDMO|V7bHOkc&V4szQmQg+;z!qVgN5;IOZo;zHDW ztc?xeN+p+0-;+uA@%}y1fDy_iUFSGETrytUrwI~GkihN7Bk*2j4#)5DQ@;=Z@NcS% zG91kQ4PKA;IOf96>6-4vjIZ!&0)yTXNcuCjB$s1JcwV~4p4rr;6qfXipLp@rCH-w| zfr~On`3!8Q(=UB(rAGq4(xnRY%g4E0AWmkHZG&iW>ISDD3AH@r@lznnL@(kw->8_< z@c#UhYb~1o?4ES&)Q+M?H1~lGkO!lMj#{Z(RjMffJ4ofIbuZwzC-7lOk@>mmJvi)gEqnCL!QrH%~W zru2_4-mbC0?6S?xZfzhMeTeqq&bNa|aq4Z{d?Qxg(n_`F?auO)PkjuTY7Cj@gMWE$ zqJ7y|*|q(VT)q+8pYnFx?V0z$%=M7+Rop*QHogk8i2Y7-BuXMXtLI|V|6|$noOM3W zTPpM1eG4^@29=7u*J`oRj7l^~c(xYCjuS%Vk6+rocRbV7uEV5b&nzJ1BPzG*KVgX~ zT>nEUdKQSU8gpQds=>F>S`}=1iD!ym)UOGR93qMgN2A+2? z<6f;@n7~L5BHZ^4Sx7Vv9dt^5lzYQ(%3O#dF&0;9oUUgi9$bG)B+quym)_X}yiLnW z(~QV2xNZEmFjmqguNG+Dk9Vi2Ft)T&g(195mXFq%BFmQg(@npn{kV@ZU{0tS=&vK} zvRjD^xBW1JBY1@Bz`QD{?>=W^CsvnaGwoE4Kv$)>XRHSek<4pxs#X{pQ2DIX{MfL( z%L>Yfkd!2~X${%8T0Bvk6p!t&x0b6Y<+VP{lHi4=M5Uh_?<-0Q04D83qq;zc$A$ib}Sm~l|j{ba;;c)C`6JQs7Xgi#`RpQ~g zv{sXrPc(sELq>o5eDaD*h2pO2x&OBilQ}RMZfGQ`|B1( z0z*>E3W?;&_{L^I-v6QOo$J)G{Cp0XxUJJm#VRjGE7pHCkupT%WF#3a@x8yfLZSbI z4qU^<^X{r9YBOmAd~{j+U_Eu1#$YA(`z;4j>C3u)Yb!5{_wLD*nobOk&rXsfPpWWU z|i)TP3{j(Ql%8jm%*5~tsX6FWPH>-?}85DBp>-eyYUnp`M zUUJxIVsF&KwmA$8GU)Ym)}3PQ*Yg~i7$K8Q=$rtQitJpIjSsDaSQVpp zIicgLNo)ZiIb+>Ko8x<%Ba44<;yBODa&8_zK-%2)Q)Gi+NWvs!{(9M1u^U2R^&m%bSDF*YWT zcF!sHzm=H>5oX3so9&38zx+w&Y}vC{4`?pt_df?gJ)8>teG&dTI-HBxvByGvOB|ln3_ocy65u`{qM`c@)EYP zPPx<`Isy{J1z$RjosE>JBfxJXFen(=`HDO3jcw5~p2q{N{?a;&*JbMB9KIiCRDHZ$ zNR+_Y`ekOvEhr|U&XvQH%`z^?_k1hA_>04>@w}@rW%UnR#9kUrPSWam$$=`l$o~CY z0Mmq*L*QjJpmNF`^X8_WvmPIqwzN*dyu3U0zHM}BXeyM%ULy}GH0_n?v&4Q!BWrP+ z=mh&sM)uL}Ck=u=EuNhrOD59)u6$2j%j6W0==f6s|JTtOV}LOQt!g+v2HV=xVnaRO zxgZ(&4okQsL9Yw>fFDa&L3~y!V#K~lsbdBk#_%nc?zh3Yhlmr@0Xvo&ZxV~12;l$~ z!L+znyf%Q>PDw~P`!NgQmB`zpmn>L!z2j!F9MyjL?W~RsYLTcfON#xiHPcCdP71j* zX+aqI$I$jB3V9J857B|<$kGZC3yU4Wn@LA1ZAo}EYF6DwFAMcpTj<(9rxoEw-tp{aPWJvxF z`S*ekF_+B00NXCyaCMu8lhk*H#%W#^>whT92(@R~$7MLK z)E;Y||E&-V#N!AT|NTUz`h%r8Yql-JP9x;tPDEWC$VbpWyq)6pu=)YY;^!UG+_9>t zpJ4ObAon1zkaQY@4FY#&7lNSuBKEFJe&F3e+pXVtV(Qk{7|8DKcJ8NdiC(=giL z_;|h-ZopcQ?W66Ht*1sxxh`FR?ljWO@9aL=Oiwm0PHnRVWlADA(#eU35I%5H1%!gW zeA?IZ=VCNSZJHBai`|@4FGp_M;+9us-_z;o%VTRJWt|t`YANg7>-pPd(RZY8Po6$u zV62|!!lz@Mxe=;M;hnNX#&OcO^{>Z`33gIKL49A2W3+DU53zm6YmaYdddxoPzaT^P zQc4qyPZI5T3I7b1bJuTWfWZHu_+D{rN=MnZGh~vRr6?10K+_fsB}=mcT5e%3%YGd?O9PrTjSgb%rUYuKdefTFfLWUmH4KHWL$r6Z_Q zvD}^hhvLwIWX}*I1yx#fAa|_Ck+&jzKk8+VwvbObr&+rZz0$`q`5($ryZ^n==q)*?yoc=bpR9i`x&BtZ!q|DI*ETPjMfU?x%PvY} z=DY^C-v<@h54PGA6u`*F_8NM{O84K_Y|l4XE_PD>an?f}`)7=l8~RM?-;u5w-aBF} zy14hE81oOs*>$pu;~%NzI#T#w=IXzX{YMH)lM{m_)*?wZ$)CPmY=lOUS?W0UALs-n zk@J%5onD9jY(v5xfQ~@|i;6a%lcMWlSs~)lx*^|}I_@~MHxk=C5JVU1ZL3&gsjLJ~ z=$)15TVv)D;El9V??~FiREGRlPjR6a5hfiPO10b>e{n%e+?a}Qgz_UiWVKc zFbne9_-(>Lv#UWeL0wJS`AUcO_RBrjx$VQJESvI09`2bz4NeKU2jb}n?t*^l%b!~? z=*?aGN4NI&REiFEK?Q;eu}aG3AVD@|9OD-6$H@>G;cO)6itVv;Ph9`+-9N)nX~JJ=gT%b}h(BMNd%>8K#6P^}iy|LC`K?zAIJKb$?BAlPLOB;StOhfv5YcvB zkyx%5^*jh-)uaz)g@2DjDMb-qvD&g|Ft6tR8c{+FHjE_lIo4-Kzd)*7f~1#kM-q3O zc;&92ZxM9g?`~`2TsZvgzS?b$Jt9rCp{+Wh0GqEVao5cOw2BQqiB-mAAhfjfrBm+= zsWj(DbF&w_Nhmro{m}YTfx_yfE+KwX$F#5A)ag9qLk(Ql_sui7TO{K=EMBsBt56{{ z%DH;3TF959E@!T;eB(ol0M4`c$zR|^Q}^gfQy0DOFhT0gZbPn434X4*W~%$zZo`yS zRRpB9D>X!`d;d!YNbw2EK)Y4Eyxi31(ObEh{1g!Sr)=vKpckU5i|cju;=H|h!7pHk zj8KMM{3FHd8>y;4cz-hftfa%5_eE#Q#$P5y_OEXGDZz*IN-rROPE~qD`ZN9fea!gCvXT(OqFufD2eKi|5~c_iqi025!`C< z{yguoZv&iL@2#bnvU(dhZ1igP1v1#SX45?$dWKR4#MUvkX1-QrdHqGous!awh|^H^ zT;N}ZBkg6zBXv23R`t(QYU|eJdWlcOm;VNkrd?<(aIHmQvUK{F=dkMZ`5(OIP%Ay} zlgd0Ebxj$~!7_ydAX=B~kdLQVM`yMk0eUE7-O>(@o%;(2I2Un8Ldi@57{+Pym%N?y zVLZ5ihKeQGNoe~Tl$hh))M)Tcxg+1et2+D7p^Id5S5Xw2?go&FVc> zs)$6$d!6odov*5A`*waT>Io4uezb3u-O#^BkG_4(r4?AX*J27!G@J>UwcGWzuCJ|? z_{#6nx8j?9(YnGY@}5&L@q3YTwMhaSuL&20;Z(li_W7qvQmIs7cfF+=jsXmKkhCqy2v zIMyl{IM#|$p(7|H+w6qco+K#6`e9J7FpA2m5w4^mFK{}A@Hy9O0+1(&myG(I2P)0R){dzy>{NlKJ>2Mim^_j;nC2mibFTeH?;9JvjormnAO=5 zy}Q}%UHHUH&i~L+{fT?Ix;%|Plcl}*frVT3)QQ4p#xYI#ZGQHS5DL0enErv{lUe{t zlTDhFTy+z^*7j4 z)xjo?t&&dF0?WT9a!1K0d|l}@m(YRq$)Fj1*ot!BLYAT;Ve6Cr-z3gpq?f6dkUHw3 zHkFC4Yqm)2n8tlEm3lH_+6|W#@^28>WkcLDsHY!WnjutYrt_+Tr%{|L}|AYMlqfB~xP{uyVgf?X(n%Pq)>Xo5%#RIf(W=a+e&YgUJ?h>{HOAveO+ z9HU23p!7;!pdAFJzOpL5&M>0ZR=vNLO!ADZdU2}%D!Xzwj0*gPBi5edZ?S@x&g-zZ zPB)d`Degf|wkdP*IZ@?3Vf~Ofl5Oh>Zon1oUpd7>@eKZ6BOi$b zcDIQaL)tz|2Z1%%bhLd8&ukELYX+w_sCBBod0FXBanS<#NjVjVSpY+hb{C2Pmtq;3 zEE~Z)y`}rAyYJ^tx6TC#^zfFKC{G5y#v^Ki(IR-_nRiH(hQ+8ZsF$`#NcrQG7zPPq z4*}8k?*)fDuY)EhdKkEVs|dd1F8Y81=TC|U50N_SW}qMp!Z~m)LD?(<#tAIiWbsME zR-Bft^^=8Wtl zNpuGLFz?w!j(i=`rmZ;r)U`5FhLc6CfF;el(cB6&oRi_;nKsunvUS`J8Hn|&qq|FC z3#3(oV@+d(1(O!ym1^c=;S>Kp@q*Up0GH;UA6{O7&dA2;@I;~;kG}}AiimaivkOUmk zfDbsWGkh>O%%IwU|L&8(CiSLe0G;YZQ+G)FRe$0fgCNVYAp4h)QtwRa{YmHi`(@E{ zZsJ8tekWRXO*G;q`In1xV^5x#=QUh7seWC)$hpRSMBB%#Ieb1iPZg-hv|jSnd@Oe= zbJWdvc*R$ZSB{_b)u(~_uu-+~XP1gXM&COjm^$gZ9?y5{JTn=L_CupV&3EQwW=JGuD|$j{=8K zRD|n|m?+#_XC%rRK7y8o#@w*W(5%JBPaT)DrpSMht=oF3Dn1vBV>}v~@U1J{ca){} zt_7=(*}C(Lz|c?%zj7UA#&^-4Ri?Y^SyQM+3I1r!Tm8{2( zcSlR7o(;2>Q@cjqe?Nzufg=r)-!#@s;2CD@i#Tu#9naWU@cH_sIb{MG#n#3|!pWMa zw_9$i?hM}n8j!rl)h@|F`8W+;p%l)S!iKj)*WX02>r^GBmAHjjnqu9Ul@P~r_oNGT zGrfbvdaB3Hea$P!Gux}mJ{Edrwtf`#U$1qDivv;Sf~t0zc%~hqz1I4%jBMe7U;d7` zRuqon=b4rkt2YuOV}s=koc-$JF_VhKl~@|tpUbT{JzG0ttmz?Q|1Aq>|1Ar~u`I?| zD*O+s9}L4XI%6h35Z9tYk#(L0N)k3qsRkXCN()?qcXH z?L5U7BN6DG=$yLv*cnk_(@OGM9>oZ+Rj_JKJ;Ynpe|{R2X78omJU>5eKCcTk7Fax_ zncUzE$>URk*2-t5W8NgflX$kZ(`@U%ZTSjH_&@3>83ZkK9TFUSAgq= zHYSaTUU0ck!(nBKcAxWT)<>?>)ew4WRQc=z70 z%NKL#y0fn7*iqDxWJ`Vx1oiDNNad6r z?+?PTFDJJR+N3rjzH_3!b{I}fLU!5{XCN4PM~W+~ZSjzGSkl?N)hrJ=o%$o?u$9T6 z8%l#iua~NOlDw2gtX+)BC-MDL=ghts+9EJr=QMLke^i4(@gY?ix+PusUdC@H7gg27 zEN-2c@wH8xO&W6Md<>=L?{$4E_A^gOMeh_xcl!GQPNW}EQ=?v|<=dX-0Ls{FCNNvn z(o8Fz>zMe{X%gS_nolIYmyUN4vb4z*9N36j&;-|vE~<~3udC%8m_aMXar-2b@P(712p#`qHzgH%7o~L+0C3g3e<;Ts-Ld2 z4jJoBiy@vWBKez1xq1M9JI|4Rsjr;nt8WJam~C<2Bv|)6*bw_mJ06_WZ6p%iOZR;0FU+SLGqz9Cq?uHk z6OlEytHss+D!wVl`9G98V9WQ0%2=sS8ymTEzx$Lzv~m1h<*QWlZNq=&f~9fvQ0<~a zW(dd%k3RyxKB;I`{7Tv|xZ*-z+10>lq8a}DW`Z=MM%k2M3h(@JEBivJ}25O zok`~#n8Cg?PTJ+^Qj-~v!nWjpRBHss;xH$~ZX4vakFO7K1UuC8S!ArdQzC0#u^-|C z7-+PKk&4=?Uk!GU8>x&~`#mavvxsSq-nC{Zf4$cpQKcvTNp|N1mXj%a4nvx3Y zqS(4oK<|eUm8t{zQajD%C`}YEV0A^S(0&tn?5je!fUdF-c-AND!-6vDv+F%-D(K9K z%KMYZ#eN@mhsg zjmd7q8pZWg=W7sS5=2xBFjOKwaZ7O*h$OKC-hgRHzne|hNLkjCV8gE-f$j-l{iZmU z8UZQb%Dd*JC#>A+GwSPgQfe?F^M@@B*Ly4`)m_Q_MLC~^_Ry9Asik%E#4n8r1TE?7 zU8Idlzg4)qF;?`xwj~4`DKo?&%X?Mp+e+z0XYvLEk1-`)>8wp>#o6=S`2+zMEqT4v z7G3=LYE>`m(3;9NLc3qIR7ypK<3sM?1$gXF5c3AY zng~zln7qOph``ea|BUpN;?3N7E{e7v?4NV3Z%dl05Z$ifkA3U5CbjnG9++zj(P~z1 z)jl`naO-+}A+r_M2#IKC}!pXh(DE`#X@DiHN*o_8Jv90p?U8%Mx+!AIcA!IZkVQKs^M_N<+HUY>iO0ksqa1phl*{n(RWk_IR%eWYHvh zLen%1wyjohCy#MSwfUk``ACY-gdqfvhe^o&*q?XcB4fbc)R5G`l*>+ z4DOgay=(0iH$-i&V(EIbUN=>33~JBfS#q9Q_DPfFyyRRdQ1K7J;}kyOVwBd>d`G**D-2BiGU+j8n7;nVROg2mJh2S(IiuT5T z2mo9rOCfv z-o_MMkc6O9y(39;t%DxxVMb_v6Y&@hx1sV%@M~ZhRr^txSLb)0I%ek{fkT`;SX(_m zv2TC)lt0_+3JZA~#;@f3+us|ZLnA zV>x&u|6rdHImI9yI%PEE&K)_^{O8Id-9SrFItg-C`tO8nIfA~hAJVZS^e>Y}t_`D?R6UC__rrM~x;=6y#Mt~2t=>pyp`V_LsS&diOMLpaO{IzowSRY@dAl&ZrRA5O-xujfhCu}4Cr zR(?;h3lcyUIx+nARQZr=pAx!jgz?|6TkGCMoL^4pRYUrX=bDsdNWic6AY1?sp9#QS zDIB|;Q983HY^-#cDN349@2fmAjxs(YK#o1Aexzt;W4|=~E+6<=TWxTEg}787GILxG z@Pup6$9iNFnW$qC?>g;A(`Hmsy}GHL{tWI)L4S>jN(zVu&m(kH)v(btE!}qcB~Z!RR}^BZ3Zs?3)1@#$v37iGV{Ul5Ltzs*XBR zc;bMhrwToo#@rVJbA5q-e~~%rC#QOYr=-F9!S4a0QDtSWZ@9k<5Z9L^C;CR>y?}kP z^suk@!DtXwua8>_d-ji>C9as zQ@&D2I&{QG%YVnKBm5+H#?Jen^_CyOzFnQ&TkpoPjyaJmDI{bFbmORd$GO^hu-g=B z>~yuJlw&RK>hJ(`y=+}L_voTZH5C{3uBz;QUkfLs!VaKat8P1OL@6baRj&zPncSe^ z>4SP7mr1IWNOQ4LW&cD`{BZ|8-5%OlGS`Oi)tl1w>21_m)t2ZWG5~k{r$#Hz_Y%hE zEEk?a##o}B@>z81+6l@{hLQz}0R&0zXm+4}v)9+#bp0#t8AP>Sq;&d0lB0B2*N;Gt zwR8%N+9%4nANuxv>#$Cm-ZK6;4C$U2f5Ozoe9D7^%=3w_lL>-o4!6eAITq_0zJAm5-`En2 zxbSO-aBG$Lszq9e>u9TfWr`s_C*WDhSC0uM4wqnD@ZHuP#N>(qB7S)YJJ7tU=69CF zXSN*IeiQM}%UDdEH8TV!5cWllL+9Hgc z@p!!w86!0EZB(JGNnVAd<;aSjf&}-Wuu3Ov+2(B|)Z#JszPNYEyB+4KA4Up99nj)k zPm3e@{uMH*nu8}HQsp|Yn1d%HHgSw+$K;4e#e(UHzadR`P}CEq3P&}J4~)6=Bqb zJahebf7R8B{4@z`873PCx#0QyU9@dLgSmD{9K}#LVW}LP=Z}{&%E4+kDnCK<3&bPJ zfj2g%m83N!4DNs;WE^MSFpAp`B^*S_{*76ICJMqwreWgv_$cJ-Vl{&7Hgh8BiC!cK3`p!O@<9R;SK(auMO65#ayKg^BQzw1hgH)m+2_YPBuUDr*ciqA z6r)FCcQAZ060~8C9=A7=9V5;r;Z%$Rz;zxsK3%`jG z;aUy1Tn=W!1ZL;rAnY;I(jk^8LD!k(uEc0e(qx+1ky&MSvWkprh5_IZ{<7lczB^~< zZKV|%@HrJh=mnn6wdXtJ#el&a=Y`yOKf>(a+}0{~q|P!;+(S!%4vOuk+50O}`qMF4 zs6lUYV{D8^HGlEI-;+Ni+JeSS!s#GB2Ki)1CF+q|V>W0v?*&eaXav%A9YMN%0cmCr zUfVWHx!2CxJWuiZ?OCb8BOxlv#(P5RqBnGb_?~FK$z?;t9&LtTC2CS?I%`Srgi(VI zu$846AV%cV=w^wG%W}5Yfy(tCz$I_Im3<<%x;d<0O25;RiFOg=+$D=WWF!a{-saN0 zlIiQLkC~F5?lcpXFy(RDVLD{x_xkEBclIx&rr;}p*U)&!_5thhMz)aCoLn%ZHuK0w zdT_#(MIl4>Q)r!)Q-G4SOp6f+vE-<`r zSUKN6wI-GE9Ms^dN!;5a;MOr9?dcMwruO4_co}X={3*M7^%K~9`jq7E4LAmiQkhdJ zoUNX-sF3`162VGYrox@w;Mp@b#*QIst@nEls|{h0{J5Es*g?x*!}p!WSx2tz#T6x@ zeD+%-+j}dpAs_0QRx&_trIROsl<;n$$SK zQTo>o1n@F|(9(N#qm*F5^I!EXOb#~tNlOpqaR=NHhBmpt&%#P7bs~ll5lYJ#<`VnmOfb#KV!>zi z#Np2?RaM5lR%=@GQy>lleB|fuLWy_o5=Xk~oev&@cc;%7BGqJ5O(uR?Lt|X@T|264 zx8FXE&?WkGs?Mb}+LESwRWl2cV4BigtMts+e9*CVra-e+tJx>u_;k?~uL_STvaucK z;-!(dH`|Ul0ZZ>l2%45n+}N?HBjKSfFIw+a)7^7WkU9Ft)s;n~uTq~T7wp3T1j2Lt z@Q>`BQ&RMcFN{;oVh2=|rlcA~+)<9c`6Yiync zMHG}!>`9>L*DQT)A`YuAdYy%uhXpzL3DPjVRypWKGgPV&=q`sx9ORDRdXrS* zkVt`n;C8K4j(JnFE=cO=M*RL23eCm<%~i)yJlg5k7~wRos`v%iq2EGb5t#6IRyUz`m4;u(VUitv4*V$x$52%l_b+7 z3Vvmj?)-&v9v$%-X%~_$z;kUl8;XpSVbJ5auRe^WM)8s~#Zh${oHacq&J>ee)$fX} zCWFhdl*;YK><}08t*2SZS}RkIG$!7v#y^OqMlH@T=M{~tUFDTBc<1Xy3Nnh+?7U*E zRcBM24%Ji6BsS4rjAC-mNSXJlBBO&# zzlQ;%2k2{#t{zk9chbZ{Jx@^5?vm$8wwZB;h<+S%TsE9_D-?L6n5;4!lqVxT-j#z@ z;Tu}!GNn#Zykx9(LgYRCvMKpX1qat3{e0G4q~&B)VO%bp;_oAGWVvjPg7D%n*-@T3 z=~eX-Ue!Qs1J<`%oYdF7Vw=t(;T05~;;IO({_8t2`qlECiYuvS`$HgtLHsJO*zJx$ zk@MKm@(x%Gl2f}SvMI8lw|$w<83)#jT-93I62*37ZOU;_G`L<=5$#F98c%dyYvA89 zDBa$vExTtS)j{lOax!^oqQr7L%Z{DRO>3u3D(+?oy^pOccNFIZrORzHv#SXZU}KS* znJlf5RO5aR-pwvla`Q!KiVq}_$-9htuS&Sr%e@XCPm#A)a{ICHHi(&ZXDC( ziWjnLsdh;tiHgZ=p63-7qV1HF5w*oQ!ZHN=lD*0l9!b~-@QTbA9#}XStyMb%vFEy1 zHW1|TkLy-6YcnR91MkLE@%Y!XN&ATT>BjP9v8x?5LJyqe^eDoN~i#|H+tFd?zJYzXb!t~ystWC;1r1y9zd_C@_l3wB_&wZ*0k z(Uk|Gu8K#tyMxPe7*woG@9)HK4Ii+~h_x>%aG++Z-`Op+fJWql*!Hbs?B}VRol1HV zTwV);2^eANRFxX7MJUoM`X`& zA~T;SY5+fo{{WLrjy$B1XAHf+9FOZ)%K8~8#r#F9hu#B$jl(}mYMqfrBpD!Lwogq? zX~xe&x0f`7W`?lccrCwP<>eOgo1nw+O(am5-SYxo<8Ifb^zmM5i)=gcTraB8(x|l-y zXlT>JYdXBZ?PXz=86fZ}>2#SdnZ>y3$j+4r38{Ow~Tf* zqxN>uoXguc??j^#J*~%*#O*yZPoC{%jyG6`$^hJ3nq0|hjX7nsiwsIP9A~99{%BFj z>Y(Pdzf&t~VM}Kqx76aPiSp$$f=w=B{hKVPNTl;QCmkv?B8+?Vr0iT%ELfwWV1#ltf72F92nW7t@~+}? zMTz12Q!9x=oOH!nhQ?b}^9bD7=Le-Pb#)`3qA2~INKra?Y zQd^a*tc>8zGwqCiwJqiBep!w@1?`Gju=l0pPbQ@pA~irzaaa#>K|IwcCLYQt+LCn& z-c+~;_)T;F0BI+Z)qcoUNYndXzEc&0?+3&5~$J<7p#j&SS{w zSq*r7YW?^Wx~TBQ}TGj-9cX-CY&{cAtZTye!mu|~gHm9;MkNbgzADfOr9kSs$qHv?26 znC=|ba;0sJpR!N!Qvwl=b5*qP&Wpjuwwes=x7Uq zDuKDc_o?;|v42@hsjVEj8LV%&I&?J>!potgakToDb*Q}9grllc7V&3{#~~j#6jSG- z^;!zIv&2Tp>R`!iU4zEk9C!8l)r+_#w_~&xUc^^=sR^ymCR8fZdllu}km5pxzu6U1 zca@GASZCBxc&&<2sPtuann2uR8L928Wu5S-YYEg#snbr2eF7tmIg0bp$CQ$m3G`PZuzZ<M`zA!56}tra6RQt# zT|^ToT$&a-y6R^v1qb#{Hh&K4u#uEj2c{3TdK7Gvrz48BnAPfY%K=BR=1*&V6E5f* zw@z!i)T~McYjiD>$Q98H!N|&4buNvZ46-^B9a(;&x>;gb7^IPmjtC;URxQ4U63e5} zoa9$I3{M`Ft9PSEAT1f%NX7+fu`zp`^2_A5FfL^WJo0N%d@(e#qWPEve8l!NJVZN= zI4YOe=B4Jl3A`QlYdeWrUD!;k9@wmYvrW4uZ}m!%O|zAa0z5J>`9op8hyMVtinNfc z#8HO?V!GUJv^@7aJ2E%6WNB^G{OWhcf4o2aSgAvXRJXWqE@=M%)l7OY^fa1JMhbF^ zx-%W6Z;+`6xIL>*YbodIfIdF(Yyv7hnX1G$xti^^HH>qQO1WoZrsK>kBuLIn_9xI& zvw9&)X==+-rZ%1c>P^Iq;Cj}!va!yt-OT9|4Utk~k_BZWt|eKdc+@sdMQX$0 z+ucq9CBK#w9DMDJ*DXogso4r`tJus-@wbqA*4!Qi)nzIJTL5~KCr*Bqm$$c5Pph}k z$rsARk_fJvFNYV!Tgpq00TacCItlt6f1Qw2d1s2;hP%MPB4O7Tp%kpt~{2&1095B;EIir4F=& zsYKa2klBC_<56k4t%afAC4x@otexAjt5l;KMP2tKj`l|A05r>Z@^M`Yn8QPFOfpM4 zXPn?*R@a2|?>6aVb@>;Z)-|T0wKk@@ot#pwrJND48^Y%}s<&?MCd_m^W3_R@o%Lp$ zyRyk?ZVpwz{A#2SOxuG$l&jRR>sQK<^jj9>Nme`-T#tHVS}`i|Pv(Gm2ZKg{nItm^ z6Zfup{t-|tUTm!L@T=I;)2P}^b6iO?;w2;w!leGuNJ$GCf$S>MZ)8h$=ud3d1%~ic z4!f#^+P9*BeQEnCG8~b5_9NwOf--r@#brmV$cL|5e!@#(Mx}^u42Z>8Z6mR(^GOb; z1{swmA9am1t)!8nsM(9VE%Q?^pq9nQnSu}TFG}7Ki`a=xgRtl-k(#$rM?0G)btgC@ z=~bF>*wM&?+DQoQwENUf{$?@Oq0?OpRbNpWb52`vjIrjDq{iz+XPU;Ob=WxOtJx9l zkkq)r>s6&7@;EedG(TvWbL~ggvpmC|wC+nZT5{D))J*=bjG#Os*#Du`pv9*WK|(_F~d)9fD zi;tAD)1O+XUX{nnS@VyrRt=gi9Hp@_^{TtEMZ=W1Cikj==(uv1mRuawK_vlV{{Y%* zfRiV!8Wnb4G8XAoM7SJLVO^Gq$5T|yx$j1Sc24pdu3t2YO^-5ETd5R@Rzse(7;IK@ zH>bI;458y2hd9M$iQCq>B|`4!p$t-rv8u@%5K04&YND5H6b`kEp+1K~#M674QG{Q* zpj9~TvT$n|QJK6k5PKSc;0zqqRhw@$gM&I$(V&h&!Kz~Iy=pz1E80ZafR;5%*7Sq5 zxdyY9IU;U5Hh)of28JhoQXW=B9o zq+}fTq=rzEkPbyssN=CRr=g1?$pob&;g+4Qxl8QiVV1k8ROgM^%{cPx2{Be2)_P3D z@M}h{BEgyjY)f_U0OQ?jig-Q5G%4ZsmX46bHY`F*nnyFS;<02 z_e&gpRZ;FGP*kZH>%~>zM?5v}-T!M0X*F5U>J896ivNWw?j@ku{MldS7O^)@- z^*dIZ8cB2ze7sg)oPBBLXllh1d(`qrv0TjPV!ByD+=Ck+UTji3}^1F)}A#$#TPq?p5yNgY| zJ0*=yJL0reUsPh~#q7^45DzcuP(!(U?n!~H=L&SD8CcO*A zVoP#cy>6h9gr?+aHOeb7UJGnzR=oU*bLRlME3=;cD^XE|!vwTydHEx+U0K9zJ|xQ=Rz-(`)|!Pd38 zormU?K<_}VB|jpVt~P>b7;dH$NMrXD#9m}a@1(Y458jn4q!`c{KR8N;pt?^AK}B9!85Yj&9cEWF^F zw=#)Pq3Pb7Ub>0hoSo(5_i_m(7~_mq?ZvVxuxQkbb`{qOv05ZyDoyn{=;b2n9Rb{n0Ky0g+k&w)v+l%0*#y=h1tYMc+{fJa(~Otn+F zt^J&n#uiu?W_J;ey>a~NyjSuo517_&6&B3VP7ib==$25pWrP3@;q6)`VKiVB$s)5= zYA0sRJ(V}AMm(`cZr>mPA4=ra;C(u^f-z;Y+s-vy|q~$({YB1`^EG+~)&k%NQfQ;g%mr>u)0DY<) zXD!Q8gmx~;r^9s#{nDd$DmqmtFWno4sdVKXQy)BDl2+f1aX>V_U&tZ5mhbWCQlTI*!PnDC-h>HtvEpBe2g2~n$R2OD>}0F6?rXLoA-n$ zr_N6!W@$Z*YDuRD`;Mpi(d`g>DHx4zxy@O+iLD=N9OkJwy@`C$+|HWcZdpGnnnoQ( z4k#4aLv=Do!;mX8Yuu1ESKes2ShliastFA{lgeKczWOd4q~+=j^jtYf_f#RM zxLC3tpbbvq0H&JlqBxb0KW{Zs=>SPB5Ha)6kH0)0^Bict%Q|;U+uqhSA zSV%H?=}N=qL|}1O5)`Lla}_#PfPyw0u4yvXr8TQY)GG{htmQi!!k;k$im=fE(y{ib zogTrVB3^1}qsgr0P}Ve!3maCIfdh>7sqJR+qYLR=HD@QW=+cBz)U$GB7(#?1^ggvL zvaHrp9_lm2b1fa2(%sgCmVj;Do|Q`0WzH1I^@Zf;eai+<0NTI6;lIMI)g55<>xywjfc6qM)YRhwm ztgiMY+m+<=P{Su;Yhw$SAy9+fp%PUSa_E;8Ov_AHE{l$9Kh5W@1G&2$^G28kvOMgL zwLZ=)uWKY|wcbuitivz8TCy{3vrFw+?^(h_Kx$|kk9wjQnwDB+;Wm5KR3|l?G)&Qo zPHFa*?ipl}eoWLhYZctmwRdJ-wQ?)i6+(~)wR1`{eGSy&x|B6rSCy7n;x66Dt#o*% zT%i>(l&SOyDAYLzNfeO4kF9QM@DCt};>#aS)zeP8le0z@DA`>zonS|Lf>(56t_q(+ zU8>CDcUz0*ZEk9;vSeqtrb%o_IVQ8k$7*SbqVvoTN~Px2T#FJ)+2*Ufyke6YxVyy?Fk*F@1AH7+kY?@rC zr8aBJW%EmFCMNY&p=8F%Ufud44;LuwK6s(8L#wa6$PTkEb514wWBbsP>s}pno*?Oo< A{{R30 literal 0 HcmV?d00001 diff --git a/modules/model_cogview.py b/modules/model_cogview.py new file mode 100644 index 000000000..f76a95ed3 --- /dev/null +++ b/modules/model_cogview.py @@ -0,0 +1,95 @@ +import transformers +import diffusers +from modules import shared, devices, sd_models + + +def load_common(diffusers_load_config={}, module=None): + from modules import model_quant, modelloader + modelloader.hf_login() + + if 'torch_dtype' not in diffusers_load_config: + diffusers_load_config['torch_dtype'] = 'torch.float16' + if 'low_cpu_mem_usage' in diffusers_load_config: + del diffusers_load_config['low_cpu_mem_usage'] + if 'load_connected_pipeline' in diffusers_load_config: + del diffusers_load_config['load_connected_pipeline'] + if 'safety_checker' in diffusers_load_config: + del diffusers_load_config['safety_checker'] + if 'requires_safety_checker' in diffusers_load_config: + del diffusers_load_config['requires_safety_checker'] + + quant_args = {} + if not quant_args: + quant_args = model_quant.create_bnb_config(quant_args, module=module) + if not quant_args: + quant_args = model_quant.create_ao_config(quant_args, module=module) + if quant_args: + shared.log.debug(f'Load model: type=CogView quantization module="{module}" {quant_args}') + + return diffusers_load_config, quant_args + + +def load_cogview3(checkpoint_info, diffusers_load_config={}): + repo_id = sd_models.path_to_repo(checkpoint_info.name) + shared.log.debug(f'Load model: type=CogView3 model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') + + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Model') + transformer = diffusers.CogView3PlusTransformer2DModel.from_pretrained( + repo_id, + subfolder="transformer", + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + **quant_args, + ) + + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Text Encoder') + text_encoder = transformers.T5EncoderModel.from_pretrained( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + **quant_args, + ) + + pipe = diffusers.CogView3PlusPipeline.from_pretrained( + repo_id, + text_encoder=text_encoder, + transformer=transformer, + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + ) + devices.torch_gc() + return pipe + + +def load_cogview4(checkpoint_info, diffusers_load_config={}): + repo_id = sd_models.path_to_repo(checkpoint_info.name) + shared.log.debug(f'Load model: type=CogView4 model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') + + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Model') + transformer = diffusers.CogView4Transformer2DModel.from_pretrained( + repo_id, + subfolder="transformer", + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + **quant_args, + ) + + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Text Encoder') + text_encoder = transformers.T5EncoderModel.from_pretrained( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + **quant_args, + ) + + pipe = diffusers.CogView4Pipeline.from_pretrained( + repo_id, + text_encoder=text_encoder, + transformer=transformer, + cache_dir=shared.opts.diffusers_dir, + **diffusers_load_config, + ) + devices.torch_gc() + return pipe diff --git a/modules/model_flux.py b/modules/model_flux.py index bdb42037b..190007819 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -110,39 +110,6 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu return transformer, text_encoder_2 -""" -def quant_flux_bnb(checkpoint_info, transformer, text_encoder_2): - repo_id = sd_models.path_to_repo(checkpoint_info.name) - cache_dir=shared.opts.diffusers_dir - if len(shared.opts.bnb_quantization) > 0 and (transformer is None or text_encoder_2 is None): - from modules.model_quant import load_bnb - load_bnb('Load model: type=FLUX') - try: - bnb_config = diffusers.BitsAndBytesConfig( - load_in_8bit=shared.opts.bnb_quantization_type in ['fp8'], - load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'], - bnb_4bit_quant_storage=shared.opts.bnb_quantization_storage, - bnb_4bit_quant_type=shared.opts.bnb_quantization_type, - bnb_4bit_compute_dtype=devices.dtype - ) - if ('Model' in shared.opts.bnb_quantization) and (transformer is None): - transformer = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) - shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if ('Text Encoder' in shared.opts.bnb_quantization) and (text_encoder_2 is None): - if repo_id == 'sayakpaul/flux.1-dev-nf4': - repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json - text_encoder_2 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) - shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - except Exception as e: - shared.log.error(f"Load model: type=FLUX failed quantize using BnB: {e}") - transformer, text_encoder_2 = None, None - if debug: - from modules import errors - errors.display(e, 'FLUX:') - return transformer, text_encoder_2 -""" - - def load_quants(kwargs, repo_id, cache_dir, allow_quant): try: if not allow_quant: diff --git a/modules/model_quant.py b/modules/model_quant.py index 23cc37aec..127254343 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -30,10 +30,10 @@ def get_quant(name): return 'none' -def create_bnb_config(kwargs = None, allow_bnb: bool = True): +def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Model'): from modules import shared, devices if len(shared.opts.bnb_quantization) > 0 and allow_bnb: - if 'Model' in shared.opts.bnb_quantization: + if 'Model' in shared.opts.bnb_quantization or (module is not None and module in shared.opts.bnb_quantization): load_bnb() if bnb is None: return kwargs @@ -53,10 +53,10 @@ def create_bnb_config(kwargs = None, allow_bnb: bool = True): return kwargs -def create_ao_config(kwargs = None, allow_ao: bool = True): +def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model'): from modules import shared if len(shared.opts.torchao_quantization) > 0 and shared.opts.torchao_quantization_mode == 'pre' and allow_ao: - if 'Model' in shared.opts.torchao_quantization: + if 'Model' in shared.opts.torchao_quantization or (module is not None and module in shared.opts.torchao_quantization): load_torchao() if ao is None: return kwargs diff --git a/modules/modeldata.py b/modules/modeldata.py index 0105771e5..012372b3b 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -37,6 +37,10 @@ def get_model_type(pipe): model_type = 'lumina' elif "OmniGen" in name: model_type = 'omnigen' + elif "CogView3" in name: + model_type = 'cogview3' + elif "CogView4" in name: + model_type = 'cogview4' elif "CogVideo" in name: model_type = 'cogvideox' elif "Sana" in name: diff --git a/modules/sd_detect.py b/modules/sd_detect.py index f011fd77a..642043f29 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -75,8 +75,10 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): guess = 'Kolors' if 'auraflow' in f.lower(): guess = 'AuraFlow' - if 'cogview' in f.lower(): - guess = 'CogView' + if 'cogview3' in f.lower(): + guess = 'CogView3' + if 'cogview4' in f.lower(): + guess = 'CogView4' if 'meissonic' in f.lower(): guess = 'Meissonic' pipeline = 'custom' diff --git a/modules/sd_models.py b/modules/sd_models.py index bb5b672f7..0d662f9df 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -295,9 +295,13 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op=' sd_model = load_lumina2(checkpoint_info, diffusers_load_config) elif model_type in ['Stable Diffusion 3']: from modules.model_sd3 import load_sd3 - shared.log.debug(f'Load {op}: model="Stable Diffusion 3"') - shared.opts.scheduler = 'Default' sd_model = load_sd3(checkpoint_info, cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None)) + elif model_type in ['CogView3']: # forced pipeline + from modules.model_cogview import load_cogview3 + sd_model = load_cogview3(checkpoint_info, diffusers_load_config) + elif model_type in ['CogView4']: # forced pipeline + from modules.model_cogview import load_cogview4 + sd_model = load_cogview4(checkpoint_info, diffusers_load_config) elif model_type in ['Meissonic']: # forced pipeline from modules.model_meissonic import load_meissonic sd_model = load_meissonic(checkpoint_info, diffusers_load_config) diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 0b50c5c6c..5eb4a5f94 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -9,7 +9,7 @@ from modules import shared, devices, processing, images, sd_vae_approx, sd_vae_t SamplerData = namedtuple('SamplerData', ['name', 'constructor', 'aliases', 'options']) approximation_indexes = { "Simple": 0, "Approximate": 1, "TAESD": 2, "Full VAE": 3 } -flow_models = ['f1', 'sd3', 'lumina', 'auraflow', 'sana', 'lumina2'] +flow_models = ['f1', 'sd3', 'lumina', 'auraflow', 'sana', 'lumina2', 'cogview4'] warned = False queue_lock = threading.Lock() diff --git a/modules/shared_items.py b/modules/shared_items.py index 5c1e3aebb..7b9940a45 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -89,7 +89,8 @@ def get_pipelines(): 'SegMoE': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser 'Kolors': getattr(diffusers, 'KolorsPipeline', None), 'AuraFlow': getattr(diffusers, 'AuraFlowPipeline', None), - 'CogView': getattr(diffusers, 'CogView3PlusPipeline', None), + 'CogView3': getattr(diffusers, 'CogView3PlusPipeline', None), + 'CogView4': getattr(diffusers, 'CogView4Pipeline', None), 'Stable Cascade': getattr(diffusers, 'StableCascadeCombinedPipeline', None), 'PixArt-Sigma': getattr(diffusers, 'PixArtSigmaPipeline', None), 'HunyuanDiT': getattr(diffusers, 'HunyuanDiTPipeline', None), diff --git a/wiki b/wiki index 910dc3083..cba8d182b 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 910dc3083c087eb806f548f05b6cf86ce4666268 +Subproject commit cba8d182b3cb18aeb4519d7c211d14cf9532021c From dbfd59434fffa1f1a2634bd828c268b322a4f5cc Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 15:30:57 -0400 Subject: [PATCH 002/122] add gemma3 Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 11 ++- installer.py | 2 +- modules/api/endpoints.py | 2 +- modules/interrogate/interrogate.py | 2 +- modules/interrogate/vqa.py | 88 ++++++++++++++----- modules/model_cogview.py | 6 +- modules/pixelsmith/pixelsmith_pipeline.py | 2 - modules/processing_vae.py | 2 +- modules/schedulers/scheduler_dpm_flowmatch.py | 19 ++-- modules/sd_offload.py | 6 +- modules/sd_samplers.py | 2 +- modules/shared.py | 3 +- modules/ui_caption.py | 8 +- requirements.txt | 6 +- 14 files changed, 108 insertions(+), 51 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f0b8b9c52..ba80790a8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,21 +2,30 @@ ## Update for 2025-03-15 +### TODO + - Gemma3 requires `git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` + - **Models** - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) - new foundation model for image generation based o T5-XXL text encoder and a flow-based diffusion transformer + new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer fully supports offloading and on-the-fly quantization simply select from *networks -> models -> reference* + *note* cogview4 is compatible with flowmatching samplers - New [zer0int CLiP-L](https://huggingface.co/zer0int/CLIP-Registers-Gated_MLP-ViT-L-14) models: download text encoders into folder set in settings -> system paths -> text encoders (default is `models/Text-encoder`) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui +- **Caption/VLM** + - [Google Gemma 3 4B](https://huggingface.co/google/gemma-3-4b-it) + simply select from list of available models in caption tab + - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* - **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - Updated SD3 - **Other** - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** + - update `diffusers` - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/installer.py b/installer.py index 153dae193..ca93e49df 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git: return - sha = 'b75b204a584e29ebf4e80a61be11458e9ed56e3e' # diffusers commit hash + sha = '82188cef0487837b8c70fc3f36ea63c05c85f341' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/api/endpoints.py b/modules/api/endpoints.py index 7bdc31c6a..ee30fac44 100644 --- a/modules/api/endpoints.py +++ b/modules/api/endpoints.py @@ -113,7 +113,7 @@ def post_vqa(req: models.ReqVQA): image = helpers.decode_base64_to_image(req.image) image = image.convert('RGB') from modules.interrogate import vqa - answer = vqa.interrogate(req.question, '', image, req.model) + answer = vqa.interrogate(req.question, req.system, '', image, req.model) return models.ResVQA(answer=answer) def post_unload_checkpoint(): diff --git a/modules/interrogate/interrogate.py b/modules/interrogate/interrogate.py index f68212d91..ce3f75193 100644 --- a/modules/interrogate/interrogate.py +++ b/modules/interrogate/interrogate.py @@ -28,7 +28,7 @@ def interrogate(image): elif shared.opts.interrogate_default_type == 'VLM': shared.log.info(f'Interrogate: type={shared.opts.interrogate_default_type} vlm="{shared.opts.interrogate_vlm_model}" prompt="{shared.opts.interrogate_vlm_prompt}"') from modules.interrogate import vqa - prompt = vqa.interrogate(image=image, model_name=shared.opts.interrogate_vlm_model, question=shared.opts.interrogate_vlm_prompt, prompt=None) + prompt = vqa.interrogate(image=image, model_name=shared.opts.interrogate_vlm_model, question=shared.opts.interrogate_vlm_prompt, prompt=None, system_prompt=shared.opts.interrogate_vlm_system) shared.log.debug(f'Interrogate: time={time.time()-t0:.2f} answer="{prompt}"') return prompt else: diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index afe5aac09..5dc88f459 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -22,19 +22,20 @@ vlm_models = { "CogFlorence 2.0 Large": "thwri/CogFlorence-2-Large-Freeze", # 1.6GB "CogFlorence 2.2 Large": "thwri/CogFlorence-2.2-Large", # 1.6GB "Moondream 2": "vikhyatk/moondream2", # 3.7GB + "Google Gemma 3 4B": "google/gemma-3-4b-it", + "Google Pix Textcaps": "google/pix2struct-textcaps-base", # 1.1GB + "Google PaliGemma 2 3B": "google/paligemma2-3b-pt-224", "Alibaba Qwen VL2 2B": "Qwen/Qwen2-VL-2B-Instruct", "Huggingface Smol VL2 0.5B": "HuggingFaceTB/SmolVLM-500M-Instruct", "Huggingface Smol VL2 2B": "HuggingFaceTB/SmolVLM-Instruct", "Salesforce BLIP Base": "Salesforce/blip-vqa-base", # 1.5GB "Salesforce BLIP Large": "Salesforce/blip-vqa-capfilt-large", # 1.5GB - "Google Pix Textcaps": "google/pix2struct-textcaps-base", # 1.1GB "Microsoft GIT TextCaps Base": "microsoft/git-base-textcaps", # 0.7GB "Microsoft GIT VQA Base": "microsoft/git-base-vqav2", # 0.7GB "Microsoft GIT VQA Large": "microsoft/git-large-vqav2", # 1.6GB "ToriiGate 0.4 2B": "Minthy/ToriiGate-v0.4-2B", "ToriiGate 0.4 7B": "Minthy/ToriiGate-v0.4-7B", "ViLT Base": "dandelin/vilt-b32-finetuned-vqa", # 0.5GB - "Google PaliGemma 2 3B": "google/paligemma2-3b-pt-224", "JoyCaption": "fancyfeast/llama-joycaption-alpha-two-hf-llava", # 0.7GB "JoyTag": "fancyfeast/joytag", # 17.4GB "AIDC Ovis2 1B": "AIDC-AI/Ovis2-1B", @@ -44,6 +45,7 @@ vlm_models = { # "DeepSeek VL2 Tiny": "deepseek-ai/deepseek-vl2-tiny", # broken # "nVidia Eagle 2 1B": "nvidia/Eagle2-1B", # not compatible with latest transformers } +vlm_system = 'You are image captioning expert, creative, unbiased and uncensored.' vlm_prompts = [ '', '', @@ -70,6 +72,7 @@ def b64(image): def clean(response, question): + strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:'] if isinstance(response, dict): if 'task' in response: response = response['task'] @@ -81,12 +84,10 @@ def clean(response, question): question = question.replace('<', '').replace('>', '').replace('_', ' ') if question in response: response = response.split(question, 1)[1] - response = response.replace('\n', '').replace('\r', '').replace('\t', '').strip() - if response.startswith('"'): - response = response[1:] - if response.endswith('"'): - response = response[:-1] - response = response.replace('Assistant:', '').strip() + while any(s in response for s in strip): + for s in strip: + response = response.replace(s, '') + response = response.replace('\n\n', '\n').replace(' ', ' ').replace('* ', '- ').strip() return response @@ -106,7 +107,7 @@ def get_kwargs(): return kwargs -def qwen(question: str, image: Image.Image, repo: str = None): +def qwen(question: str, image: Image.Image, repo: str = None, system_prompt: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') @@ -118,12 +119,11 @@ def qwen(question: str, image: Image.Image, repo: str = None): loaded = repo model = model.to(devices.device, devices.dtype) question = question.replace('<', '').replace('>', '').replace('_', ' ') + system_prompt = system_prompt or shared.opts.vlm_system conversation = [ { "role": "system", - "content": [ - {"type": "text", "text": "You are image captioning expert, creative, unbiased and uncensored."} - ], + "content": [{"type": "text", "text": system_prompt}], }, { "role": "user", @@ -134,7 +134,6 @@ def qwen(question: str, image: Image.Image, repo: str = None): } ] text_prompt = processor.apply_chat_template(conversation, add_generation_prompt=True) - # '<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>Describe this image.<|im_end|>\n<|im_start|>assistant\n' inputs = processor(text=[text_prompt], images=[image], padding=True, return_tensors="pt") inputs = inputs.to(devices.device, devices.dtype) output_ids = model.generate( @@ -149,6 +148,47 @@ def qwen(question: str, image: Image.Image, repo: str = None): return response +def gemma(question: str, image: Image.Image, repo: str = None, system_prompt: str = None): + global processor, model, loaded # pylint: disable=global-statement + if model is None or loaded != repo: + shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = transformers.Gemma3ForConditionalGeneration.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + loaded = repo + model = model.to(devices.device, devices.dtype) + question = question.replace('<', '').replace('>', '').replace('_', ' ') + system_prompt = system_prompt or shared.opts.vlm_system + conversation = [ + { + "role": "system", + "content": [{"type": "text", "text": system_prompt}] + }, + { + "role": "user", + "content": [ + {"type": "image", "image": b64(image)}, + {"type": "text", "text": question} + ] + } + ] + inputs = processor.apply_chat_template( + conversation, + add_generation_prompt=True, + tokenize=True, + return_dict=True, + return_tensors="pt", + ).to(device=devices.device, dtype=devices.dtype) + input_len = inputs["input_ids"].shape[-1] + with devices.inference_context(): + generation = model.generate( + **inputs, + **get_kwargs(), + ) + generation = generation[0][input_len:] + response = processor.decode(generation, skip_special_tokens=True) + return response + + def paligemma(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: @@ -219,7 +259,7 @@ def ovis(question: str, image: Image.Image, repo: str = None): return response -def smol(question: str, image: Image.Image, repo: str = None): +def smol(question: str, image: Image.Image, repo: str = None, system_prompt: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') @@ -233,12 +273,11 @@ def smol(question: str, image: Image.Image, repo: str = None): loaded = repo model.to(devices.device, devices.dtype) question = question.replace('<', '').replace('>', '').replace('_', ' ') + system_prompt = system_prompt or shared.opts.vlm_system conversation = [ { "role": "system", - "content": [ - {"type": "text", "text": "You are image captioning expert, creative, unbiased and uncensored."} - ], + "content": [{"type": "text", "text": system_prompt}], }, { "role": "user", @@ -249,7 +288,6 @@ def smol(question: str, image: Image.Image, repo: str = None): } ] text_prompt = processor.apply_chat_template(conversation, add_generation_prompt=True) - # '<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>Describe this image.<|im_end|>\n<|im_start|>assistant\n' inputs = processor(text=text_prompt, images=[image], padding=True, return_tensors="pt") inputs = inputs.to(devices.device, devices.dtype) output_ids = model.generate( @@ -410,7 +448,7 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str return response -def interrogate(question, prompt, image, model_name, quiet:bool=False): +def interrogate(question, system_prompt, prompt, image, model_name, quiet:bool=False): if not quiet: shared.state.begin('Interrogate') t0 = time.time() @@ -457,9 +495,9 @@ def interrogate(question, prompt, image, model_name, quiet:bool=False): elif 'florence' in vqa_model.lower(): answer = florence(question, image, vqa_model) elif 'qwen' in vqa_model.lower() or 'torii' in vqa_model.lower(): - answer = qwen(question, image, vqa_model) + answer = qwen(question, image, vqa_model, system_prompt) elif 'smol' in vqa_model.lower(): - answer = smol(question, image, vqa_model) + answer = smol(question, image, vqa_model, system_prompt) elif 'joytag' in vqa_model.lower(): from modules.interrogate import joytag answer = joytag.predict(image) @@ -471,6 +509,8 @@ def interrogate(question, prompt, image, model_name, quiet:bool=False): answer = deepseek.predict(question, image, vqa_model) elif 'paligemma' in vqa_model.lower(): answer = paligemma(question, image, vqa_model) + elif 'gemma' in vqa_model.lower(): + answer = gemma(question, image, vqa_model, system_prompt) elif 'ovis' in vqa_model.lower(): answer = ovis(question, image, vqa_model) else: @@ -481,7 +521,9 @@ def interrogate(question, prompt, image, model_name, quiet:bool=False): if shared.opts.interrogate_offload and model is not None: model.to(devices.cpu) devices.torch_gc() + print('HERE1', answer) answer = clean(answer, question) + print('HERE2', answer) t1 = time.time() if not quiet: shared.log.debug(f'Interrogate: type=vlm model="{model_name}" repo="{vqa_model}" args={get_kwargs()} time={t1-t0:.2f}') @@ -489,7 +531,7 @@ def interrogate(question, prompt, image, model_name, quiet:bool=False): return answer -def batch(model_name, batch_files, batch_folder, batch_str, question, prompt, write, append, recursive): +def batch(model_name, system_prompt, batch_files, batch_folder, batch_str, question, prompt, write, append, recursive): class BatchWriter: def __init__(self, folder, mode='w'): self.folder = folder @@ -536,7 +578,7 @@ def batch(model_name, batch_files, batch_folder, batch_str, question, prompt, wr if shared.state.interrupted: break image = Image.open(file) - prompt = interrogate(question, prompt, image, model_name, quiet=True) + prompt = interrogate(question, system_prompt, prompt, image, model_name, quiet=True) prompts.append(prompt) if write: writer.add(file, prompt) diff --git a/modules/model_cogview.py b/modules/model_cogview.py index f76a95ed3..8ced40ce2 100644 --- a/modules/model_cogview.py +++ b/modules/model_cogview.py @@ -76,7 +76,7 @@ def load_cogview4(checkpoint_info, diffusers_load_config={}): ) diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Text Encoder') - text_encoder = transformers.T5EncoderModel.from_pretrained( + text_encoder = transformers.AutoModelForCausalLM.from_pretrained( repo_id, subfolder="text_encoder", cache_dir=shared.opts.diffusers_dir, @@ -91,5 +91,9 @@ def load_cogview4(checkpoint_info, diffusers_load_config={}): cache_dir=shared.opts.diffusers_dir, **diffusers_load_config, ) + if shared.opts.diffusers_eval: + pipe.text_encoder.eval() + pipe.transformer.eval() + pipe.enable_model_cpu_offload() # TODO cogview4: balanced offload does not work for GlmModel devices.torch_gc() return pipe diff --git a/modules/pixelsmith/pixelsmith_pipeline.py b/modules/pixelsmith/pixelsmith_pipeline.py index 4e04b2d92..702ee67f6 100644 --- a/modules/pixelsmith/pixelsmith_pipeline.py +++ b/modules/pixelsmith/pixelsmith_pipeline.py @@ -133,7 +133,6 @@ class PAGIdentitySelfAttnProcessor: value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 hidden_states_org = F.scaled_dot_product_attention( query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False ) @@ -248,7 +247,6 @@ class PAGCFGIdentitySelfAttnProcessor: value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 hidden_states_org = F.scaled_dot_product_attention( query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False ) diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 54dafb940..be819d658 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -139,7 +139,7 @@ def full_vae_decode(latents, model): if latents_mean and latents_std: latents_mean = (torch.tensor(latents_mean).view(1, 4, 1, 1).to(latents.device, latents.dtype)) latents_std = (torch.tensor(latents_std).view(1, 4, 1, 1).to(latents.device, latents.dtype)) - latents = latents * latents_std / scaling_factor + latents_mean + latents = ((latents * latents_std) / scaling_factor) + latents_mean else: latents = latents / scaling_factor if shift_factor: diff --git a/modules/schedulers/scheduler_dpm_flowmatch.py b/modules/schedulers/scheduler_dpm_flowmatch.py index 69452aca9..ab9aa47a9 100644 --- a/modules/schedulers/scheduler_dpm_flowmatch.py +++ b/modules/schedulers/scheduler_dpm_flowmatch.py @@ -230,6 +230,7 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): device: Union[str, torch.device] = None, sigmas: Optional[List[float]] = None, mu: Optional[float] = None, + timesteps: Optional[torch.Tensor] = None, ): """ Sets the discrete timesteps used for the diffusion chain (to be run before inference). @@ -355,12 +356,12 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): sigma_min = sigmas[-1] sigmas = np.linspace(1.0, sigma_min, num_inference_steps) sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device) - + if self.config.use_dynamic_shifting: sigmas = self.time_shift(mu, 1.0, sigmas) else: sigmas = self.config.shift * sigmas / (1 + (self.config.shift - 1) * sigmas) - + timesteps = sigmas * self.config.num_train_timesteps self.timesteps = timesteps.to(device=device) self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)]) @@ -517,7 +518,7 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): model_output = sample - sigma * model_output d = (sample - model_output) / sigma dt = sigma_next - sigma - sample = sample + d * dt + sample = sample + d * dt else: # DPM-Solver2 sigma_mid = sigma.log().lerp(sigma_next.log(), 0.5).exp() @@ -596,7 +597,7 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): elif self.config.algorithm_type == "dpmsolver++2M": if self.config.solver_order == 2: t, t_next = t_fn(sigma), t_fn(sigma_next) - h = t_next - t + h = t_next - t if self.model_outputs[-2] is None or sigma_next == 0: sample = (sigma_fn(t_next) / sigma_fn(t)) * sample - (-h).expm1() * model_output else: @@ -703,7 +704,7 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): if self.config.use_noise_sampler: sample = sample + self.noise_sampler(sigma_fn(t), sigma_fn(t_next)) * self.config.s_noise * su else: - sample = sample + noise * self.config.s_noise * su + sample = sample + noise * self.config.s_noise * su del x_2 del denoised_2 del d @@ -745,13 +746,13 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): t, s = -sigma.log(), -sigma_next.log() h = s - t h_eta = h * 2 - + # 3. Delta timestep dt = sigma_next - sigma sample = sample + model_output * dt sample = torch.exp(-h_eta) * sample + (-h_eta).expm1().neg() * model_output - + if self.h_2 is not None: r0 = self.h_1 / h r1 = self.h_2 / h @@ -780,13 +781,13 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): sample = sample + self.noise_sampler(sigma, sigma_next) * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise else: sample = sample + noise * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise - + self.h_2 = self.h_1 self.h_1 = h if not self.config.use_noise_sampler and noise is not None: del noise prev_sample = sample - + # Cast sample back to expected dtype prev_sample = prev_sample.to(model_output.dtype) diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 33ce1b2c0..2c4126209 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -10,8 +10,9 @@ from modules.timer import process as process_timer debug_move = shared.log.trace if os.environ.get('SD_MOVE_DEBUG', None) is not None else lambda *args, **kwargs: None -should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'hunyuanvideo', 'cogvideox', 'mochi'] +should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'hunyuanvideo', 'cogvideox', 'mochi', 'cogview4'] offload_hook_instance = None +balanced_offload_exclude = ['OmniGenPipeline', 'CogView4Pipeline'] def get_signature(cls): @@ -193,8 +194,7 @@ def apply_balanced_offload(sd_model, exclude=[]): if sd_model is None: return sd_model t0 = time.time() - excluded = ['OmniGenPipeline'] - if sd_model.__class__.__name__ in excluded: + if sd_model.__class__.__name__ in balanced_offload_exclude: return sd_model cached = True checkpoint_name = sd_model.sd_checkpoint_info.name if getattr(sd_model, "sd_checkpoint_info", None) is not None else None diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index accd6b0ed..67ed8e8ee 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -79,7 +79,7 @@ def create_sampler(name, model): shared.log.debug(f'Sampler: "{name}" config={config.options}') return sampler elif shared.native: - FlowModels = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow', 'Sana', 'HunyuanVideoPipeline'] + FlowModels = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow', 'Sana', 'HunyuanVideoPipeline', 'CogView4Pipeline'] if 'KDiffusion' in model.__class__.__name__: return None if not any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' in name: diff --git a/modules/shared.py b/modules/shared.py index 962bf05d0..10cb1007e 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -19,7 +19,7 @@ from modules.dml import memory_providers, default_memory_provider, directml_do_h from modules.onnx_impl import initialize_onnx, execution_providers from modules.memstats import memory_stats, ram_stats # pylint: disable=unused-import from modules.interrogate.openclip import caption_models, caption_types, get_clip_models, refresh_clip_models, category_types -from modules.interrogate.vqa import vlm_models, vlm_prompts +from modules.interrogate.vqa import vlm_models, vlm_prompts, vlm_system from modules.ui_components import DropdownEditable from modules.options import OptionInfo import modules.memmon @@ -855,6 +855,7 @@ options_templates.update(options_section(('interrogate', "Interrogate"), { "interrogate_vlm_sep": OptionInfo("

VLM

", "", gr.HTML), "interrogate_vlm_model": OptionInfo(list(vlm_models)[0], "VLM: default model", gr.Dropdown, {"choices": list(vlm_models)}), "interrogate_vlm_prompt": OptionInfo(vlm_prompts[2], "VLM: default prompt", DropdownEditable, {"choices": vlm_prompts }), + "interrogate_vlm_system": OptionInfo(vlm_system, "VLM: default prompt"), "interrogate_vlm_num_beams": OptionInfo(3, "VLM: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}), "interrogate_vlm_max_length": OptionInfo(512, "VLM: max length", gr.Slider, {"minimum": 1, "maximum": 4096, "step": 1, "visible": False}), "interrogate_vlm_do_sample": OptionInfo(False, "VLM: use sample method"), diff --git a/modules/ui_caption.py b/modules/ui_caption.py index 04a16702c..474427e9d 100644 --- a/modules/ui_caption.py +++ b/modules/ui_caption.py @@ -35,6 +35,8 @@ def create_ui(): with gr.Tabs(elem_id="mode_caption"): with gr.Tab("VLM Caption", elem_id="tab_vlm_caption"): from modules.interrogate import vqa + with gr.Row(): + vlm_system = gr.Textbox(label="System prompt", value=vqa.vlm_system, lines=1, elem_id='vlm_system') with gr.Row(): vlm_question = gr.Dropdown(label="Predefined question", allow_custom_value=False, choices=vqa.vlm_prompts, value=vqa.vlm_prompts[2], elem_id='vlm_question') with gr.Row(): @@ -114,7 +116,7 @@ def create_ui(): btn_clip_analyze_img = gr.Button("Analyze", variant='primary', elem_id="btn_clip_analyze_img") with gr.Column(variant='compact', elem_id='interrogate_output'): with gr.Row(elem_id='interrogate_output_prompt'): - prompt = gr.Textbox(label="Answer", lines=8, placeholder="ai generated image description") + prompt = gr.Textbox(label="Answer", lines=12, placeholder="ai generated image description") with gr.Row(elem_id='interrogate_output_classes'): medium = gr.Label(elem_id="interrogate_label_medium", label="Medium", num_top_classes=5, visible=False) artist = gr.Label(elem_id="interrogate_label_artist", label="Artist", num_top_classes=5, visible=False) @@ -127,8 +129,8 @@ def create_ui(): btn_clip_interrogate_img.click(openclip.interrogate_image, inputs=[image, clip_model, blip_model, clip_mode], outputs=[prompt]) btn_clip_analyze_img.click(openclip.analyze_image, inputs=[image, clip_model, blip_model], outputs=[medium, artist, movement, trending, flavor]) btn_clip_interrogate_batch.click(fn=openclip.interrogate_batch, inputs=[clip_batch_files, clip_batch_folder, clip_batch_str, clip_model, blip_model, clip_mode, clip_save_output, clip_save_append, clip_folder_recursive], outputs=[prompt]) - btn_vlm_caption.click(fn=vqa.interrogate, inputs=[vlm_question, vlm_prompt, image, vlm_model], outputs=[prompt]) - btn_vlm_caption_batch.click(fn=vqa.batch, inputs=[vlm_model, vlm_batch_files, vlm_batch_folder, vlm_batch_str, vlm_question, vlm_prompt, vlm_save_output, vlm_save_append, vlm_folder_recursive], outputs=[prompt]) + btn_vlm_caption.click(fn=vqa.interrogate, inputs=[vlm_question, vlm_system, vlm_prompt, image, vlm_model], outputs=[prompt]) + btn_vlm_caption_batch.click(fn=vqa.batch, inputs=[vlm_model, vlm_system, vlm_batch_files, vlm_batch_folder, vlm_batch_str, vlm_question, vlm_prompt, vlm_save_output, vlm_save_append, vlm_folder_recursive], outputs=[prompt]) for tabname, button in copy_interrogate_buttons.items(): generation_parameters_copypaste.register_paste_params_button(generation_parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=prompt, source_image_component=image,)) diff --git a/requirements.txt b/requirements.txt index f61890cbe..e64ada91b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -41,18 +41,18 @@ torchsde==0.2.6 antlr4-python3-runtime==4.9.3 requests==2.32.3 tqdm==4.67.1 -accelerate==1.3.0 +accelerate==1.5.2 opencv-contrib-python-headless==4.9.0.80 einops==0.4.1 gradio==3.43.2 -huggingface_hub==0.28.1 +huggingface_hub==0.29.3 numexpr==2.8.8 numpy==1.26.4 numba==0.59.1 protobuf==4.25.3 pytorch_lightning==1.9.4 tokenizers==0.21.0 -transformers==4.48.3 +transformers==4.49.0 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 From b7666802548c5300094d3ee54af4299bc91061f5 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 15:55:33 -0400 Subject: [PATCH 003/122] remote vae scaling-factor and shift-factor Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/omnigen/utils.py | 5 +---- modules/sd_vae_remote.py | 39 ++++++++++++++++++++++++++++++++++++--- 3 files changed, 38 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index ba80790a8..4b84016c1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,6 +26,7 @@ - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** - update `diffusers` + - remote vae use `scaling_factor` and `shift_factor` - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/omnigen/utils.py b/modules/omnigen/utils.py index 8ace4fab6..bf0a6de62 100644 --- a/modules/omnigen/utils.py +++ b/modules/omnigen/utils.py @@ -28,8 +28,6 @@ def update_ema(ema_model, model, decay=0.9999): ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay) - - def requires_grad(model, flag=True): """ Set requires_grad flag for all parameters in a model. @@ -59,7 +57,6 @@ def center_crop_arr(pil_image, image_size): return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) - def crop_arr(pil_image, max_image_size): while min(*pil_image.size) >= 2 * max_image_size: pil_image = pil_image.resize( @@ -89,7 +86,6 @@ def crop_arr(pil_image, max_image_size): return Image.fromarray(arr) - def vae_encode(vae, x, weight_dtype): if x is not None: if vae.config.shift_factor is not None: @@ -100,6 +96,7 @@ def vae_encode(vae, x, weight_dtype): x = x.to(weight_dtype) return x + def vae_encode_list(vae, x, weight_dtype): latents = [] for img in x: diff --git a/modules/sd_vae_remote.py b/modules/sd_vae_remote.py index 7d2645b9a..d69a616db 100644 --- a/modules/sd_vae_remote.py +++ b/modules/sd_vae_remote.py @@ -7,12 +7,17 @@ from PIL import Image from safetensors.torch import _tobytes -hf_endpoints = { +hf_decode_endpoints = { 'sd': 'https://q1bj3bpq6kzilnsu.us-east-1.aws.endpoints.huggingface.cloud', 'sdxl': 'https://x2dmsqunjd6k9prw.us-east-1.aws.endpoints.huggingface.cloud', 'f1': 'https://whhx50ex1aryqvw6.us-east-1.aws.endpoints.huggingface.cloud', 'hunyuanvideo': 'https://o7ywnmrahorts457.us-east-1.aws.endpoints.huggingface.cloud', } +hf_encode_endpoints = { + 'sd': 'https://qc6479g0aac6qwy9.us-east-1.aws.endpoints.huggingface.cloud', + 'sdxl': 'https://xjqqhmyn62rog84g.us-east-1.aws.endpoints.huggingface.cloud', + 'f1': 'https://ptccx55jz97f9zgo.us-east-1.aws.endpoints.huggingface.cloud', +} dtypes = { "float16": torch.float16, "float32": torch.float32, @@ -26,18 +31,19 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ tensors = [] content = 0 model_type = model_type or shared.sd_model_type - url = hf_endpoints.get(model_type, None) + url = hf_decode_endpoints.get(model_type, None) if url is None: shared.log.error(f'Decode: type="remote" type={model_type} unsuppported') return tensors t0 = time.time() modelloader.hf_login() latents = latents.unsqueeze(0) if len(latents.shape) == 3 else latents + from diffusers.utils.remote_utils import remote_decode + for i in range(latents.shape[0]): try: latent = latents[i].detach().clone().to(device=devices.cpu, dtype=devices.dtype).unsqueeze(0) params = { - "do_scaling": True, "input_tensor_type": "binary", "shape": list(latent.shape), "dtype": str(latent.dtype).split(".", maxsplit=1)[-1], @@ -59,6 +65,9 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ if (model_type == 'f1') and (width > 0) and (height > 0): params['width'] = width params['height'] = height + if shared.sd_model.vae is not None and shared.sd_model.vae.config is not None: + params['scaling_factor'] = shared.sd_model.vae.config.get("scaling_factor", None) + params['shift_factor'] = shared.sd_model.vae.config.get("shift_factor", None) response = requests.post( url=url, headers=headers, @@ -86,3 +95,27 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ t1 = time.time() shared.log.debug(f'Decode: type="remote" model={model_type} mode={shared.opts.remote_vae_type} args={params} bytes={content} time={t1-t0:.3f}s') return tensors + + +def remote_encode(image: Image.Image, model_type: str = None) -> torch.Tensor: + from modules import devices, shared, errors, modelloader + tensors = [] + model_type = model_type or shared.sd_model_type + url = hf_decode_endpoints.get(model_type, None) + if url is None: + shared.log.error(f'Decode: type="remote" type={model_type} unsuppported') + return tensors + t0 = time.time() + modelloader.hf_login() + + try: + params = {} + content = 0 + tensor = None + except Exception as e: + shared.log.error(f'Encode: type="remote" model={model_type} {e}') + errors.display(e, 'VAE') + + t1 = time.time() + shared.log.debug(f'Encode: type="remote" model={model_type} mode={shared.opts.remote_vae_type} args={params} image={image} bytes={content} time={t1-t0:.3f}s') + return tensor From a91c95870d85ee2f0a104192743157f20ead6944 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 17:03:37 -0400 Subject: [PATCH 004/122] remote vae encode Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 5 +++++ modules/interrogate/vqa.py | 5 +++-- modules/postprocess/yolo.py | 1 + modules/processing_args.py | 18 +++++++++++++--- modules/sd_vae_remote.py | 41 +++++++++++++++++++++++++------------ modules/shared.py | 1 + 6 files changed, 53 insertions(+), 18 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4b84016c1..7d46b5173 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,7 @@ ### TODO - Gemma3 requires `git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` + - Remote VAE encode for SD15 and Flux.1: - **Models** - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) @@ -15,6 +16,10 @@ download text encoders into folder set in settings -> system paths -> text encoders (default is `models/Text-encoder`) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui +- **Remote VAE** + - add support for remote vae encode in addition to remote vae decode + - used by *img2img, inpaint, hires, detailer* + - remote vae encode is disabled by default, you can enable it in *settings -> variable auto-encoder* - **Caption/VLM** - [Google Gemma 3 4B](https://huggingface.co/google/gemma-3-4b-it) simply select from list of available models in caption tab diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index 5dc88f459..2bdec3543 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -150,6 +150,9 @@ def qwen(question: str, image: Image.Image, repo: str = None, system_prompt: str def gemma(question: str, image: Image.Image, repo: str = None, system_prompt: str = None): global processor, model, loaded # pylint: disable=global-statement + if not hasattr(transformers, 'Gemma3ForConditionalGeneration'): + shared.log.error(f'Interrogate: vlm="{repo}" gemma is not available') + return '' if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') model = transformers.Gemma3ForConditionalGeneration.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) @@ -521,9 +524,7 @@ def interrogate(question, system_prompt, prompt, image, model_name, quiet:bool=F if shared.opts.interrogate_offload and model is not None: model.to(devices.cpu) devices.torch_gc() - print('HERE1', answer) answer = clean(answer, question) - print('HERE2', answer) t1 = time.time() if not quiet: shared.log.debug(f'Interrogate: type=vlm model="{model_name}" repo="{vqa_model}" args={get_kwargs()} time={t1-t0:.2f}') diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index 4aa80c613..8054e03ea 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -265,6 +265,7 @@ class YoloRestorer(Detailer): 'inpaint_full_res_padding': shared.opts.detailer_padding, 'width': resolution, 'height': resolution, + 'vae_type': orig_p.get('vae_type', 'Full'), } if args['denoising_strength'] == 0: shared.log.debug(f'Detailer: model="{name}" strength=0 skip') diff --git a/modules/processing_args.py b/modules/processing_args.py index ff693a03e..c0e201f77 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -38,13 +38,20 @@ def task_specific_kwargs(p, model): model.register_to_config(requires_aesthetics_score = False) if 'hires' not in p.ops: p.ops.append('img2img') + if p.vae_type == 'Remote': + from modules.sd_vae_remote import remote_encode + p.init_images = remote_encode(p.init_images) task_args = { 'image': p.init_images, 'strength': p.denoising_strength, } if model.__class__.__name__ == 'FluxImg2ImgPipeline': # needs explicit width/height - p.width = 8 * math.ceil(p.init_images[0].width / 8) - p.height = 8 * math.ceil(p.init_images[0].height / 8) + if torch.is_tensor(p.init_images[0]): + p.width = p.init_images[0].shape[-1] * 16 + p.height = p.init_images[0].shape[-2] * 16 + else: + p.width = 8 * math.ceil(p.init_images[0].width / 8) + p.height = 8 * math.ceil(p.init_images[0].height / 8) task_args['width'], task_args['height'] = p.width, p.height if model.__class__.__name__ == 'OmniGenPipeline': p.width = 16 * math.ceil(p.init_images[0].width / 16) @@ -70,9 +77,14 @@ def task_specific_kwargs(p, model): else: p.ops.append('inpaint') width, height = processing_helpers.resize_init_images(p) + mask_image = p.task_args.get('image_mask', None) or getattr(p, 'image_mask', None) or getattr(p, 'mask', None) + if p.vae_type == 'Remote': + from modules.sd_vae_remote import remote_encode + p.init_images = remote_encode(p.init_images) + # mask_image = remote_encode(mask_image) task_args = { 'image': p.init_images, - 'mask_image': p.task_args.get('image_mask', None) or getattr(p, 'image_mask', None) or getattr(p, 'mask', None), + 'mask_image': mask_image, 'strength': p.denoising_strength, 'height': height, 'width': width, diff --git a/modules/sd_vae_remote.py b/modules/sd_vae_remote.py index d69a616db..55b8aa7ff 100644 --- a/modules/sd_vae_remote.py +++ b/modules/sd_vae_remote.py @@ -1,3 +1,4 @@ +from typing import List import io import time import json @@ -38,7 +39,6 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ t0 = time.time() modelloader.hf_login() latents = latents.unsqueeze(0) if len(latents.shape) == 3 else latents - from diffusers.utils.remote_utils import remote_decode for i in range(latents.shape[0]): try: @@ -97,25 +97,40 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ return tensors -def remote_encode(image: Image.Image, model_type: str = None) -> torch.Tensor: +def remote_encode(images: List[Image.Image], model_type: str = None) -> torch.Tensor: + from diffusers.utils import remote_utils from modules import devices, shared, errors, modelloader + if not shared.opts.remote_vae_encode: + return images tensors = [] model_type = model_type or shared.sd_model_type - url = hf_decode_endpoints.get(model_type, None) + url = hf_encode_endpoints.get(model_type, None) if url is None: shared.log.error(f'Decode: type="remote" type={model_type} unsuppported') - return tensors + return images t0 = time.time() modelloader.hf_login() - try: - params = {} - content = 0 - tensor = None - except Exception as e: - shared.log.error(f'Encode: type="remote" model={model_type} {e}') - errors.display(e, 'VAE') + if isinstance(images, Image.Image): + images = [images] + for init_image in images: + try: + init_latent = remote_utils.remote_encode( + endpoint=url, + image=init_image, + scaling_factor = shared.sd_model.vae.config.get("scaling_factor", None), + shift_factor = shared.sd_model.vae.config.get("shift_factor", None), + ) + tensors.append(init_latent) + except Exception as e: + shared.log.error(f'Encode: type="remote" model={model_type} {e}') + errors.display(e, 'VAE') + if len(tensors) > 0 and torch.is_tensor(tensors[0]): + tensors = torch.cat(tensors, dim=0) + tensors = tensors.to(dtype=devices.dtype) + else: + return images t1 = time.time() - shared.log.debug(f'Encode: type="remote" model={model_type} mode={shared.opts.remote_vae_type} args={params} image={image} bytes={content} time={t1-t0:.3f}s') - return tensor + shared.log.debug(f'Encode: type="remote" model={model_type} mode={shared.opts.remote_vae_type} image={images} latent={tensors.shape} time={t1-t0:.3f}s') + return tensors diff --git a/modules/shared.py b/modules/shared.py index 10cb1007e..ed3ef031d 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -424,6 +424,7 @@ options_templates.update(options_section(('vae_encoder', "Variable Auto Encoder" "sd_vae_sliced_encode": OptionInfo(False, "VAE sliced encode", gr.Checkbox, {"visible": not native}), "nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox), "remote_vae_type": OptionInfo('raw', "Remote VAE image type", gr.Dropdown, {"choices": ['raw', 'jpg', 'png']}), + "remote_vae_encode": OptionInfo(False, "Remote VAE for encode"), "rollback_vae": OptionInfo(False, "Attempt VAE roll back for NaN values", gr.Checkbox, {"visible": not native}), })) From d4a67dd946b3470339a3785c116d49839e94d319 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 19:04:03 -0400 Subject: [PATCH 005/122] lora logging Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + cli/prompt-detect.py | 32 ++++++++++++++++ .../Lora/extra_networks_lora.py | 4 +- extensions-builtin/Lora/networks.py | 36 +++++++++--------- modules/extra_networks.py | 2 +- modules/lora/extra_networks_lora.py | 12 +++--- modules/lora/networks.py | 38 +++++++++---------- modules/processing_callbacks.py | 11 ++++-- modules/processing_diffusers.py | 7 ++-- .../textual_inversion/textual_inversion.py | 2 +- 10 files changed, 91 insertions(+), 54 deletions(-) create mode 100644 cli/prompt-detect.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 7d46b5173..133e0d39c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -38,6 +38,7 @@ - fix cuda errors with *directml* - fix memory stats not displaying the ram usage - fix **RunPod** memory limit reporting + - fix flux ipadapter with start/stop values - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/cli/prompt-detect.py b/cli/prompt-detect.py new file mode 100644 index 000000000..cd3d0189a --- /dev/null +++ b/cli/prompt-detect.py @@ -0,0 +1,32 @@ +# Example: +# > python cli/lang-detect.py "have a good day" +# > ['eng_latn:1.00'] +# eng=language, latn=latin alphabet, 1.00=confidence + +import sys +import fasttext +from huggingface_hub import hf_hub_download + + +repo_id = "facebook/fasttext-language-identification" +model = None + + +def detect(text:str, top:int=1, threshold:float=0.25) -> str: + try: + global model # pylint: disable=global-statement + if model is None: + model_path = hf_hub_download(repo_id, filename="model.bin") + model = fasttext.load_model(model_path) + lang, score = model.predict(text, k=top, threshold=threshold, on_unicode_error="ignore") + result = [f"{l.replace("__label__", "").lower()}:{s:.2f}" for l, s in zip(lang, score) if s > threshold][:top] + return result + except Exception as e: + return str(e) + + +if __name__ == "__main__": + if len(sys.argv) < 2: + print(f"Usage: {sys.argv[0]} ") + else: + print(detect(sys.argv[1])) diff --git a/extensions-builtin/Lora/extra_networks_lora.py b/extensions-builtin/Lora/extra_networks_lora.py index 76d490eda..2cbdaea60 100644 --- a/extensions-builtin/Lora/extra_networks_lora.py +++ b/extensions-builtin/Lora/extra_networks_lora.py @@ -55,7 +55,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): loaded.tags = loaded.tags[:shared.opts.lora_apply_tags] all_tags.extend(loaded.tags) if len(all_tags) > 0: - shared.log.debug(f"Load network: type=LoRA tags={all_tags} max={shared.opts.lora_apply_tags} apply") + shared.log.debug(f"Network load: type=LoRA tags={all_tags} max={shared.opts.lora_apply_tags} apply") all_tags = ', '.join(all_tags) p.extra_generation_params["LoRA tags"] = all_tags if '_tags_' in p.prompt: @@ -129,7 +129,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if len(networks.loaded_networks) > 0 and step == 0: self.infotext(p) self.prompt(p) - shared.log.info(f'Load network: type=LoRA apply={[n.name for n in networks.loaded_networks]} method=legacy te={te_multipliers} unet={unet_multipliers} dims={dyn_dims} load={t1-t0:.2f}') + shared.log.info(f'Network load: type=LoRA apply={[n.name for n in networks.loaded_networks]} method=legacy te={te_multipliers} unet={unet_multipliers} dims={dyn_dims} load={t1-t0:.2f}') def deactivate(self, p): t0 = time.time() diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 1f02f3846..e59555993 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -95,13 +95,13 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ t0 = time.time() name = name.replace(".", "_") #cached = lora_cache.get(name, None) - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') # if cached is not None: # return cached if not shared.native: return None if not hasattr(shared.sd_model, 'load_lora_weights'): - shared.log.error(f'Load network: type=LoRA class={shared.sd_model.__class__} does not implement load lora') + shared.log.error(f'Network load: type=LoRA class={shared.sd_model.__class__} does not implement load lora') return None try: shared.sd_model.load_lora_weights(network_on_disk.filename, adapter_name=name) @@ -110,9 +110,9 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ pass else: if 'The following keys have not been correctly renamed' in str(e): - shared.log.error(f'Load network: type=LoRA name="{name}" diffusers unsupported format') + shared.log.error(f'Network load: type=LoRA name="{name}" diffusers unsupported format') else: - shared.log.error(f'Load network: type=LoRA name="{name}" {e}') + shared.log.error(f'Network load: type=LoRA name="{name}" {e}') if debug: errors.display(e, "LoRA") return None @@ -133,7 +133,7 @@ def load_network(name, network_on_disk) -> network.Network: t0 = time.time() cached = lora_cache.get(name, None) if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') if cached is not None: return cached net = network.Network(name, network_on_disk) @@ -182,11 +182,11 @@ def load_network(name, network_on_disk) -> network.Network: else: net.modules[key] = net_module if len(keys_failed_to_match) > 0: - shared.log.warning(f'Load network: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') + shared.log.warning(f'Network load: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" unmatched={keys_failed_to_match}') + shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') else: - shared.log.debug(f'Load network: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)}') + shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)}') if len(matched_networks) == 0: return None lora_cache[name] = net @@ -233,7 +233,7 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No if network_on_disk is not None: shorthash = getattr(network_on_disk, 'shorthash', '').lower() if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') try: if recompile_model: shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}") @@ -245,13 +245,13 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No net.mentioned_name = name network_on_disk.read_hash() except Exception as e: - shared.log.error(f'Load network: type=LoRA file="{network_on_disk.filename}" {e}') + shared.log.error(f'Network load: type=LoRA file="{network_on_disk.filename}" {e}') if debug: errors.display(e, 'LoRA') continue if net is None: failed_to_load_networks.append(name) - shared.log.error(f'Load network: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') + shared.log.error(f'Network load: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') continue if shared.native: shared.sd_model.embedding_db.load_diffusers_embedding(None, net.bundle_embeddings) @@ -265,24 +265,24 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No lora_cache.pop(name, None) if len(diffuser_loaded) > 0: - shared.log.debug(f'Load network: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') + shared.log.debug(f'Network load: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') try: shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) if shared.opts.lora_fuse_diffusers: shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling shared.sd_model.unload_lora_weights() except Exception as e: - shared.log.error(f'Load network: type=LoRA {e}') + shared.log.error(f'Network load: type=LoRA {e}') if debug: errors.display(e, 'LoRA') if len(loaded_networks) > 0 and debug: - shared.log.debug(f'Load network: type=LoRA loaded={len(loaded_networks)} cache={list(lora_cache)}') + shared.log.debug(f'Network load: type=LoRA loaded={len(loaded_networks)} cache={list(lora_cache)}') devices.torch_gc() if recompile_model: - shared.log.info("Load network: type=LoRA recompiling model") + shared.log.info("Network load: type=LoRA recompiling model") backup_lora_model = shared.compiled_model_state.lora_model if 'Model' in shared.opts.cuda_compile: shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) @@ -310,7 +310,7 @@ def network_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Li self.weight = torch.nn.Parameter(weights_backup.to(self.weight.device, copy=True)) self.freeze() elif getattr(self, "quant_type", None) in ['nf4', 'fp4']: - bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: device = self.weight.device self.weight = bnb.nn.Params4bit(weights_backup, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) @@ -339,7 +339,7 @@ def maybe_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if isinstance(self, torch.nn.MultiheadAttention): weights_backup = (self.in_proj_weight.clone().to(devices.cpu), self.out_proj.weight.clone().to(devices.cpu)) elif getattr(self.weight, "quant_type", None) in ['nf4', 'fp4']: - bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: with devices.inference_context(): weights_backup = bnb.functional.dequantize_4bit(self.weight, quant_state=self.weight.quant_state, quant_type=self.weight.quant_type, blocksize=self.weight.blocksize,).to(devices.cpu) @@ -390,7 +390,7 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn # inpainting model. zero pad updown to make channel[1] 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if getattr(self.weight, "quant_type", None) in ['nf4', 'fp4']: # or self.weight.numel() != updown.numel(): - bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: device = self.weight.device weight = bnb.functional.dequantize_4bit(self.weight, quant_state=self.weight.quant_state, quant_type=self.weight.quant_type, blocksize=self.weight.blocksize) diff --git a/modules/extra_networks.py b/modules/extra_networks.py index 420f3beda..e882b113c 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -87,7 +87,7 @@ def activate(p, extra_network_data=None, step=0, include=[], exclude=[]): stepwise = stepwise or is_stepwise(extra_network_args) functional = shared.opts.lora_functional if shared.opts.lora_force_diffusers and stepwise: - shared.log.warning("Load network: type=LoRA method=composable loader=diffusers not compatible") + shared.log.warning("Network load: type=LoRA method=composable loader=diffusers not compatible") stepwise = False shared.opts.data['lora_functional'] = stepwise or functional diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 357c5291f..2167f97ac 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -52,7 +52,7 @@ def prompt(p): all_tags = list(set(all_tags)) all_tags = [t for t in all_tags if t not in p.prompt] if len(all_tags) > 0: - shared.log.debug(f"Load network: type=LoRA tags={all_tags} max={shared.opts.lora_apply_tags} apply") + shared.log.debug(f"Network load: type=LoRA tags={all_tags} max={shared.opts.lora_apply_tags} apply") all_tags = ', '.join(all_tags) p.extra_generation_params["LoRA tags"] = all_tags if '_tags_' in p.prompt: @@ -129,7 +129,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): sd_model.loaded_loras = {} key = f'{",".join(include)}:{",".join(exclude)}' loaded = sd_model.loaded_loras.get(key, []) - # shared.log.trace(f'Load network: type=LoRA key="{key}" requested={requested} loaded={loaded}') + # shared.log.trace(f'Network load: type=LoRA key="{key}" requested={requested} loaded={loaded}') if (len(requested) == 0) or (len(requested) != len(loaded)): sd_model.loaded_loras[key] = requested return True @@ -153,7 +153,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if debug: import sys fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - debug_log(f'Load network: type=LoRA include={include} exclude={exclude} requested={requested} fn={fn}') + debug_log(f'Network load: type=LoRA include={include} exclude={exclude} requested={requested} fn={fn}') force_diffusers = network_overrides.check_override() if force_diffusers: @@ -166,18 +166,18 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if has_changed: networks.network_deactivate(include, exclude) networks.network_activate(include, exclude) - debug_log(f'Load network: type=LoRA previous={[n.name for n in networks.previously_loaded_networks]} current={[n.name for n in networks.loaded_networks]} changed') + debug_log(f'Network load: type=LoRA previous={[n.name for n in networks.previously_loaded_networks]} current={[n.name for n in networks.loaded_networks]} changed') if len(networks.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or force_diffusers) and step == 0: infotext(p) prompt(p) if (has_changed or force_diffusers) and len(include) == 0: # print only once - shared.log.info(f'Load network: type=LoRA apply={[n.name for n in networks.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') + shared.log.info(f'Network load: type=LoRA apply={[n.name for n in networks.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') def deactivate(self, p): if shared.native: networks.previously_loaded_networks = networks.loaded_networks.copy() - debug_log(f'Load network: type=LoRA active={[n.name for n in networks.previously_loaded_networks]} deactivate') + debug_log(f'Network load: type=LoRA active={[n.name for n in networks.previously_loaded_networks]} deactivate') if shared.native and len(networks.diffuser_loaded) > 0: if not (shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled is True): if hasattr(shared.sd_model, "unfuse_lora"): diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 9e981a234..f6e8acadb 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -45,11 +45,11 @@ module_types = [ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_default_multiplier) -> Union[network.Network, None]: t0 = time.time() name = name.replace(".", "_") - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') if not shared.native: return None if not hasattr(shared.sd_model, 'load_lora_weights'): - shared.log.error(f'Load network: type=LoRA class={shared.sd_model.__class__} does not implement load lora') + shared.log.error(f'Network load: type=LoRA class={shared.sd_model.__class__} does not implement load lora') return None try: shared.sd_model.load_lora_weights(network_on_disk.filename, adapter_name=name) @@ -58,9 +58,9 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ pass else: if 'The following keys have not been correctly renamed' in str(e): - shared.log.error(f'Load network: type=LoRA name="{name}" diffusers unsupported format') + shared.log.error(f'Network load: type=LoRA name="{name}" diffusers unsupported format') else: - shared.log.error(f'Load network: type=LoRA name="{name}" {e}') + shared.log.error(f'Network load: type=LoRA name="{name}" {e}') if debug: errors.display(e, "LoRA") return None @@ -79,7 +79,7 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: cached = lora_cache.get(name, None) if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') if cached is not None: return cached net = network.Network(name, network_on_disk) @@ -132,11 +132,11 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: else: net.modules[key] = net_module if len(keys_failed_to_match) > 0: - shared.log.warning(f'Load network: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') + shared.log.warning(f'Network load: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" unmatched={keys_failed_to_match}') + shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') else: - shared.log.debug(f'Load network: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} direct={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} direct={shared.opts.lora_fuse_diffusers}') if len(matched_networks) == 0: return None lora_cache[name] = net @@ -247,7 +247,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non if network_on_disk is not None: shorthash = getattr(network_on_disk, 'shorthash', '').lower() if debug: - shared.log.debug(f'Load network: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') try: if recompile_model: shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}") @@ -259,13 +259,13 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non net.mentioned_name = name network_on_disk.read_hash() except Exception as e: - shared.log.error(f'Load network: type=LoRA file="{network_on_disk.filename}" {e}') + shared.log.error(f'Network load: type=LoRA file="{network_on_disk.filename}" {e}') if debug: errors.display(e, 'LoRA') continue if net is None: failed_to_load_networks.append(name) - shared.log.error(f'Load network: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') + shared.log.error(f'Network load: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') continue if hasattr(shared.sd_model, 'embedding_db'): shared.sd_model.embedding_db.load_diffusers_embedding(None, net.bundle_embeddings) @@ -279,7 +279,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non lora_cache.pop(name, None) if not skip_lora_load and len(diffuser_loaded) > 0: - shared.log.debug(f'Load network: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') + shared.log.debug(f'Network load: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') try: t0 = time.time() shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) @@ -288,15 +288,15 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non shared.sd_model.unload_lora_weights() timer.activate += time.time() - t0 except Exception as e: - shared.log.error(f'Load network: type=LoRA {e}') + shared.log.error(f'Network load: type=LoRA {e}') if debug: errors.display(e, 'LoRA') if len(loaded_networks) > 0 and debug: - shared.log.debug(f'Load network: type=LoRA loaded={[n.name for n in loaded_networks]} cache={list(lora_cache)}') + shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in loaded_networks]} cache={list(lora_cache)}') if recompile_model: - shared.log.info("Load network: type=LoRA recompiling model") + shared.log.info("Network load: type=LoRA recompiling model") backup_lora_model = shared.compiled_model_state.lora_model if 'Model' in shared.opts.cuda_compile: shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) @@ -330,7 +330,7 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n self.network_weights_backup = None if getattr(weight, "quant_type", None) in ['nf4', 'fp4']: if bnb is None: - bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: with devices.inference_context(): if shared.opts.lora_fuse_diffusers: @@ -430,7 +430,7 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) self.weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) except Exception as e: - shared.log.error(f'Load network: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') + shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') else: try: new_weight = model_weights.to(devices.device) + lora_weights.to(devices.device) @@ -556,7 +556,7 @@ def network_deactivate(include=[], exclude=[]): timer.deactivate = time.time() - t0 if debug and len(previously_loaded_networks) > 0: weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403 # pylint: disable=R1718 - shared.log.debug(f'Deactivate network: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() if shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") @@ -619,7 +619,7 @@ def network_activate(include=[], exclude=[]): timer.activate += time.time() - t0 if debug and len(loaded_networks) > 0: weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403 # pylint: disable=R1718 - shared.log.debug(f'Load network: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() if shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index b78b1e6a1..cb90a5950 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -80,10 +80,13 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {} ip_adapter_starts = list(p.ip_adapter_starts) ip_adapter_ends = list(p.ip_adapter_ends) if any(end != 1 for end in ip_adapter_ends) or any(start != 0 for start in ip_adapter_starts): - for i in range(len(ip_adapter_scales)): - ip_adapter_scales[i] *= float(step >= pipe.num_timesteps * ip_adapter_starts[i]) - ip_adapter_scales[i] *= float(step <= pipe.num_timesteps * ip_adapter_ends[i]) - debug_callback(f"Callback: IP Adapter scales={ip_adapter_scales}") + if 'Flux' in pipe.__class__.__name__: + ip_adapter_scales = [(ip_adapter_starts[0] + (ip_adapter_ends[0] - ip_adapter_starts[0]) * (i / (19 - 1))) for i in range(19)] + else: + for i in range(len(ip_adapter_scales)): + ip_adapter_scales[i] *= float(step >= pipe.num_timesteps * ip_adapter_starts[i]) + ip_adapter_scales[i] *= float(step <= pipe.num_timesteps * ip_adapter_ends[i]) + debug_callback(f"Callback: IP Adapter scales={ip_adapter_scales}") pipe.set_ip_adapter_scale(ip_adapter_scales) if step != getattr(pipe, 'num_timesteps', 0): kwargs = processing_correction.correction_callback(p, timestep, kwargs, initial=step == 0) diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index acea34874..a5f6b87c9 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -107,10 +107,7 @@ def process_base(p: processing.StableDiffusionProcessing): if hasattr(output, 'images'): shared.history.add(output.images, info=processing.create_infotext(p), ops=p.ops) timer.process.record('pipeline') - ras.unapply(shared.sd_model) - hidiffusion.unapply() sd_models_compile.openvino_post_compile(op="base") # only executes on compiled vino models - sd_models_compile.check_deepcache(enable=False) if shared.cmd_opts.profile: t1 = time.time() shared.log.debug(f'Profile: pipeline call: {t1-t0:.2f}') @@ -142,6 +139,10 @@ def process_base(p: processing.StableDiffusionProcessing): shared.log.error(f'Processing: step=base args={err_args} {e}') errors.display(e, 'Processing') modelstats.analyze() + finally: + ras.unapply(shared.sd_model) + hidiffusion.unapply() + sd_models_compile.check_deepcache(enable=False) if hasattr(shared.sd_model, 'embedding_db') and len(shared.sd_model.embedding_db.embeddings_used) > 0: # register used embeddings p.extra_generation_params['Embeddings'] = ', '.join(shared.sd_model.embedding_db.embeddings_used) diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 86c7cd260..a3f16ab9c 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -421,7 +421,7 @@ class EmbeddingDatabase: if self.previously_displayed_embeddings != displayed_embeddings and shared.opts.diffusers_enable_embed: self.previously_displayed_embeddings = displayed_embeddings t1 = time.time() - shared.log.info(f"Load network: type=embeddings loaded={len(self.word_embeddings)} skipped={len(self.skipped_embeddings)} time={t1-t0:.2f}") + shared.log.info(f"Network load: type=embeddings loaded={len(self.word_embeddings)} skipped={len(self.skipped_embeddings)} time={t1-t0:.2f}") def find_embedding_at_position(self, tokens, offset): From f9a692daf604807b04f2c5936e50a107f6a409bd Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 15 Mar 2025 22:10:41 -0400 Subject: [PATCH 006/122] update changelog and wiki Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 18 +++++++++++------- modules/scripts_postprocessing.py | 2 +- wiki | 2 +- 3 files changed, 13 insertions(+), 9 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 133e0d39c..4510cc983 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,10 +13,10 @@ simply select from *networks -> models -> reference* *note* cogview4 is compatible with flowmatching samplers - New [zer0int CLiP-L](https://huggingface.co/zer0int/CLIP-Registers-Gated_MLP-ViT-L-14) models: - download text encoders into folder set in settings -> system paths -> text encoders (default is `models/Text-encoder`) - load using *settings -> text encoder* + download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) + load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui -- **Remote VAE** +- **Remote VAE** - add support for remote vae encode in addition to remote vae decode - used by *img2img, inpaint, hires, detailer* - remote vae encode is disabled by default, you can enable it in *settings -> variable auto-encoder* @@ -24,14 +24,18 @@ - [Google Gemma 3 4B](https://huggingface.co/google/gemma-3-4b-it) simply select from list of available models in caption tab - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* -- **Wiki/Docs** +- **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - - Updated SD3 + - updated SD3 content +- [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) updates + - add detection of prompt language and alphabet and filter based on those values + - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) + against top-10 standard harmful content categories - **Other** - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** - - update `diffusers` - - remote vae use `scaling_factor` and `shift_factor` + - update `diffusers` and other requirements + - remote vae use `scaling_factor` and `shift_factor` - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/scripts_postprocessing.py b/modules/scripts_postprocessing.py index 3afaf68d3..da65ee551 100644 --- a/modules/scripts_postprocessing.py +++ b/modules/scripts_postprocessing.py @@ -65,7 +65,7 @@ class ScriptPostprocessingRunner: script.args_from = len(inputs) script.args_to = len(inputs) script.controls = wrap_call(script.ui, script.filename, "ui") - for control in script.controls.values(): + for control in script.controls.values() if script.controls is not None else []: control.custom_script_source = os.path.basename(script.filename) inputs += list(script.controls.values()) script.args_to = len(inputs) diff --git a/wiki b/wiki index cba8d182b..3676f5628 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit cba8d182b3cb18aeb4519d7c211d14cf9532021c +Subproject commit 3676f5628e0a6048f21cac03ada9065a1500a354 From a3e10fd24c5134b69344268548875cbafed2d3ad Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 11:22:29 -0400 Subject: [PATCH 007/122] fix progress eta reporting Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + cli/api-progress.py | 19 +++++++++++++------ modules/api/server.py | 4 ++-- 3 files changed, 16 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4510cc983..d9320ff9e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -43,6 +43,7 @@ - fix memory stats not displaying the ram usage - fix **RunPod** memory limit reporting - fix flux ipadapter with start/stop values + - fix progress api `eta_relative` - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/cli/api-progress.py b/cli/api-progress.py index 00ed618d2..cb81293da 100755 --- a/cli/api-progress.py +++ b/cli/api-progress.py @@ -18,7 +18,7 @@ class Dot(dict): opts = Dot({ "timeout": 3600, - "frequency": 60, + "frequency": 1, "action": "sudo shutdown now", "url": "http://127.0.0.1:7860", "user": "", @@ -46,15 +46,22 @@ log.info(f'sdnext monitor started: {opts}') while True: try: status = progress() + # {'progress': 0.0, 'eta_relative': 0.0, 'state': {'skipped': False, 'interrupted': False, 'job': '', 'job_count': 0, 'job_timestamp': '20250316110822', 'job_no': 0, 'sampling_step': 20, 'sampling_steps': 20}, 'current_image': None, 'textinfo': None} state = status.get('state', {}) - last_job = state.get('job_timestamp', None) - if last_job is None: + job_timestamp = state.get('job_timestamp', None) + job_progress = status.get('progress', 0) + eta_relative = status.get('eta_relative', 0) + job = state.get('job', '') + job_timestamp = state.get('job_timestamp', None) + sampling_step = state.get('sampling_step', 0) + sampling_steps = state.get('sampling_steps', 0) + if job_timestamp is None: log.warning(f'sdnext montoring cannot get last job info: {status}') else: - last_job = datetime.datetime.strptime(last_job, "%Y%m%d%H%M%S") - elapsed = datetime.datetime.now() - last_job + job_timestamp = datetime.datetime.strptime(job_timestamp, "%Y%m%d%H%M%S") if job_timestamp != '0' else datetime.datetime.now() + elapsed = datetime.datetime.now() - job_timestamp timeout = round(opts.timeout - elapsed.total_seconds()) - log.info(f'sdnext: last_job={last_job} elapsed={elapsed} timeout={timeout}') + log.info(f'sdnext: last="{job_timestamp}" elapsed={elapsed} timeout={timeout} progress={job_progress} eta={eta_relative} step={sampling_step}/{sampling_steps} job="{job}"') if timeout < 0: log.warning(f'sdnext reached: timeout={opts.timeout} action={opts.action}') os.system(opts.action) diff --git a/modules/api/server.py b/modules/api/server.py index 828d3fd95..28c429e29 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -91,10 +91,10 @@ def get_progress(req: models.ReqProgress = Depends()): step_y = max(shared.state.sampling_steps, 1) current = step_y * batch_x + step_x total = step_y * batch_y - progress = current / total if current > 0 and total > 0 else 0 + progress = min((current / total) if current > 0 and total > 0 else 0, 1) time_since_start = time.time() - shared.state.time_start eta_relative = (time_since_start / progress) - time_since_start if progress > 0 else 0 - res = models.ResProgress(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo) + res = models.ResProgress(progress=round(progress, 2), eta_relative=round(eta_relative, 2), current_image=current_image, textinfo=shared.state.textinfo, state=shared.state.dict(), ) return res def get_status(): From f113efc6f55f8990f4d8c752bdb1421beb55dbb1 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 12:14:13 -0400 Subject: [PATCH 008/122] insightface loader Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 ++++-- modules/face/faceswap.py | 14 ++++++++++---- modules/face/insightface.py | 3 ++- 3 files changed, 16 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d9320ff9e..a0869ca8d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,10 +27,11 @@ - **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - updated SD3 content -- [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) updates - - add detection of prompt language and alphabet and filter based on those values +- [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) extension updates + - add detection of prompt language and alphabet and filter based on those values - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) against top-10 standard harmful content categories + - add banned words/expressions check against prompt variations - **Other** - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** @@ -44,6 +45,7 @@ - fix **RunPod** memory limit reporting - fix flux ipadapter with start/stop values - fix progress api `eta_relative` + - fix `insightface` loader - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/modules/face/faceswap.py b/modules/face/faceswap.py index d7d0a32a5..df3765fb2 100644 --- a/modules/face/faceswap.py +++ b/modules/face/faceswap.py @@ -13,13 +13,19 @@ swapper = None def face_swap(p: processing.StableDiffusionProcessing, app, input_images: List[Image.Image], source_image: Image.Image, cache: bool): - import insightface.model_zoo global swapper # pylint: disable=global-statement if swapper is None: - model_path = hf.hf_hub_download(repo_id='ezioruan/inswapper_128.onnx', filename='inswapper_128.onnx', cache_dir=shared.opts.hfcache_dir) + import insightface.model_zoo + repo_id = 'ezioruan/inswapper_128.onnx' + model_path = hf.hf_hub_download(repo_id=repo_id, filename='inswapper_128.onnx', cache_dir=shared.opts.hfcache_dir) + shared.log.debug(f'FaceSwap load: repo="{repo_id}" path="{model_path}"') # model_path = hf.hf_hub_download(repo_id='somanchiu/reswapper', filename='reswapper_256-1567500_originalInswapperClassCompatible.onnx', cache_dir=shared.opts.hfcache_dir) - router: insightface.model_zoo.model_zoo.INSwapper = insightface.model_zoo.model_zoo.ModelRouter(model_path) - swapper = router.get_model() + try: + router: insightface.model_zoo.model_zoo.INSwapper = insightface.model_zoo.model_zoo.ModelRouter(model_path) + swapper = router.get_model() + except Exception as e: + shared.log.error(f'FaceSwap load: {e}') + return None np_image = cv2.cvtColor(np.array(source_image), cv2.COLOR_RGB2BGR) faces = app.get(np_image) diff --git a/modules/face/insightface.py b/modules/face/insightface.py index 529e4be32..7c9ec5af9 100644 --- a/modules/face/insightface.py +++ b/modules/face/insightface.py @@ -20,13 +20,14 @@ def get_app(mp_name, threshold=0.5, resolution=640): install('git+https://github.com/tencent-ailab/IP-Adapter.git', 'ip_adapter', ignore=False) if insightface_app is None or mp_name != instightface_mp: + import insightface from insightface.model_zoo import model_zoo from insightface.app import face_analysis model_zoo.print = lambda *args, **kwargs: None face_analysis.print = lambda *args, **kwargs: None import huggingface_hub as hf import zipfile - log.debug(f"InsightFace: mp={mp_name} provider={devices.onnx}") + log.debug(f"InsightFace: version={insightface.__version__} mp={mp_name} provider={devices.onnx}") root_dir = os.path.join(opts.diffusers_dir, 'models--vladmandic--insightface-faceanalysis') local_dir = os.path.join(root_dir, 'models') extract_dir = os.path.join(local_dir, mp_name) From 86e55ef11d4734efc47e70023ba865702608edd3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 12:27:41 -0400 Subject: [PATCH 009/122] rename vae none to default Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/model_flux.py | 4 ++-- modules/model_sd3.py | 4 ++-- modules/sd_vae.py | 2 +- modules/shared_items.py | 2 +- 5 files changed, 7 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a0869ca8d..a5182571d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -37,6 +37,7 @@ - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements - remote vae use `scaling_factor` and `shift_factor` + - rename vae *None* to *Default* to avoid confusion - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/model_flux.py b/modules/model_flux.py index 190007819..8e102afeb 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -265,7 +265,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch if debug: from modules import errors errors.display(e, 'FLUX T5:') - if shared.opts.sd_vae != 'None' and shared.opts.sd_vae != 'Automatic': + if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': try: debug(f'Load model: type=FLUX vae="{shared.opts.sd_vae}"') from modules import sd_vae @@ -276,7 +276,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch vae = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, **diffusers_load_config) except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}") - shared.opts.sd_vae = 'None' + shared.opts.sd_vae = 'Default' if debug: from modules import errors errors.display(e, 'FLUX VAE:') diff --git a/modules/model_sd3.py b/modules/model_sd3.py index bf8644284..962a6d6db 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -37,7 +37,7 @@ def load_overrides(kwargs, cache_dir): except Exception as e: shared.log.error(f"Load model: type=SD3 failed to load T5: {e}") shared.opts.sd_text_encoder = 'None' - if shared.opts.sd_vae != 'None' and shared.opts.sd_vae != 'Automatic': + if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': try: from modules import sd_vae vae_file = sd_vae.vae_dict[shared.opts.sd_vae] @@ -47,7 +47,7 @@ def load_overrides(kwargs, cache_dir): shared.log.debug(f'Load model: type=SD3 vae="{shared.opts.sd_vae}"') except Exception as e: shared.log.error(f"Load model: type=SD3 failed to load VAE: {e}") - shared.opts.sd_vae = 'None' + shared.opts.sd_vae = 'Default' return kwargs diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 8c34e29a8..17faf12d8 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -117,7 +117,7 @@ def resolve_vae(checkpoint_file): return None, None if shared.cmd_opts.vae is not None: # 1st return shared.cmd_opts.vae, 'forced' - if shared.opts.sd_vae == "None": # 2nd + if shared.opts.sd_vae == "Default": # 2nd return None, None vae_near_checkpoint = find_vae_near_checkpoint(checkpoint_file) if vae_near_checkpoint is not None: # 3rd diff --git a/modules/shared_items.py b/modules/shared_items.py index 7b9940a45..91299850e 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -5,7 +5,7 @@ def postprocessing_scripts(): def sd_vae_items(): import modules.sd_vae - return ["Automatic", "None"] + list(modules.sd_vae.vae_dict) + return ["Automatic", "Default"] + list(modules.sd_vae.vae_dict) def sd_taesd_items(): From 2c726c1d5e2f12e4496acacf2a1f2b2f4970d61c Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 12:42:13 -0400 Subject: [PATCH 010/122] fix remote vae for flux Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 3 ++- modules/sd_vae_remote.py | 6 ++++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a5182571d..fef78d91c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -46,7 +46,8 @@ - fix **RunPod** memory limit reporting - fix flux ipadapter with start/stop values - fix progress api `eta_relative` - - fix `insightface` loader + - fix `insightface` loader + - fix remove vae for flux.1 - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/modules/sd_vae_remote.py b/modules/sd_vae_remote.py index 55b8aa7ff..c3591af7b 100644 --- a/modules/sd_vae_remote.py +++ b/modules/sd_vae_remote.py @@ -42,7 +42,9 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ for i in range(latents.shape[0]): try: - latent = latents[i].detach().clone().to(device=devices.cpu, dtype=devices.dtype).unsqueeze(0) + latent = latents[i].detach().clone().to(device=devices.cpu, dtype=devices.dtype) + if model_type != 'f1': + latent = latent.unsqueeze(0) params = { "input_tensor_type": "binary", "shape": list(latent.shape), @@ -76,7 +78,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ timeout=300, ) if not response.ok: - shared.log.error(f'Decode: type="remote" model={model_type} code={response.status_code} headers={response.headers} {response.json()}') + shared.log.error(f'Decode: type="remote" model={model_type} code={response.status_code} shape={latent.shape} url="{url}" args={params} headers={response.headers} response={response.json()}') else: content += len(response.content) if shared.opts.remote_vae_type == 'raw': From 281afcde1a21ebf1b3e60e9cd56934d2cf5717c8 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 12:54:17 -0400 Subject: [PATCH 011/122] guard against git returining invalid timestamp Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + installer.py | 8 +++++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index fef78d91c..367a70c75 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -48,6 +48,7 @@ - fix progress api `eta_relative` - fix `insightface` loader - fix remove vae for flux.1 + - guard against git returining invalid timestamp - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/installer.py b/installer.py index ca93e49df..fb54f25e4 100644 --- a/installer.py +++ b/installer.py @@ -1433,14 +1433,16 @@ def check_timestamp(): if 'Setup complete without errors' in line: setup_time = int(line.split(' ')[-1]) try: - version_time = int(git('log -1 --pretty=format:"%at"')) + version_time = git('log -1 --pretty=format:"%at"') + version_time = ''.join(filter(str.isdigit, version_time)) + version_time = int(version_time) if len(version_time) > 0 else -1 + log.debug(f'Timestamp repository update time: {time.ctime(version_time)}') except Exception as e: log.error(f'Timestamp local repository version: {e}') - log.debug(f'Timestamp repository update time: {time.ctime(int(version_time))}') if setup_time == -1: return False log.debug(f'Timestamp previous setup time: {time.ctime(setup_time)}') - if setup_time < version_time: + if setup_time < version_time or version_time == -1: ok = False extension_time = check_extensions() log.debug(f'Timestamp latest extensions time: {time.ctime(extension_time)}') From a4b26e7ddbada41bb521a3c1aecf4707a50497b7 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 13:25:26 -0400 Subject: [PATCH 012/122] fix hires with latent upscale Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/processing_args.py | 2 ++ modules/processing_helpers.py | 22 +++++++++++----------- 3 files changed, 14 insertions(+), 11 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 367a70c75..5386b3565 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -49,6 +49,7 @@ - fix `insightface` loader - fix remove vae for flux.1 - guard against git returining invalid timestamp + - fix hires with latent upscale - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/modules/processing_args.py b/modules/processing_args.py index c0e201f77..261a996b0 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -298,6 +298,8 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t p.init_images = kwargs['image'] if isinstance(kwargs['image'], Image.Image): p.init_images = [kwargs['image']] + if isinstance(kwargs['image'], torch.Tensor): + p.init_images = kwargs['image'] # handle remaining args for arg in kwargs: diff --git a/modules/processing_helpers.py b/modules/processing_helpers.py index a5abba6a2..8fffa8313 100644 --- a/modules/processing_helpers.py +++ b/modules/processing_helpers.py @@ -401,24 +401,24 @@ def resize_init_images(p): def resize_hires(p, latents): # input=latents output=pil if not latent_upscaler else latent if not torch.is_tensor(latents): shared.log.warning('Hires: input is not tensor') - first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, vae_type=p.vae_type, output_type='pil', width=p.width, height=p.height) - return first_pass_images + decoded = processing_vae.vae_decode(latents=latents, model=shared.sd_model, vae_type=p.vae_type, output_type='pil', width=p.width, height=p.height) + return decoded if (p.hr_upscale_to_x == 0 or p.hr_upscale_to_y == 0) and hasattr(p, 'init_hr'): shared.log.error('Hires: missing upscaling dimensions') - return first_pass_images + return decoded if p.hr_upscaler.lower().startswith('latent'): - resized_image = images.resize_image(p.hr_resize_mode, latents, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler, context=p.hr_resize_context) - return resized_image + resized = images.resize_image(p.hr_resize_mode, latents, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler, context=p.hr_resize_context) + return resized - first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, vae_type=p.vae_type, output_type='pil', width=p.width, height=p.height) - resized_images = [] - for img in first_pass_images: - resized_image = images.resize_image(p.hr_resize_mode, img, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler, context=p.hr_resize_context) - resized_images.append(resized_image) + decoded = processing_vae.vae_decode(latents=latents, model=shared.sd_model, vae_type=p.vae_type, output_type='pil', width=p.width, height=p.height) + resized = [] + for image in decoded: + resize = images.resize_image(p.hr_resize_mode, image, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler, context=p.hr_resize_context) + resized.append(resize) devices.torch_gc() - return resized_images + return resized def fix_prompts(p, prompts, negative_prompts, prompts_2, negative_prompts_2): From ed602b173db337a74d63d90364c1bec3b69c8cd0 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 13:37:32 -0400 Subject: [PATCH 013/122] fix legacy diffusion latent upscalers Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/postprocess/sdupscaler_model.py | 10 +++++++--- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5386b3565..3f8557e8d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -50,6 +50,7 @@ - fix remove vae for flux.1 - guard against git returining invalid timestamp - fix hires with latent upscale + - fix legacy diffusion latent upscalers - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/modules/postprocess/sdupscaler_model.py b/modules/postprocess/sdupscaler_model.py index 5ec7168d3..7b8c7a8ca 100644 --- a/modules/postprocess/sdupscaler_model.py +++ b/modules/postprocess/sdupscaler_model.py @@ -24,15 +24,19 @@ class UpscalerDiffusion(Upscaler): def load_model(self, path: str): from modules.sd_models import set_diffuser_options - scaler: UpscalerData = [x for x in self.scalers if x.data_path == path][0] + scaler: UpscalerData = [x for x in self.scalers if x.data_path == path or x.name == path] + if len(scaler) == 0: + shared.log.error(f"Upscaler cannot match model: type={self.name} model={path}") + return None + scaler = scaler[0] if self.models.get(path, None) is not None: shared.log.debug(f"Upscaler cached: type={scaler.name} model={path}") return self.models[path] else: - model = diffusers.DiffusionPipeline.from_pretrained(path, cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype) + model = diffusers.DiffusionPipeline.from_pretrained(scaler.data_path, cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype) if hasattr(model, "set_progress_bar_config"): model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + 'Upscale', ncols=80, colour='#327fba') - set_diffuser_options(scaler.model, vae=None, op='upscaler') + set_diffuser_options(model, vae=None, op='upscaler') self.models[path] = model return self.models[path] From 942553a504981be23caefbdde7d730cd2d9721b5 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 18:46:26 -0400 Subject: [PATCH 014/122] rename vae and unet none to default Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 2 +- modules/model_flux.py | 12 ++++++------ modules/model_sd3.py | 8 ++++---- modules/model_stablecascade.py | 4 ++-- modules/processing_info.py | 4 ++-- modules/sd_models.py | 2 +- modules/sd_unet.py | 2 +- modules/shared.py | 4 ++-- 8 files changed, 19 insertions(+), 19 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 3f8557e8d..c659343df 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -37,7 +37,7 @@ - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements - remote vae use `scaling_factor` and `shift_factor` - - rename vae *None* to *Default* to avoid confusion + - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/model_flux.py b/modules/model_flux.py index 8e102afeb..79d81d031 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -216,7 +216,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change transformer = diffusers.FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config) if transformer is None: shared.log.error('Failed to load UNet model') - shared.opts.sd_unet = 'None' + shared.opts.sd_unet = 'Default' return transformer @@ -238,20 +238,20 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch devices.torch_gc(force=True) # load overrides if any - if shared.opts.sd_unet != 'None': + if shared.opts.sd_unet != 'Default': try: debug(f'Load model: type=FLUX unet="{shared.opts.sd_unet}"') transformer = load_transformer(sd_unet.unet_dict[shared.opts.sd_unet]) if transformer is None: - shared.opts.sd_unet = 'None' + shared.opts.sd_unet = 'Default' sd_unet.failed_unet.append(shared.opts.sd_unet) except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load UNet: {e}") - shared.opts.sd_unet = 'None' + shared.opts.sd_unet = 'Default' if debug: from modules import errors errors.display(e, 'FLUX UNet:') - if shared.opts.sd_text_encoder != 'None': + if shared.opts.sd_text_encoder != 'Default': try: debug(f'Load model: type=FLUX te="{shared.opts.sd_text_encoder}"') from modules.model_te import load_t5, load_vit_l @@ -261,7 +261,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load T5: {e}") - shared.opts.sd_text_encoder = 'None' + shared.opts.sd_text_encoder = 'Default' if debug: from modules import errors errors.display(e, 'FLUX T5:') diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 962a6d6db..baf936c1a 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -5,7 +5,7 @@ from modules import shared, devices, sd_models, sd_unet, model_quant, model_tool def load_overrides(kwargs, cache_dir): - if shared.opts.sd_unet != 'None': + if shared.opts.sd_unet != 'Default': try: fn = sd_unet.unet_dict[shared.opts.sd_unet] if fn.endswith('.safetensors'): @@ -20,9 +20,9 @@ def load_overrides(kwargs, cache_dir): shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}" fmt=gguf') except Exception as e: shared.log.error(f"Load model: type=SD3 failed to load UNet: {e}") - shared.opts.sd_unet = 'None' + shared.opts.sd_unet = 'Default' sd_unet.failed_unet.append(shared.opts.sd_unet) - if shared.opts.sd_text_encoder != 'None': + if shared.opts.sd_text_encoder != 'Default': try: from modules.model_te import load_t5, load_vit_l, load_vit_g if 'vit-l' in shared.opts.sd_text_encoder.lower(): @@ -36,7 +36,7 @@ def load_overrides(kwargs, cache_dir): shared.log.debug(f'Load model: type=SD3 variant="t5" te="{shared.opts.sd_text_encoder}"') except Exception as e: shared.log.error(f"Load model: type=SD3 failed to load T5: {e}") - shared.opts.sd_text_encoder = 'None' + shared.opts.sd_text_encoder = 'Default' if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': try: from modules import sd_vae diff --git a/modules/model_stablecascade.py b/modules/model_stablecascade.py index 3c3339dca..0c767d33b 100644 --- a/modules/model_stablecascade.py +++ b/modules/model_stablecascade.py @@ -93,7 +93,7 @@ def load_cascade_combined(checkpoint_info, diffusers_load_config): if 'cascade' in checkpoint_info.name.lower(): diffusers_load_config["variant"] = 'bf16' - if shared.opts.sd_unet != "None" or 'stabilityai' in checkpoint_info.name.lower(): + if shared.opts.sd_unet != "Default" or 'stabilityai' in checkpoint_info.name.lower(): if 'cascade' in checkpoint_info.name and ('lite' in checkpoint_info.name or (checkpoint_info.hash is not None and 'abc818bb0d' in checkpoint_info.hash)): decoder_folder = 'decoder_lite' prior_folder = 'prior_lite' @@ -107,7 +107,7 @@ def load_cascade_combined(checkpoint_info, diffusers_load_config): decoder = StableCascadeDecoderPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, text_encoder=None, **diffusers_load_config) # shared.log.debug(f'StableCascade {decoder_folder}: scale={decoder.latent_dim_scale}') prior_text_encoder = None - if shared.opts.sd_unet != "None": + if shared.opts.sd_unet != "Default": prior_unet, prior_text_encoder = load_prior(unet_dict[shared.opts.sd_unet]) else: prior_unet = StableCascadeUNet.from_pretrained("stabilityai/stable-cascade-prior", subfolder=prior_folder, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) diff --git a/modules/processing_info.py b/modules/processing_info.py index fa084a2fb..5677a538c 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -87,8 +87,8 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No args['Grid'] = grid if shared.native: args['Pipeline'] = shared.sd_model.__class__.__name__ - args['TE'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'None') else shared.opts.sd_text_encoder - args['UNet'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_unet is None or shared.opts.sd_unet == 'None') else shared.opts.sd_unet + args['TE'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'Default') else shared.opts.sd_text_encoder + args['UNet'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_unet is None or shared.opts.sd_unet == 'Default') else shared.opts.sd_unet if 'txt2img' in p.ops: args["Variation seed"] = all_subseeds[index] if p.subseed_strength > 0 else None args["Variation strength"] = p.subseed_strength if p.subseed_strength > 0 else None diff --git a/modules/sd_models.py b/modules/sd_models.py index 0d662f9df..8c4356793 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -948,7 +948,7 @@ def get_native(pipe: diffusers.DiffusionPipeline): def reload_text_encoder(initial=False): - if initial and (shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'None'): + if initial and (shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'Default'): return # dont unload signature = get_signature(shared.sd_model) t5 = [k for k, v in signature.items() if 'T5EncoderModel' in str(v)] diff --git a/modules/sd_unet.py b/modules/sd_unet.py index deb0b24b0..cfba470a1 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -10,7 +10,7 @@ debug = os.environ.get('SD_LOAD_DEBUG', None) is not None def load_unet(model): global loaded_unet # pylint: disable=global-statement - if shared.opts.sd_unet == 'None': + if shared.opts.sd_unet == 'Default': return if shared.opts.sd_unet not in list(unet_dict): shared.log.error(f'UNet model not found: {shared.opts.sd_unet}') diff --git a/modules/shared.py b/modules/shared.py index ed3ef031d..4328eb787 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -389,7 +389,7 @@ options_templates.update(options_section(('sd', "Models & Loading"), { "diffusers_pipeline": OptionInfo('Autodetect', 'Model pipeline', gr.Dropdown, lambda: {"choices": list(shared_items.get_pipelines()), "visible": native}), "sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", DropdownEditable, lambda: {"choices": list_checkpoint_titles()}, refresh=refresh_checkpoints), "sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints), - "sd_unet": OptionInfo("None", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list), + "sd_unet": OptionInfo("Default", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list), "latent_history": OptionInfo(16, "Latent history size", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}), "offload_sep": OptionInfo("

Model Offloading

", "", gr.HTML), @@ -429,7 +429,7 @@ options_templates.update(options_section(('vae_encoder', "Variable Auto Encoder" })) options_templates.update(options_section(('text_encoder', "Text Encoder"), { - "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": shared_items.sd_te_items()}, refresh=shared_items.refresh_te_list), + "sd_text_encoder": OptionInfo('Default', "Text encoder model", gr.Dropdown, lambda: {"choices": shared_items.sd_te_items()}, refresh=shared_items.refresh_te_list), "prompt_attention": OptionInfo("native", "Prompt attention parser", gr.Radio, {"choices": ["native", "compel", "xhinker", "a1111", "fixed"] }), "prompt_mean_norm": OptionInfo(False, "Prompt attention normalization", gr.Checkbox), "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results", gr.Checkbox, {"visible": False}), From d4dff967b3a9cb4be0f19869fba9a89275703d88 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 19:51:38 -0400 Subject: [PATCH 015/122] asymmetric vae v2 and libvips support Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 3 ++ modules/upscaler_simple.py | 69 ++++++++++++++++++++++++++++--- scripts/postprocessing_upscale.py | 2 +- wiki | 2 +- 4 files changed, 69 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c659343df..fbc4fe616 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,6 +33,8 @@ against top-10 standard harmful content categories - add banned words/expressions check against prompt variations - **Other** + - **upscale**: new [asymmetric vae v2](Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method + - **upscale**: new experimental support for `libvips` upscaling - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements @@ -51,6 +53,7 @@ - guard against git returining invalid timestamp - fix hires with latent upscale - fix legacy diffusion latent upscalers + - fix upscaler selection in postprocessing - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - add xpu to profiler diff --git a/modules/upscaler_simple.py b/modules/upscaler_simple.py index a28d540f5..bb342d5c8 100644 --- a/modules/upscaler_simple.py +++ b/modules/upscaler_simple.py @@ -98,21 +98,29 @@ class UpscalerAsymmetricVAE(Upscaler): super().__init__(False) self.name = "Asymmetric VAE" self.vae = None + self.selected = None self.scalers = [ - UpscalerData("Asymmetric VAE", None, self), + UpscalerData("Asymmetric VAE v1", None, self), + UpscalerData("Asymmetric VAE v2", None, self), ] def do_upscale(self, img: Image, selected_model=None): + if selected_model is None: + return img import torchvision.transforms.functional as F import diffusers from modules import shared, devices - - if self.vae is None: - self.vae = diffusers.AsymmetricAutoencoderKL.from_pretrained("Heasterian/AsymmetricAutoencoderKLUpscaler", cache_dir=shared.opts.hfcache_dir) + if self.vae is None or selected_model != self.selected: + if 'v1' in selected_model: + repo_id = 'Heasterian/AsymmetricAutoencoderKLUpscaler' + else: + repo_id = 'Heasterian/AsymmetricAutoencoderKLUpscaler_v2' + self.vae = diffusers.AsymmetricAutoencoderKL.from_pretrained(repo_id, cache_dir=shared.opts.hfcache_dir) + shared.log.debug(f'Upscaler load: vae="{repo_id}"') self.vae.requires_grad_(False) self.vae = self.vae.to(device=devices.device, dtype=devices.dtype) self.vae.eval() - img = img.resize((8 * (img.width // 8), 8 * (img.height // 8)), resample=Image.Resampling.BILINEAR).convert('RGB') + img = img.resize((8 * (img.width // 8), 8 * (img.height // 8)), resample=Image.Resampling.LANCZOS).convert('RGB') tensor = (F.pil_to_tensor(img).unsqueeze(0) / 255.0).to(device=devices.device, dtype=devices.dtype) self.vae = self.vae.to(device=devices.device) tensor = self.vae(tensor).sample @@ -141,3 +149,54 @@ class UpscalerDCC(Upscaler): upscaled = (255.0 * upscaled).astype(np.uint8) upscaled = Image.fromarray(upscaled) return upscaled + + +class UpscalerVIPS(Upscaler): + def __init__(self, dirname=None): # pylint: disable=unused-argument + super().__init__(False) + self.name = "VIPS" + self.scalers = [ + UpscalerData("VIPS Lanczos 2", None, self), + UpscalerData("VIPS Lanczos 3", None, self), + UpscalerData("VIPS Mitchell", None, self), + UpscalerData("VIPS MagicKernelSharp 2013", None, self), + UpscalerData("VIPS MagicKernelSharp 2021", None, self), + ] + + def do_upscale(self, img: Image, selected_model=None): + if selected_model is None: + return img + from installer import install + from modules.shared import log + install('pyvips') + try: + import pyvips + except Exception as e: + log.error(f"Upscaler: vips {e}") + return img + vips_image = pyvips.Image.new_from_array(img) + # import numpy as np + # np_image = np.array(img) + # h, w, c = np_image.shape + # np_linear = np_image.reshape(w * h * c) + # vips_image = pyvips.Image.new_from_memory(np_linear.data, w, h, c, 'uchar') + try: + if selected_model is None: + return img + elif selected_model == "VIPS Lanczos 2": + vips_image = vips_image.resize(2, kernel='lanczos2') + elif selected_model == "VIPS Lanczos 3": + vips_image = vips_image.resize(2, kernel='lanczos3') + elif selected_model == "VIPS Mitchell": + vips_image = vips_image.resize(2, kernel='mitchell') + elif selected_model == "VIPS MagicKernelSharp 2013": + vips_image = vips_image.resize(2, kernel='mks2013') + elif selected_model == "VIPS MagicKernelSharp 2021": + vips_image = vips_image.resize(2, kernel='mks2021') + else: + return img + except Exception as e: + log.error(f"Upscaler: vips {e}") + return img + upscaled = Image.fromarray(vips_image.numpy()) + return upscaled diff --git a/scripts/postprocessing_upscale.py b/scripts/postprocessing_upscale.py index 104a0fb37..cffff8ed4 100644 --- a/scripts/postprocessing_upscale.py +++ b/scripts/postprocessing_upscale.py @@ -54,7 +54,7 @@ class ScriptPostprocessingUpscale(scripts_postprocessing.ScriptPostprocessing): info["Postprocess upscale to"] = f"{upscale_to_width}x{upscale_to_height}" else: info["Postprocess upscale by"] = upscale_by - image = upscaler.scaler.upscale(image, upscale_by, upscaler.data_path) + image = upscaler.scaler.upscale(image, upscale_by, upscaler.name) if upscale_mode == 1 and upscale_crop: cropped = Image.new("RGB", (upscale_to_width, upscale_to_height)) cropped.paste(image, box=(upscale_to_width // 2 - image.width // 2, upscale_to_height // 2 - image.height // 2)) diff --git a/wiki b/wiki index 3676f5628..62636f56b 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 3676f5628e0a6048f21cac03ada9065a1500a354 +Subproject commit 62636f56b6ee7b623eed50f150ad2dea77976948 From 4f56f4aa333f9dec618652fd261f9b0a1121c03f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 16 Mar 2025 21:45:05 -0400 Subject: [PATCH 016/122] add new optimum-quanto on-the-fly and simplify quantization loading Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 ++- cli/prompt-detect.py | 32 ------------- modules/model_cogview.py | 6 +-- modules/model_flux.py | 38 +++++----------- modules/model_lumina.py | 4 +- modules/model_quant.py | 97 +++++++++++++++++++++++++++++----------- modules/model_sana.py | 8 +--- modules/model_sd3.py | 12 +---- modules/model_tools.py | 6 +-- modules/shared.py | 8 +++- scripts/allegrovideo.py | 9 +--- scripts/hunyuanvideo.py | 9 +--- scripts/ltxvideo.py | 13 +----- scripts/mochivideo.py | 4 +- 14 files changed, 104 insertions(+), 148 deletions(-) delete mode 100644 cli/prompt-detect.py diff --git a/CHANGELOG.md b/CHANGELOG.md index fbc4fe616..e237f7f02 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,8 +33,10 @@ against top-10 standard harmful content categories - add banned words/expressions check against prompt variations - **Other** - - **upscale**: new [asymmetric vae v2](Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method - - **upscale**: new experimental support for `libvips` upscaling + - **upscale**: new [asymmetric vae v2](Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method + - **upscale**: new experimental support for `libvips` upscaling + - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models + note: previous method for quanto is still valid and is noted in settings as post-load quantization - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements diff --git a/cli/prompt-detect.py b/cli/prompt-detect.py deleted file mode 100644 index cd3d0189a..000000000 --- a/cli/prompt-detect.py +++ /dev/null @@ -1,32 +0,0 @@ -# Example: -# > python cli/lang-detect.py "have a good day" -# > ['eng_latn:1.00'] -# eng=language, latn=latin alphabet, 1.00=confidence - -import sys -import fasttext -from huggingface_hub import hf_hub_download - - -repo_id = "facebook/fasttext-language-identification" -model = None - - -def detect(text:str, top:int=1, threshold:float=0.25) -> str: - try: - global model # pylint: disable=global-statement - if model is None: - model_path = hf_hub_download(repo_id, filename="model.bin") - model = fasttext.load_model(model_path) - lang, score = model.predict(text, k=top, threshold=threshold, on_unicode_error="ignore") - result = [f"{l.replace("__label__", "").lower()}:{s:.2f}" for l, s in zip(lang, score) if s > threshold][:top] - return result - except Exception as e: - return str(e) - - -if __name__ == "__main__": - if len(sys.argv) < 2: - print(f"Usage: {sys.argv[0]} ") - else: - print(detect(sys.argv[1])) diff --git a/modules/model_cogview.py b/modules/model_cogview.py index 8ced40ce2..a0594cc16 100644 --- a/modules/model_cogview.py +++ b/modules/model_cogview.py @@ -18,11 +18,7 @@ def load_common(diffusers_load_config={}, module=None): if 'requires_safety_checker' in diffusers_load_config: del diffusers_load_config['requires_safety_checker'] - quant_args = {} - if not quant_args: - quant_args = model_quant.create_bnb_config(quant_args, module=module) - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args, module=module) + quant_args = model_quant.create_config(module=module) if quant_args: shared.log.debug(f'Load model: type=CogView quantization module="{module}" {quant_args}') diff --git a/modules/model_flux.py b/modules/model_flux.py index 79d81d031..97b39aedc 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -5,7 +5,7 @@ import diffusers import transformers from safetensors.torch import load_file from huggingface_hub import hf_hub_download -from modules import shared, devices, modelloader, sd_models, sd_unet, model_te, model_quant +from modules import shared, errors, devices, modelloader, sd_models, sd_unet, model_te, model_quant debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -44,7 +44,6 @@ def load_flux_quanto(checkpoint_info): except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load Quanto transformer: {e}") if debug: - from modules import errors errors.display(e, 'FLUX Quanto:') try: @@ -72,7 +71,6 @@ def load_flux_quanto(checkpoint_info): except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load Quanto text encoder: {e}") if debug: - from modules import errors errors.display(e, 'FLUX Quanto:') return transformer, text_encoder_2 @@ -105,33 +103,25 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu shared.log.error(f"Load model: type=FLUX failed to load BnB transformer: {e}") transformer, text_encoder_2 = None, None if debug: - from modules import errors errors.display(e, 'FLUX:') return transformer, text_encoder_2 def load_quants(kwargs, repo_id, cache_dir, allow_quant): try: - if not allow_quant: - return kwargs - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=FLUX quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=FLUX quant={quant_args}') + quant_args = model_quant.create_config(allow=allow_quant) if not quant_args: return kwargs - if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): + if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization): kwargs['transformer'] = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'text_encoder_2' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization): + quant_args = model_quant.create_config(allow=allow_quant, module='Text Encoder') + if not quant_args: + return kwargs + if 'text_encoder_2' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization or 'Text Encoder' in shared.opts.quanto_quantization): kwargs['text_encoder_2'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') except Exception as e: shared.log.error(f'Quantization: {e}') + errors.display(e, 'Quantization:') return kwargs @@ -197,15 +187,13 @@ def load_transformer(file_path): # triggered by opts.sd_unet change else: quant_args = model_quant.create_bnb_config({}) if quant_args: - model_quant.load_bnb(f'Load model: type=FLUX quant={quant_args}') shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant=bnb dtype={devices.dtype}') from modules.model_flux_nf4 import load_flux_nf4 transformer, _text_encoder_2 = load_flux_nf4(file_path, prequantized=False) if transformer is not None: return transformer - quant_args = model_quant.create_ao_config({}) + quant_args = model_quant.create_config() if quant_args: - model_quant.load_torchao(f'Load model: type=FLUX quant={quant_args}') shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant=torchao dtype={devices.dtype}') transformer = diffusers.FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config, **quant_args) if transformer is not None: @@ -249,7 +237,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.log.error(f"Load model: type=FLUX failed to load UNet: {e}") shared.opts.sd_unet = 'Default' if debug: - from modules import errors errors.display(e, 'FLUX UNet:') if shared.opts.sd_text_encoder != 'Default': try: @@ -263,7 +250,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.log.error(f"Load model: type=FLUX failed to load T5: {e}") shared.opts.sd_text_encoder = 'Default' if debug: - from modules import errors errors.display(e, 'FLUX T5:') if shared.opts.sd_vae != 'Default' and shared.opts.sd_vae != 'Automatic': try: @@ -278,7 +264,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}") shared.opts.sd_vae = 'Default' if debug: - from modules import errors errors.display(e, 'FLUX VAE:') # load quantized components if any @@ -293,7 +278,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load NF4 components: {e}") if debug: - from modules import errors errors.display(e, 'FLUX NF4:') if quant == 'qint8' or quant == 'qint4': try: @@ -305,7 +289,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch except Exception as e: shared.log.error(f"Load model: type=FLUX failed to load Quanto components: {e}") if debug: - from modules import errors errors.display(e, 'FLUX Quanto:') # initialize pipeline with pre-loaded components @@ -346,8 +329,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch fn = checkpoint_info.path if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)): kwargs = load_quants(kwargs, repo_id, cache_dir=shared.opts.diffusers_dir, allow_quant=allow_quant) - kwargs = model_quant.create_bnb_config(kwargs, allow_quant) - kwargs = model_quant.create_ao_config(kwargs, allow_quant) + # kwargs = model_quant.create_config(kwargs, allow_quant) if fn.endswith('.safetensors') and os.path.isfile(fn): pipe = diffusers.FluxPipeline.from_single_file(fn, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config) else: diff --git a/modules/model_lumina.py b/modules/model_lumina.py index 9dff4dccd..9e9ca8eff 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -32,9 +32,7 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): if quant_args: model_quant.load_bnb(f'Load model: type=Lumina quant={quant_args}') if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=Lumina quant={quant_args}') + quant_args = model_quant.create_config() kwargs = {} repo_id = sd_models.path_to_repo(checkpoint_info.name) if ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): diff --git a/modules/model_quant.py b/modules/model_quant.py index 127254343..0430dc9e0 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -1,3 +1,4 @@ +import os import sys import copy import time @@ -9,9 +10,9 @@ ao = None bnb = None intel_nncf = None optimum_quanto = None - quant_last_model_name = None quant_last_model_device = None +debug = os.environ.get('SD_QUANT_DEBUG', None) is not None def get_quant(name): @@ -44,7 +45,7 @@ def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Mode bnb_4bit_quant_type=shared.opts.bnb_quantization_type, bnb_4bit_compute_dtype=devices.dtype ) - shared.log.debug(f'Quantization: module=all type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') + log.debug(f'Quantization: module=all type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') if kwargs is None: return bnb_config else: @@ -60,9 +61,8 @@ def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model' load_torchao() if ao is None: return kwargs - diffusers.utils.import_utils.is_torchao_available = lambda: True ao_config = diffusers.TorchAoConfig(shared.opts.torchao_quantization_type) - shared.log.debug(f'Quantization: module=all type=torchao dtype={shared.opts.torchao_quantization_type}') + log.debug(f'Quantization: module=all type=torchao dtype={shared.opts.torchao_quantization_type}') if kwargs is None: return ao_config else: @@ -71,6 +71,47 @@ def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model' return kwargs +def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = 'Model'): + from modules import shared + if len(shared.opts.quanto_quantization) > 0 and allow_quanto: + if 'Model' in shared.opts.quanto_quantization or (module is not None and module in shared.opts.quanto_quantization): + load_quanto(silent=True) + if optimum_quanto is None: + return kwargs + quanto_config = diffusers.QuantoConfig( + weights_dtype=shared.opts.quanto_quantization_type, + ) + quanto_config.activations = None # patch so it works with transformers + log.debug(f'Quantization: module=all type=quanto dtype={shared.opts.quanto_quantization_type}') + if kwargs is None: + return quanto_config + else: + kwargs['quantization_config'] = quanto_config + return kwargs + return kwargs + + +def create_config(kwargs = None, allow: bool = True, module: str = 'Model'): + if kwargs is None: + kwargs = {} + kwargs = create_bnb_config(kwargs, allow_bnb=allow, module=module) + if kwargs is not None and 'quantization_config' in kwargs: + if debug: + log.trace(f'Quantization: type=bnb config={kwargs.get("quantization_config", None)}') + return kwargs + kwargs = create_ao_config(kwargs, allow_ao=allow, module=module) + if kwargs is not None and 'quantization_config' in kwargs: + if debug: + log.trace(f'Quantization: type=torchao config={kwargs.get("quantization_config", None)}') + return kwargs + kwargs = create_quanto_config(kwargs, allow_quanto=allow, module=module) + if kwargs is not None and 'quantization_config' in kwargs: + if debug: + log.trace(f'Quantization: type=quanto config={kwargs.get("quantization_config", None)}') + return kwargs + return kwargs + + def load_torchao(msg='', silent=False): global ao # pylint: disable=global-statement if ao is not None: @@ -81,6 +122,9 @@ def load_torchao(msg='', silent=False): ao = torchao fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.debug(f'Quantization: type=torchao version={ao.__version__} fn={fn}') # pylint: disable=protected-access + from diffusers.utils import import_utils + import_utils.is_torchao_available = lambda: True + import_utils._torchao_available = True # pylint: disable=protected-access return ao except Exception as e: if len(msg) > 0: @@ -102,9 +146,10 @@ def load_bnb(msg='', silent=False): try: import bitsandbytes bnb = bitsandbytes - diffusers.utils.import_utils._bitsandbytes_available = True # pylint: disable=protected-access - diffusers.utils.import_utils._bitsandbytes_version = '0.43.3' # pylint: disable=protected-access - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + from diffusers.utils import import_utils + import_utils._bitsandbytes_available = True # pylint: disable=protected-access + import_utils._bitsandbytes_version = '0.43.3' # pylint: disable=protected-access + fn = f'{sys._getframe(3).f_code.co_name}:{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.debug(f'Quantization: type=bitsandbytes version={bnb.__version__} fn={fn}') # pylint: disable=protected-access return bnb except Exception as e: @@ -117,18 +162,20 @@ def load_bnb(msg='', silent=False): def load_quanto(msg='', silent=False): - from modules import shared global optimum_quanto # pylint: disable=global-statement if optimum_quanto is not None: return optimum_quanto - install('optimum-quanto==0.2.6', quiet=True) + install('optimum-quanto==0.2.7', quiet=True) try: from optimum import quanto # pylint: disable=no-name-in-module optimum_quanto = quanto - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + fn = f'{sys._getframe(3).f_code.co_name}:{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access log.debug(f'Quantization: type=quanto version={quanto.__version__} fn={fn}') # pylint: disable=protected-access - if shared.opts.diffusers_offload_mode in {'balanced', 'sequential'}: - shared.log.error(f'Quantization: type=quanto offload={shared.opts.diffusers_offload_mode} not supported') + from diffusers.utils import import_utils + import_utils.is_optimum_quanto_available = lambda: True + import_utils._optimum_quanto_available = True # pylint: disable=protected-access + import_utils._optimum_quanto_version = quanto.__version__ # pylint: disable=protected-access + import_utils._replace_with_quanto_layers = diffusers.quantizers.quanto.utils._replace_with_quanto_layers # pylint: disable=protected-access return optimum_quanto except Exception as e: if len(msg) > 0: @@ -169,7 +216,7 @@ def apply_layerwise(sd_model, quiet:bool=False): storage_dtype = torch.float8_e5m2 else: storage_dtype = None - shared.log.warning(f'Quantization: type=layerwise storage={shared.opts.layerwise_quantization_storage} not supported') + log.warning(f'Quantization: type=layerwise storage={shared.opts.layerwise_quantization_storage} not supported') return non_blocking = False if not hasattr(quantization_config.QuantizationMethod, 'LAYERWISE'): @@ -198,7 +245,7 @@ def apply_layerwise(sd_model, quiet:bool=False): m.quantization_method = quantization_config.QuantizationMethod.LAYERWISE # pylint: disable=no-member log.quiet(quiet, f'Quantization: type=layerwise module={module} cls={cls} storage={storage_dtype} compute={devices.dtype} blocking={not non_blocking}') except Exception as e: - shared.log.error(f'Quantization: type=layerwise {e}') + log.error(f'Quantization: type=layerwise {e}') def nncf_send_to_device(model, device): @@ -244,7 +291,7 @@ def nncf_compress_weights(sd_model): try: t0 = time.time() from modules import shared, devices, sd_models - shared.log.info(f"Quantization: type=NNCF modules={shared.opts.nncf_compress_weights}") + log.info(f"Quantization: type=NNCF modules={shared.opts.nncf_compress_weights}") global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement sd_model = sd_models.apply_function_to_model(sd_model, nncf_compress_model, shared.opts.nncf_compress_weights, op="nncf") @@ -259,9 +306,9 @@ def nncf_compress_weights(sd_model): quant_last_model_device = None t1 = time.time() - shared.log.info(f"Quantization: type=NNCF time={t1-t0:.2f}") + log.info(f"Quantization: type=NNCF time={t1-t0:.2f}") except Exception as e: - shared.log.warning(f"Quantization: type=NNCF {e}") + log.warning(f"Quantization: type=NNCF {e}") return sd_model @@ -312,9 +359,9 @@ def optimum_quanto_weights(sd_model): t0 = time.time() from modules import shared, devices, sd_models if shared.opts.diffusers_offload_mode in {"balanced", "sequential"}: - shared.log.warning(f"Quantization: type=Optimum.quanto offload={shared.opts.diffusers_offload_mode} not compatible") + log.warning(f"Quantization: type=Optimum.quanto offload={shared.opts.diffusers_offload_mode} not compatible") return sd_model - shared.log.info(f"Quantization: type=Optimum.quanto: modules={shared.opts.optimum_quanto_weights}") + log.info(f"Quantization: type=Optimum.quanto: modules={shared.opts.optimum_quanto_weights}") global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement quanto = load_quanto() quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) @@ -361,9 +408,9 @@ def optimum_quanto_weights(sd_model): devices.torch_gc(force=True) t1 = time.time() - shared.log.info(f"Quantization: type=Optimum.quanto time={t1-t0:.2f}") + log.info(f"Quantization: type=Optimum.quanto time={t1-t0:.2f}") except Exception as e: - shared.log.warning(f"Quantization: type=Optimum.quanto {e}") + log.warning(f"Quantization: type=Optimum.quanto {e}") return sd_model @@ -374,19 +421,19 @@ def torchao_quantization(sd_model): fn = getattr(q, shared.opts.torchao_quantization_type, None) if fn is None: - shared.log.error(f"Quantization: type=TorchAO type={shared.opts.torchao_quantization_type} not supported") + log.error(f"Quantization: type=TorchAO type={shared.opts.torchao_quantization_type} not supported") return sd_model def torchao_model(model, op=None, sd_model=None): # pylint: disable=unused-argument q.quantize_(model, fn(), device=devices.device) return model - shared.log.info(f"Quantization: type=TorchAO pipe={sd_model.__class__.__name__} quant={shared.opts.torchao_quantization_type} fn={fn} targets={shared.opts.torchao_quantization}") + log.info(f"Quantization: type=TorchAO pipe={sd_model.__class__.__name__} quant={shared.opts.torchao_quantization_type} fn={fn} targets={shared.opts.torchao_quantization}") try: t0 = time.time() sd_models.apply_function_to_model(sd_model, torchao_model, shared.opts.torchao_quantization, op="torchao") t1 = time.time() - shared.log.info(f"Quantization: type=TorchAO time={t1-t0:.2f}") + log.info(f"Quantization: type=TorchAO time={t1-t0:.2f}") except Exception as e: - shared.log.error(f"Quantization: type=TorchAO {e}") + log.error(f"Quantization: type=TorchAO {e}") setup_logging() # torchao uses dynamo which messes with logging so reset is needed return sd_model diff --git a/modules/model_sana.py b/modules/model_sana.py index 79a13592d..54f2681fa 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -8,13 +8,7 @@ from modules import shared, sd_models, devices, modelloader, model_quant def load_quants(kwargs, repo_id, cache_dir): quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=Sana quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=Sana quant={quant_args}') + quant_args = model_quant.create_config() if not quant_args: return kwargs load_args = kwargs.copy() diff --git a/modules/model_sd3.py b/modules/model_sd3.py index baf936c1a..59789155e 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -52,14 +52,7 @@ def load_overrides(kwargs, cache_dir): def load_quants(kwargs, repo_id, cache_dir): - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=SD3 quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=SD3 quant={quant_args}') + quant_args = model_quant.create_config() if not quant_args: return kwargs if 'Model' in shared.opts.bnb_quantization and 'transformer' not in kwargs: @@ -157,8 +150,7 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): shared.log.debug(f'Load model: type=SD3 kwargs={list(kwargs)} repo="{repo_id}"') - kwargs = model_quant.create_bnb_config(kwargs) - kwargs = model_quant.create_ao_config(kwargs) + kwargs = model_quant.create_config(kwargs) pipe = loader( repo_id, torch_dtype=devices.dtype, diff --git a/modules/model_tools.py b/modules/model_tools.py index fdeda5c2a..8e473ba23 100644 --- a/modules/model_tools.py +++ b/modules/model_tools.py @@ -69,13 +69,11 @@ def load_modules(repo_id: str, params: dict): subfolder = 'text_encoder_2' if cls == transformers.T5EncoderModel: # t5-xxl subfolder = 'text_encoder_3' - kwargs = model_quant.create_bnb_config(kwargs) - kwargs = model_quant.create_ao_config(kwargs) + kwargs = model_quant.create_config(kwargs) kwargs['variant'] = 'fp16' if cls == diffusers.SD3Transformer2DModel: subfolder = 'transformer' - kwargs = model_quant.create_bnb_config(kwargs) - kwargs = model_quant.create_ao_config(kwargs) + kwargs = model_quant.create_config(kwargs) if subfolder is None: continue shared.log.debug(f'Load: module={name} class={cls.__name__} repo={repo_id} location={subfolder}') diff --git a/modules/shared.py b/modules/shared.py index 4328eb787..8b2b2b9b0 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -513,7 +513,11 @@ options_templates.update(options_section(('quantization', "Quantization Settings "bnb_quantization_type": OptionInfo("nf4", "Quantization type", gr.Dropdown, {"choices": ['nf4', 'fp8', 'fp4'], "visible": native}), "bnb_quantization_storage": OptionInfo("uint8", "Backend storage", gr.Dropdown, {"choices": ["float16", "float32", "int8", "uint8", "float64", "bfloat16"], "visible": native}), - "optimum_quanto_sep": OptionInfo("

Optimum Quanto

", "", gr.HTML), + "quanto_quantization_sep": OptionInfo("

Optimum Quanto

", "", gr.HTML), + "quanto_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder"], "visible": native}), + "quanto_quantization_type": OptionInfo("int8", "Quantization weights type", gr.Dropdown, {"choices": ["float8", "int8", "int4", "int2"], "visible": native}), + + "optimum_quanto_sep": OptionInfo("

Optimum Quanto: post-load

", "", gr.HTML), "optimum_quanto_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "ControlNet"], "visible": native}), "optimum_quanto_weights_type": OptionInfo("qint8", "Quantization weights type", gr.Dropdown, {"choices": ['qint8', 'qfloat8_e4m3fn', 'qfloat8_e5m2', 'qint4', 'qint2'], "visible": native}), "optimum_quanto_activations_type": OptionInfo("none", "Quantization activations type ", gr.Dropdown, {"choices": ['none', 'qint8', 'qfloat8_e4m3fn', 'qfloat8_e5m2'], "visible": native}), @@ -524,7 +528,7 @@ options_templates.update(options_section(('quantization', "Quantization Settings "torchao_quantization_mode": OptionInfo("pre", "Quantization mode", gr.Dropdown, {"choices": ['pre', 'post'], "visible": native}), "torchao_quantization_type": OptionInfo("int8_weight_only", "Quantization type", gr.Dropdown, {"choices": ['int4_weight_only', 'int8_dynamic_activation_int4_weight', 'int8_weight_only', 'int8_dynamic_activation_int8_weight', 'float8_weight_only', 'float8_dynamic_activation_float8_weight', 'float8_static_activation_float8_weight'], "visible": native}), - "nncf_compress_sep": OptionInfo("

NNCF

", "", gr.HTML), + "nncf_compress_sep": OptionInfo("

NNCF: Neural Network Compression Framework

", "", gr.HTML), "nncf_compress_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "ControlNet"], "visible": native}), "nncf_compress_weights_mode": OptionInfo("INT8", "Quantization type", gr.Dropdown, {"choices": ['INT8', 'INT8_SYM', 'INT4_ASYM', 'INT4_SYM', 'NF4'] if cmd_opts.use_openvino else ['INT8']}), "nncf_compress_weights_raito": OptionInfo(0, "Compress ratio", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": cmd_opts.use_openvino}), diff --git a/scripts/allegrovideo.py b/scripts/allegrovideo.py index 003211a60..fdcd52adf 100644 --- a/scripts/allegrovideo.py +++ b/scripts/allegrovideo.py @@ -61,14 +61,7 @@ class Script(scripts.Script): if shared.sd_model.__class__ != diffusers.AllegroPipeline: sd_models.unload_model_weights() t0 = time.time() - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=Allegro quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=Allegro quant={quant_args}') + quant_args = model_quant.create_config() transformer = diffusers.AllegroTransformer3DModel.from_pretrained( repo_id, subfolder="transformer", diff --git a/scripts/hunyuanvideo.py b/scripts/hunyuanvideo.py index d67fd7d31..dfd33e8ca 100644 --- a/scripts/hunyuanvideo.py +++ b/scripts/hunyuanvideo.py @@ -91,14 +91,7 @@ class Script(scripts.Script): if shared.sd_model.__class__ != diffusers.HunyuanVideoPipeline or model != loaded_model: sd_models.unload_model_weights() t0 = time.time() - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=HunyuanVideo quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=HunyuanVideo quant={quant_args}') + quant_args = model_quant.create_config() transformer = diffusers.HunyuanVideoTransformer3DModel.from_pretrained( pretrained_model_name_or_path='tencent/HunyuanVideo', subfolder="transformer", diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py index 5599ceed4..135256ef2 100644 --- a/scripts/ltxvideo.py +++ b/scripts/ltxvideo.py @@ -16,14 +16,7 @@ repos = { def load_quants(kwargs, repo_id): - quant_args = {} - quant_args = model_quant.create_bnb_config(quant_args) - if quant_args: - model_quant.load_bnb(f'Load model: type=LTXVideo quant={quant_args}') - if not quant_args: - quant_args = model_quant.create_ao_config(quant_args) - if quant_args: - model_quant.load_torchao(f'Load model: type=LTXVideo quant={quant_args}') + quant_args = model_quant.create_config() if not quant_args: return kwargs model_quant.load_bnb(f'Load model: type=LTX quant={quant_args}') @@ -119,9 +112,7 @@ class Script(scripts.Script): repo_id = model_custom if shared.sd_model.__class__ != cls: sd_models.unload_model_weights() - kwargs = {} - kwargs = model_quant.create_bnb_config(kwargs) - kwargs = model_quant.create_ao_config(kwargs) + kwargs = model_quant.create_config() diffusers.LTXVideoTransformer3DModel.forward = teacache_forward if os.path.isfile(repo_id): shared.sd_model = cls.from_single_file( diff --git a/scripts/mochivideo.py b/scripts/mochivideo.py index cbc9dad20..e2c193c24 100644 --- a/scripts/mochivideo.py +++ b/scripts/mochivideo.py @@ -42,9 +42,7 @@ class Script(scripts.Script): cls = diffusers.MochiPipeline if shared.sd_model.__class__ != cls: sd_models.unload_model_weights() - kwargs = {} - kwargs = model_quant.create_bnb_config(kwargs) - kwargs = model_quant.create_ao_config(kwargs) + kwargs = model_quant.create_config() shared.sd_model = cls.from_pretrained( repo_id, cache_dir = shared.opts.hfcache_dir, From 024635eb20e7e32c20da126755df7d9de365e8da Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 17 Mar 2025 10:33:57 -0400 Subject: [PATCH 017/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 +++--- wiki | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e237f7f02..3f62f251c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,7 +33,7 @@ against top-10 standard harmful content categories - add banned words/expressions check against prompt variations - **Other** - - **upscale**: new [asymmetric vae v2](Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method + - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method - **upscale**: new experimental support for `libvips` upscaling - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models note: previous method for quanto is still valid and is noted in settings as post-load quantization @@ -286,7 +286,7 @@ Just one week after latest release and what a week it was with over 50 commits! with detailed defaults for each model type also configurable - select between 150+ *OpenCLiP* supported models, 20+ built-in *VLMs*, *DeepDanbooru* - **VLM**: now that we can use VLMs freely, we've also added support for few more out-of-the-box - [Alibaba Qwen VL2](https://huggingface.co/Qwen/Qwen2-VL-2B), [Huggingface Smol VL2](HuggingFaceTB/SmolVLM-Instruct), [ToriiGate 0.4](Minthy/ToriiGate-v0.4-2B) + [Alibaba Qwen VL2](https://huggingface.co/Qwen/Qwen2-VL-2B), [Huggingface Smol VL2](https://huggingface.co/HuggingFaceTB/SmolVLM-Instruct), [ToriiGate 0.4](https://huggingface.co/Minthy/ToriiGate-v0.4-2B) - **Postprocess** - new sota remove background model: [BEN2](https://huggingface.co/PramaLLC/BEN2) select in *process -> remove background* or enable postprocessing for txt2img/img2img operations @@ -394,7 +394,7 @@ Two weeks since last release, time for update! - piecewise rectified flow as model acceleration - use `perflow` scheduler combined with one of the available pre-trained [models](https://huggingface.co/hansyan) - **Other**: - - **upscale**: new [asymmetric vae](Heasterian/AsymmetricAutoencoderKLUpscaler) upscaling method + - **upscale**: new [asymmetric vae](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler) upscaling method - **gallery**: add http fallback for slow/unreliable links - **splash**: add legacy mode indicator on splash screen - **network**: extract thumbnail from model metadata if present diff --git a/wiki b/wiki index 62636f56b..38ccb19d3 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 62636f56b6ee7b623eed50f150ad2dea77976948 +Subproject commit 38ccb19d3c3586d0f511a02d07e25f186c83f006 From d8c82eddd9282ba65074837e88a2461f967a35bf Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 17 Mar 2025 12:21:39 -0400 Subject: [PATCH 018/122] lint fixes Signed-off-by: Vladimir Mandic --- .ruff.toml | 3 +++ CHANGELOG.md | 34 ++++++++++++++++++------------ modules/sd_hijack.py | 6 +++--- modules/sd_hijack_optimizations.py | 2 +- modules/ui_extensions.py | 2 +- wiki | 2 +- 6 files changed, 30 insertions(+), 19 deletions(-) diff --git a/.ruff.toml b/.ruff.toml index 8e3d13e64..0fc9de8b3 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -70,6 +70,7 @@ select = [ ignore = [ "B006", # Do not use mutable data structures for argument defaults "B008", # Do not perform function call in argument defaults + "C420", # Unnecessary dict comprehension for iterable; use `dict.fromkeys` instead "C408", # Unnecessary `dict` call "I001", # Import block is un-sorted or un-formatted "E402", # Module level import not at top of file @@ -84,6 +85,8 @@ ignore = [ "RUF012", # Mutable class attributes "RUF013", # PEP 484 prohibits implicit `Optional` "RUF015", # Prefer `next(...)` over single element slice + "RUF046", # Value being cast to `int` is already an integer + "RUF051", # Prefer pop over del ] fixable = ["ALL"] unfixable = [] diff --git a/CHANGELOG.md b/CHANGELOG.md index 3f62f251c..ca168767f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,11 +1,17 @@ # Change Log for SD.Next -## Update for 2025-03-15 +## Update for 2025-03-17 ### TODO - Gemma3 requires `git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - Remote VAE encode for SD15 and Flux.1: +### Highlights for 2025-03-17 + +Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP models, improvements to remote VAE, additional docs/guides. + +### Details for 2025-03-17 + - **Models** - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer @@ -16,17 +22,21 @@ download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui +- **Wiki/Docs** + - new [Caption](https://github.com/vladmandic/sdnext/wiki/Caption) guide + - new [VAE](https://github.com/vladmandic/sdnext/wiki/VAE) guide + - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info + - updated [SD3](https://github.com/vladmandic/sdnext/wiki/SD3) guide - **Remote VAE** - add support for remote vae encode in addition to remote vae decode - used by *img2img, inpaint, hires, detailer* - remote vae encode is disabled by default, you can enable it in *settings -> variable auto-encoder* + - add remote vae info to metadata, thanks @iDeNoh + - remote vae use `scaling_factor` and `shift_factor` - **Caption/VLM** - [Google Gemma 3 4B](https://huggingface.co/google/gemma-3-4b-it) simply select from list of available models in caption tab - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* -- **Wiki/Docs** - - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - - updated SD3 content - [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) extension updates - add detection of prompt language and alphabet and filter based on those values - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) @@ -37,11 +47,16 @@ - **upscale**: new experimental support for `libvips` upscaling - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models note: previous method for quanto is still valid and is noted in settings as post-load quantization - - add remote vae info to metadata, thanks @iDeNoh - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements - - remote vae use `scaling_factor` and `shift_factor` - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion +- **IPEX** + - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* + - add xpu to profiler + - fix untyped_storage, torch.eye and torch.cuda.device ops + - fix torch 2.7 compatibility + - fix performance with balanced offload + - fix triton and torch.compile - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled @@ -56,13 +71,6 @@ - fix hires with latent upscale - fix legacy diffusion latent upscalers - fix upscaler selection in postprocessing -- **IPEX** - - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* - - add xpu to profiler - - fix untyped_storage, torch.eye and torch.cuda.device ops - - fix torch 2.7 compatibility - - fix performance with balanced offload - - fix triton and torch.compile ## Update for 2025-02-28 diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index d4456945a..57573a493 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -375,11 +375,11 @@ if devices.backend != "ipex": # disable_compile for AutoencoderKLOutput is the only change if torch.__version__.startswith("2.6"): from dataclasses import dataclass - from torch.compiler import disable as disable_compile - import diffusers.models.autoencoders.autoencoder_kl + from torch.compiler import disable as disable_compile # pylint: disable=ungrouped-imports + import diffusers.models.autoencoders.autoencoder_kl # pylint: disable=ungrouped-imports @dataclass @disable_compile class AutoencoderKLOutput(diffusers.utils.BaseOutput): - latent_dist: "DiagonalGaussianDistribution" # noqa: F821 + latent_dist: "DiagonalGaussianDistribution" # noqa: F821 diffusers.models.autoencoders.autoencoder_kl.AutoencoderKLOutput = AutoencoderKLOutput diff --git a/modules/sd_hijack_optimizations.py b/modules/sd_hijack_optimizations.py index ce6537925..f89465e36 100644 --- a/modules/sd_hijack_optimizations.py +++ b/modules/sd_hijack_optimizations.py @@ -315,7 +315,7 @@ def get_xformers_flash_attention_op(q, k, v): if 'Flash attention' not in shared.opts.xformers_options: return None try: - flash_attention_op = xformers.ops.MemoryEfficientAttentionFlashAttentionOp # pylint: disable=used-before-assignment + flash_attention_op = xformers.ops.MemoryEfficientAttentionFlashAttentionOp # pylint: disable=possibly-used-before-assignment, used-before-assignment fw, _bw = flash_attention_op if fw.supports(xformers.ops.fmha.Inputs(query=q, key=k, value=v, attn_bias=None)): return flash_attention_op diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index f2e36004b..2b75df0e5 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -98,7 +98,7 @@ def apply_changes(disable_list, update_list, disable_all): def check_updates(_id_task, disable_list, search_text, sort_column): if shared.cmd_opts.disable_extension_access: shared.log.error('Extension: apply changes disallowed because public access is enabled and insecure is not specified') - return + return create_html(search_text, sort_column) disabled = json.loads(disable_list) assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}" exts = [ext for ext in extensions.extensions if ext.remote is not None and ext.name not in disabled] diff --git a/wiki b/wiki index 38ccb19d3..37bc610f4 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 38ccb19d3c3586d0f511a02d07e25f186c83f006 +Subproject commit 37bc610f46affcdf417cfea61324acd505cad723 From 966eb3210a3eced98ffa11daac6063793b9e6a0c Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 17 Mar 2025 14:35:15 -0400 Subject: [PATCH 019/122] add api-grid Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 166 +++++++++++++++++++++++++++++++++++++++++++++ cli/api-txt2img.py | 3 + 2 files changed, 169 insertions(+) create mode 100755 cli/api-grid.py diff --git a/cli/api-grid.py b/cli/api-grid.py new file mode 100755 index 000000000..bd45ea066 --- /dev/null +++ b/cli/api-grid.py @@ -0,0 +1,166 @@ +#!/usr/bin/env python +from dataclasses import dataclass +import io +import os +import time +import math +import base64 +import logging +import argparse +import requests +import urllib3 +from PIL import Image, ImageDraw, ImageFont + + +@dataclass +class Options: # set default parameters here + prompt: str = '' + negative_prompt: str = '' + seed: int = -1 + steps: int = 20 + sampler_name: str = "Default" + width: int = 1024 + height: int = 1024 + save_images: bool = False + send_images: bool = True + + +@dataclass +class Server: # set server and save options here or use command line arguments + url: str = 'http://127.0.0.1:7860' + api: str = '/sdapi/v1/txt2img' + user: str = None + password: str = None + folder: str = '/tmp' + images: bool = False + grids: bool = False + labels: bool = False + + +logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(levelname)s: %(message)s') +log = logging.getLogger(__name__) +urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) +server = Server() +options = Options() + + +def post(): + req = requests.post(f'{server.url}{server.api}', + json=vars(options), + timeout=300, + verify=False, + auth=requests.auth.HTTPBasicAuth(server.user, server.password) if (server.user is not None) and (server.password is not None) else None) + return { 'error': req.status_code, 'reason': req.reason, 'url': req.url } if req.status_code != 200 else req.json() + + +def generate(ts: float, x: int, y: int): # pylint: disable=redefined-outer-name + t0 = time.time() + log.info(f'x={x} y={y} {options}') + data = post() + t1 = time.time() + images = [] + if 'images' in data: + for i in range(len(data['images'])): + b64 = data['images'][i].split(',',1)[0] + image = Image.open(io.BytesIO(base64.b64decode(b64))) + images.append(image) + info = data['info'] + fn = os.path.join(server.folder, f'{round(ts)}-{x}-{y}.jpg') if server.images else None + log.info(f'image: time={t1-t0:.2f} size={image.size} fn="{fn}" info="{info}"') + if fn is not None: + image.save(fn) + else: + log.warning(data) + return images + + +def merge(images: list[Image.Image], horizontal: bool, labels: list[str] = None): + rows = 1 if horizontal else len(images) + cols = math.ceil(len(images) / rows) + w = max([i.size[0] for i in images]) + h = max([i.size[1] for i in images]) + image = Image.new('RGB', size = (cols * w, rows * h), color = 'black') + font = ImageFont.truetype('DejaVuSansMono', 1024 // 32) + for i, img in enumerate(images): + x = i % cols * w + y = i // cols * h + img.thumbnail((w, h), Image.Resampling.LANCZOS) + image.paste(img, box=(x, y)) + if labels is not None and len(images) == len(labels): + ctx = ImageDraw.Draw(image) + ctx.text((x + 1, y + 1), labels[i], font = font, fill = (0, 0, 0)) + ctx.text((x, y), labels[i], font = font, fill = (255, 255, 255)) + # log.info({ 'grid': { 'images': len(images), 'rows': rows, 'cols': cols, 'cell': [w, h] } }) + return image + + +def grid(x_file: str, y_file: str): + def set_param(line): + param = line.split(':', maxsplit=1) + if param[0] == 'prompt': + options.prompt += f'{param[1]} ' # prompt is appended so its not overwritten + elif param[0] == 'lora': + options.prompt += f' ' # lora is appended to prompt + else: + setattr(options, param[0].strip(), param[1].strip()) + + log.info(server) + x = open(x_file, encoding='utf8').read().splitlines() if x_file is not None else [] + y = open(y_file, encoding='utf8').read().splitlines() if y_file is not None else [] + t0 = time.time() + log.info(f'grid: x={len(x)} y={len(y)} prefix={round(t0)}') + vertical = [] + Image.MAX_IMAGE_PIXELS = None + for j in range(max(1, len(y))): + horizontal = [] + labels = [] + for i in range(max(1, len(x))): + if len(x) > i: + set_param(x[i]) + if len(y) > i: + set_param(y[i]) + images = generate(t0, i, j) + if images is not None and len(images) > 0: + horizontal.extend(images) + labels.append(f'{x[i] if len(x) > i else ""}\n{y[j] if len(y) > j else ""}') + options.prompt = '' # reset prompt + if server.grids: + if len(horizontal) == 0: + log.warning(f'grid: empty row={j}') + continue + merged = merge(horizontal, horizontal=True, labels=labels if server.labels else None) + vertical.append(merged) + if server.grids: + if len(vertical) == 0: + log.warning('grid: empty grid') + return + merged = merge(vertical, horizontal=False) + fn = os.path.join(server.folder, f'{round(t0)}.jpg') + merged.save(fn) + log.info(f'grid: size={merged.size} fn="{fn}"') + t1 = time.time() + log.info(f'done: time={t1-t0:.2f}') + + +if __name__ == "__main__": + log.info(__file__) + parser = argparse.ArgumentParser(description = 'api-txt2img') + parser.add_argument('--x', required=False, default=None, help='file to use for x-axis values') + parser.add_argument('--y', required=False, default=None, help='file to use for y-axis values') + parser.add_argument('--folder', required=False, default='/tmp', help='folder to use for saving images') + parser.add_argument('--image', required=False, default=False, help='save individual images') + parser.add_argument('--grid', required=False, default=True, help='save image grids') + parser.add_argument('--labels', required=False, default=True, help='draw image labels') + parser.add_argument('--url', required=False, default='http://127.0.0.1:7860', help='server url') + parser.add_argument('--user', required=False, default=None, help='server user') + parser.add_argument('--password', required=False, default=None, help='server password') + args = parser.parse_args() + log.info(args) + server.folder = args.folder + server.images = args.image + server.grids = args.grid + server.labels = args.labels + server.url = args.url + server.user = args.user + server.password = args.password + grid(args.x, args.y) diff --git a/cli/api-txt2img.py b/cli/api-txt2img.py index 868b13eee..02bb876d6 100755 --- a/cli/api-txt2img.py +++ b/cli/api-txt2img.py @@ -54,10 +54,12 @@ def generate(args): # pylint: disable=redefined-outer-name options['hr_sampler_name'] = args.sampler data = post('/sdapi/v1/txt2img', options) t1 = time.time() + images = [] if 'images' in data: for i in range(len(data['images'])): b64 = data['images'][i].split(',',1)[0] image = Image.open(io.BytesIO(base64.b64decode(b64))) + images.append(image) info = data['info'] log.info(f'image received: size={image.size} time={t1-t0:.2f} info="{info}"') if args.output: @@ -65,6 +67,7 @@ def generate(args): # pylint: disable=redefined-outer-name log.info(f'image saved: size={image.size} filename={args.output}') else: log.warning(f'no images received: {data}') + return images if __name__ == "__main__": From d8044136b9e48ab38206c87340eea3f25fb82ca4 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 17 Mar 2025 22:29:03 -0400 Subject: [PATCH 020/122] prototype video tab Signed-off-by: Vladimir Mandic --- TODO.md | 27 +++---- javascript/amethyst-nightfall.css | 2 - javascript/black-gray.css | 5 -- javascript/black-orange.css | 2 - javascript/black-teal.css | 5 -- javascript/emerald-paradise.css | 2 - javascript/invoked.css | 2 - javascript/light-teal.css | 4 - javascript/midnight-barbie.css | 2 - javascript/orchid-dreams.css | 2 - javascript/sdnext.css | 17 ++-- javascript/timeless-beige.css | 2 - javascript/ui.js | 12 +++ modules/ui.py | 11 ++- modules/ui_common.py | 36 +++++---- modules/ui_sections.py | 21 ++--- modules/ui_txt2img.py | 4 +- modules/ui_video.py | 130 ++++++++++++++++++++++++++++++ modules/video_models/hunyuan.py | 13 +++ 19 files changed, 220 insertions(+), 79 deletions(-) create mode 100644 modules/ui_video.py create mode 100644 modules/video_models/hunyuan.py diff --git a/TODO.md b/TODO.md index 6acb369c5..0575e97c7 100644 --- a/TODO.md +++ b/TODO.md @@ -4,28 +4,23 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Current -- - ## Future Candidates - Redesign postprocessing - Flux NF4 loader: - IPAdapter negative: - Control API enhance scripts compatibility -- CogView4 ## Code TODO -- flux: loader for civitai nf4 models (fixme) -- hypertile: vae breaks when using non-standard sizes (fixme) -- install: enable ROCm for windows when available (fixme) -- lora make support quantized flux (fixme) -- lora: add other quantization types (fixme) -- model load: force-reloading entire model as loading transformers only leads to massive memory usage (fixme) -- model loader: implement model in-memory caching (fixme) -- modernui: monkey-patch for missing tabs.select event (fixme) -- processing: remove duplicate mask params (fixme) -- resize image: enable full VAE mode for resize-latent (fixme) -- sana: fails when quantized (fixme) -- support scripts via api (fixme) -- transformer from-single-file with quant (fixme) +- enable ROCm for windows when available +- resize image: enable full VAE mode for resize-latent +- infotext: handle using regex instead +- processing: remove duplicate mask params +- model loader: implement model in-memory caching +- hypertile: vae breaks when using non-standard sizes +- force-reloading entire model as loading transformers only leads to massive memory usage +- add other quantization types +- lora make support quantized flux +- control: support scripts via api +- modernui: monkey-patch for missing tabs.select event diff --git a/javascript/amethyst-nightfall.css b/javascript/amethyst-nightfall.css index ef582848d..929bd9e77 100644 --- a/javascript/amethyst-nightfall.css +++ b/javascript/amethyst-nightfall.css @@ -88,8 +88,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/black-gray.css b/javascript/black-gray.css index c262a3bf4..784626d60 100644 --- a/javascript/black-gray.css +++ b/javascript/black-gray.css @@ -103,11 +103,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } - -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt { background-color: var(--background-color); box-shadow: none !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea, #control_prompt > label > textarea, #control_neg_prompt > label > textarea { font-size: 1.0em; line-height: 1.4em; } -#txt2img_styles, #img2img_styles, #control_styles { padding: 0; } -#txt2img_styles_refresh, #img2img_styles_refresh, #control_styles_refresh { padding: 0; margin-top: 1em; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: var(--primary-950); padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/black-orange.css b/javascript/black-orange.css index 54b98b1df..467c53d1d 100644 --- a/javascript/black-orange.css +++ b/javascript/black-orange.css @@ -105,8 +105,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/black-teal.css b/javascript/black-teal.css index 0a6db4fa2..851c03953 100644 --- a/javascript/black-teal.css +++ b/javascript/black-teal.css @@ -142,11 +142,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } - -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt { background-color: var(--background-color); box-shadow: none !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea, #control_prompt > label > textarea, #control_neg_prompt > label > textarea { font-size: 1.0em; line-height: 1.4em; } -#txt2img_styles, #img2img_styles, #control_styles { padding: 0; margin-top: 2px; } -#txt2img_styles_refresh, #img2img_styles_refresh, #control_styles_refresh { padding: 0; margin-top: 1em; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: var(--neutral-950); padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/emerald-paradise.css b/javascript/emerald-paradise.css index f951356cc..411b37774 100644 --- a/javascript/emerald-paradise.css +++ b/javascript/emerald-paradise.css @@ -101,8 +101,6 @@ button.selected {background: var(--button-primary-background-fill);} /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-top: -2px; height: 2.4em; } diff --git a/javascript/invoked.css b/javascript/invoked.css index 72e78d31a..a5954adbf 100644 --- a/javascript/invoked.css +++ b/javascript/invoked.css @@ -101,8 +101,6 @@ button.selected {background: var(--button-primary-background-fill);} /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-top: -2px; height: 2.4em; } diff --git a/javascript/light-teal.css b/javascript/light-teal.css index 7dc9e5950..df8a3ab51 100644 --- a/javascript/light-teal.css +++ b/javascript/light-teal.css @@ -101,10 +101,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea, #control_prompt > label > textarea, #control_neg_prompt > label > textarea { font-size: 1.0em; line-height: 1.4em; } -#txt2img_styles, #img2img_styles, #control_styles { padding: 0; } -#txt2img_styles_refresh, #img2img_styles_refresh, #control_styles_refresh { padding: 0; margin-top: 1em; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/midnight-barbie.css b/javascript/midnight-barbie.css index ea78e5cab..9facd698d 100644 --- a/javascript/midnight-barbie.css +++ b/javascript/midnight-barbie.css @@ -94,8 +94,6 @@ svg.feather.feather-image, .feather .feather-image { display: none } /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } diff --git a/javascript/orchid-dreams.css b/javascript/orchid-dreams.css index 4b121c761..915823bb3 100644 --- a/javascript/orchid-dreams.css +++ b/javascript/orchid-dreams.css @@ -101,8 +101,6 @@ button.selected {background: var(--button-primary-background-fill);} /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-top: -2px; height: 2.4em; } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 3a564679e..d9aa365e7 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -99,20 +99,20 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_gallery, #img2img_gallery { height: 50vh; } #control-result { background: var(--button-secondary-background-fill); padding: 0.2em; } #control-inputs { margin-top: 1em; } -#txt2img_prompt_container, #img2img_prompt_container, #control_prompt_container { margin-right: var(--layout-gap) } +#txt2img_prompt_container, #img2img_prompt_container, #control_prompt_container, #video_prompt_container { margin-right: var(--layout-gap) } #txt2img_footer, #img2img_footer, #control_footer { height: fit-content; display: none; } #txt2img_generate_box, #img2img_generate_box { gap: 0.5em; flex-wrap: unset; min-width: unset; width: 66.6%; } #control_generate_box { gap: 0.5em; flex-wrap: unset; min-width: unset; width: 100%; } #control_generate_box button:nth-child(1) { flex-grow: 2; } #control_generate_box button:nth-child(2) { flex-grow: 1; } -#txt2img_actions_column, #img2img_actions_column, #control_actions_column { gap: 0.3em; height: fit-content; } -#txt2img_generate_box>button, #img2img_generate_box>button, #control_generate_box>button, #txt2img_enqueue, #img2img_enqueue, #txt2img_enqueue>button, #img2img_enqueue>button { min-height: 44px !important; max-height: 44px !important; line-height: 1em; white-space: break-spaces; min-width: unset; } +#txt2img_actions_column, #img2img_actions_column, #control_actions_column, #video_actions_column { gap: 0.3em; height: fit-content; } +#txt2img_generate_box>button, #img2img_generate_box>button, #control_generate_box>button, #video_generate_box>button, #txt2img_enqueue, #img2img_enqueue, #txt2img_enqueue>button, #img2img_enqueue>button { min-height: 44px !important; max-height: 44px !important; line-height: 1em; white-space: break-spaces; min-width: unset; } #txt2img_enqueue_wrapper, #img2img_enqueue_wrapper, #control_enqueue_wrapper { min-width: unset !important; width: 31%; } #txt2img_generate_line2, #img2img_generate_line2, #txt2img_tools, #img2img_tools, #control_generate_line2, #control_tools { display: flex; } #txt2img_generate_line2>button, #img2img_generate_line2>button, #extras_generate_box>button, #control_generate_line2>button, #txt2img_tools>button, #img2img_tools>button, #control_tools>button { height: 2em; line-height: 0; font-size: var(--text-md); min-width: unset; display: block !important; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt { display: contents; } -#txt2img_actions_column, #img2img_actions_column, #control_actions { flex-flow: wrap; justify-content: space-between; } +#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt, #video_prompt, #video_neg_prompt { display: contents; } +#txt2img_actions_column, #img2img_actions_column, #control_actions, #video_actions { flex-flow: wrap; justify-content: space-between; } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } .interrogate:hover { background: var(--button-primary-background-fill-hover) !important; } @@ -140,6 +140,11 @@ div#extras_scale_to_tab div.form { flex-direction: row; } #txt2img_advanced_options, #img2img_advanced_options, #control_advanced_options { min-width: 100%; } #txt2img_advanced_options .gradio-checkbox, #img2img_advanced_options .gradio-checkbox, #control_advanced_options .gradio-checkbox { min-width: unset !important; max-width: fit-content; } +#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt, #video_prompt, #video_neg_prompt { background-color: var(--background-color); box-shadow: none !important; } +#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea, #control_prompt > label > textarea, #control_neg_prompt > label > textarea, #video_prompt > label > textarea, #video_neg_prompt > label > textarea { font-size: 1.0em; line-height: 1.4em; } +#txt2img_styles, #img2img_styles, #control_styles, #video_styles { padding: 0; margin-top: 2px; } +#txt2img_styles_refresh, #img2img_styles_refresh, #control_styles_refresh, #video_styles_refresh { padding: 0; margin-top: 1em; } + /* settings */ #si-sparkline-memo, #si-sparkline-load { background-color: #111; } #quicksettings { width: fit-content; } @@ -388,7 +393,7 @@ div:has(>#tab-gallery-folders) { flex-grow: 0 !important; background-color: var( #txt2img_results, #extras_results, #txt2im g_footer p { text-wrap: wrap; max-width: 100% !important; } /* maintain side by side split on larger mobile displays for from text */ } #scripts_alwayson_txt2img div, #scripts_alwayson_img2img div { max-width: 100%; } - #txt2img_prompt_container, #img2img_prompt_container, #control_prompt_container { resize: vertical !important; } + #txt2img_prompt_container, #img2img_prompt_container, #control_prompt_container, #video_prompt_container { resize: vertical !important; } #txt2img_generate_box, #txt2img_enqueue_wrapper { min-width: 100% !important;} /* make generate and enqueue buttons take up the entire width of their rows. */ #img2img_toprow>div.gradio-column { flex-grow: 1 !important;} /*make interrogate buttons take up appropriate space. */ #img2img_actions_column { display: flex; min-width: fit-content !important; flex-direction: row;justify-content: space-evenly; align-items: center;} diff --git a/javascript/timeless-beige.css b/javascript/timeless-beige.css index 4b0f7d9e4..a8a9c1536 100644 --- a/javascript/timeless-beige.css +++ b/javascript/timeless-beige.css @@ -101,8 +101,6 @@ button.selected {background: var(--button-primary-background-fill);} /* gradio elements overrides */ #div.gradio-container { overflow-x: hidden; } #img2img_label_copy_to_img2img { font-weight: normal; } -#txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: var(--background-color); box-shadow: 4px 4px 4px 0px #333333 !important; } -#txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.1rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } #interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-top: -2px; height: 2.4em; } diff --git a/javascript/ui.js b/javascript/ui.js index 9df1a26e4..badec0f01 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -240,6 +240,18 @@ function submit_control(...args) { return res; } +function submit_video(...args) { + log('submitVideo'); + clearGallery('video'); + const id = randomId(); + requestProgress(id, null, gradioApp().getElementById('video_gallery')); + const res = create_submit_args(args); + res[0] = id; + res[1] = window.submit_state; + window.submit_state = ''; + return res; +} + function submit_postprocessing(...args) { log('SubmitExtras'); clearGallery('extras'); diff --git a/modules/ui.py b/modules/ui.py index 6d52187c8..371a26945 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -146,6 +146,12 @@ def create_ui(startup_timer = None): ui_control.create_ui() timer.startup.record("ui-control") + with gr.Blocks(analytics_enabled=False) as video_interface: + if shared.native and shared.cmd_opts.experimental: + from modules import ui_video + ui_video.create_ui() + timer.startup.record("ui-video") + with gr.Blocks(analytics_enabled=False) as extras_interface: from modules import ui_postprocessing ui_postprocessing.create_ui() @@ -398,7 +404,10 @@ def create_ui(startup_timer = None): interfaces = [] interfaces += [(txt2img_interface, "Text", "txt2img")] interfaces += [(img2img_interface, "Image", "img2img")] - interfaces += [(control_interface, "Control", "control")] if control_interface is not None else [] + if control_interface is not None: + interfaces += [(control_interface, "Control", "control")] + if video_interface is not None: + interfaces += [(video_interface, "Video", "video")] interfaces += [(extras_interface, "Process", "process")] interfaces += [(caption_interface, "Caption", "caption")] interfaces += [(gallery_interface, "Gallery", "gallery")] diff --git a/modules/ui_common.py b/modules/ui_common.py index a8167e457..f5260d10c 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -237,8 +237,8 @@ def interrogate_booru(image): # legacy function return gr.update() if prompt is None else prompt -def create_output_panel(tabname, preview=True, prompt=None, height=None): - with gr.Column(variant='panel', elem_id=f"{tabname}_results"): +def create_output_panel(tabname, preview=True, prompt=None, height=None, transfer=True, scale=1): + with gr.Column(variant='panel', elem_id=f"{tabname}_results", scale=scale): with gr.Group(elem_id=f"{tabname}_gallery_container"): if tabname == "txt2img": gr.HTML(value="", elem_id="main_info", visible=False, elem_classes=["main-info"]) @@ -270,10 +270,13 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None): clip_files.click(fn=None, _js='clip_gallery_urls', inputs=[result_gallery], outputs=[]) save = gr.Button('Save', elem_id=f'save_{tabname}') delete = gr.Button('Delete', elem_id=f'delete_{tabname}') - if not shared.native: - buttons = generation_parameters_copypaste.create_buttons(["img2img", "inpaint", "extras"]) + if transfer: + if not shared.native: + buttons = generation_parameters_copypaste.create_buttons(["img2img", "inpaint", "extras"]) + else: + buttons = generation_parameters_copypaste.create_buttons(["txt2img", "img2img", "control", "extras", "caption"]) else: - buttons = generation_parameters_copypaste.create_buttons(["txt2img", "img2img", "control", "extras", "caption"]) + buttons = None download_files = gr.File(None, file_count="multiple", interactive=False, show_label=False, visible=False, elem_id=f'download_files_{tabname}') with gr.Group(): @@ -309,17 +312,18 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None): else: paste_field_names = [] debug(f'Paste field: tab={tabname} fields={paste_field_names}') - for paste_tabname, paste_button in buttons.items(): - debug(f'Create output panel: source={tabname} target={paste_tabname} button={paste_button}') - bindings = generation_parameters_copypaste.ParamBinding( - paste_button=paste_button, - tabname=paste_tabname, - source_tabname=tabname, - source_image_component=result_gallery, - paste_field_names=paste_field_names, - source_text_component=prompt or generation_info - ) - generation_parameters_copypaste.register_paste_params_button(bindings) + if buttons is not None: + for paste_tabname, paste_button in buttons.items(): + debug(f'Create output panel: source={tabname} target={paste_tabname} button={paste_button}') + bindings = generation_parameters_copypaste.ParamBinding( + paste_button=paste_button, + tabname=paste_tabname, + source_tabname=tabname, + source_image_component=result_gallery, + paste_field_names=paste_field_names, + source_text_component=prompt or generation_info + ) + generation_parameters_copypaste.register_paste_params_button(bindings) return result_gallery, generation_info, html_info, html_info_formatted, html_log diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 0d0e0eb8d..20b1cfc50 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -4,7 +4,7 @@ from modules.ui_components import ToolButton from modules.interrogate import interrogate -def create_toprow(is_img2img: bool = False, id_part: str = None): +def create_toprow(is_img2img: bool = False, id_part: str = None, negative_visible: bool = True, reprocess_visible: bool = True): def apply_styles(prompt, prompt_neg, styles): prompt = shared.prompt_styles.apply_styles_to_prompt(prompt, styles, wildcards=not shared.opts.extra_networks_apply_unparsed) prompt_neg = shared.prompt_styles.apply_negative_styles_to_prompt(prompt_neg, styles, wildcards=not shared.opts.extra_networks_apply_unparsed) @@ -21,19 +21,20 @@ def create_toprow(is_img2img: bool = False, id_part: str = None): with gr.Row(): with gr.Column(scale=80): with gr.Row(elem_id=f"{id_part}_prompt_row"): - prompt = gr.Textbox(elem_id=f"{id_part}_prompt", label="Prompt", show_label=False, lines=3, placeholder="Prompt", elem_classes=["prompt"]) + prompt = gr.Textbox(elem_id=f"{id_part}_prompt", label="Prompt", show_label=False, lines=3 if negative_visible else 5, placeholder="Prompt", elem_classes=["prompt"]) with gr.Row(): with gr.Column(scale=80): with gr.Row(elem_id=f"{id_part}_negative_row"): - negative_prompt = gr.Textbox(elem_id=f"{id_part}_neg_prompt", label="Negative prompt", show_label=False, lines=3, placeholder="Negative prompt", elem_classes=["prompt"]) + negative_prompt = gr.Textbox(elem_id=f"{id_part}_neg_prompt", label="Negative prompt", show_label=False, lines=3, placeholder="Negative prompt", elem_classes=["prompt"], visible=negative_visible) with gr.Column(scale=1, elem_id=f"{id_part}_actions_column"): with gr.Row(elem_id=f"{id_part}_generate_box"): reprocess = [] submit = gr.Button('Generate', elem_id=f"{id_part}_generate", variant='primary') - reprocess.append(gr.Button('Reprocess', elem_id=f"{id_part}_reprocess", variant='primary', visible=True)) - reprocess.append(gr.Button('Reprocess decode', elem_id=f"{id_part}_reprocess_decode", variant='primary', visible=False)) - reprocess.append(gr.Button('Reprocess refine', elem_id=f"{id_part}_reprocess_refine", variant='primary', visible=False)) - reprocess.append(gr.Button('Reprocess face', elem_id=f"{id_part}_reprocess_detail", variant='primary', visible=False)) + if reprocess_visible: + reprocess.append(gr.Button('Reprocess', elem_id=f"{id_part}_reprocess", variant='primary', visible=True)) + reprocess.append(gr.Button('Reprocess decode', elem_id=f"{id_part}_reprocess_decode", variant='primary', visible=False)) + reprocess.append(gr.Button('Reprocess refine', elem_id=f"{id_part}_reprocess_refine", variant='primary', visible=False)) + reprocess.append(gr.Button('Reprocess face', elem_id=f"{id_part}_reprocess_detail", variant='primary', visible=False)) with gr.Row(elem_id=f"{id_part}_generate_line2"): interrupt = gr.Button('Stop', elem_id=f"{id_part}_interrupt") interrupt.click(fn=lambda: shared.state.interrupt(), _js="requestInterrupt", inputs=[], outputs=[]) @@ -79,9 +80,9 @@ def ar_change(ar, width, height): return gr.update(), gr.update() -def create_resolution_inputs(tab): - width = gr.Slider(minimum=64, maximum=4096, step=8, label="Width", value=1024, elem_id=f"{tab}_width") - height = gr.Slider(minimum=64, maximum=4096, step=8, label="Height", value=1024, elem_id=f"{tab}_height") +def create_resolution_inputs(tab, default_width=1024, default_height=1024): + width = gr.Slider(minimum=64, maximum=4096, step=8, label="Width", value=default_width, elem_id=f"{tab}_width") + height = gr.Slider(minimum=64, maximum=4096, step=8, label="Height", value=default_height, elem_id=f"{tab}_height") ar_list = ['AR'] + [x.strip() for x in shared.opts.aspect_ratios.split(',') if x.strip() != ''] ar_dropdown = gr.Dropdown(show_label=False, interactive=True, choices=ar_list, value=ar_list[0], elem_id=f"{tab}_ar", elem_classes=["ar-dropdown"]) for c in [ar_dropdown, width, height]: diff --git a/modules/ui_txt2img.py b/modules/ui_txt2img.py index 30f9ab5ad..42f57fc26 100644 --- a/modules/ui_txt2img.py +++ b/modules/ui_txt2img.py @@ -1,6 +1,6 @@ import gradio as gr from modules.call_queue import wrap_gradio_gpu_call, wrap_queued_call -from modules import timer, shared, ui_common, ui_sections, generation_parameters_copypaste, processing, processing_vae, devices +from modules import timer, shared, ui_common, ui_sections, generation_parameters_copypaste, processing, processing_vae, devices, images from modules.ui_components import ToolButton # pylint: disable=unused-import @@ -23,7 +23,7 @@ def create_ui(): txt2img_prompt, txt2img_prompt_styles, txt2img_negative_prompt, txt2img_submit, txt2img_reprocess, txt2img_paste, txt2img_extra_networks_button, txt2img_token_counter, txt2img_token_button, txt2img_negative_token_counter, txt2img_negative_token_button = ui_sections.create_toprow(is_img2img=False, id_part="txt2img") txt_prompt_img = gr.File(label="", elem_id="txt2img_prompt_image", file_count="single", type="binary", visible=False) - txt_prompt_img.change(fn=modules.images.image_data, inputs=[txt_prompt_img], outputs=[txt2img_prompt, txt_prompt_img]) + txt_prompt_img.change(fn=images.image_data, inputs=[txt_prompt_img], outputs=[txt2img_prompt, txt_prompt_img]) with gr.Row(variant='compact', elem_id="txt2img_extra_networks", elem_classes=["extra_networks_root"], visible=False) as extra_networks_ui: from modules import ui_extra_networks diff --git a/modules/ui_video.py b/modules/ui_video.py new file mode 100644 index 000000000..130aebd0d --- /dev/null +++ b/modules/ui_video.py @@ -0,0 +1,130 @@ +# TODO hunyuanvideo: seed, scheduler, scheduler_shift, guidance_scale=1.0, true_cfg_scale=6.0, num_inference_steps=30, prompt_template, vae, offloading +# TODO modernui video tab + +from dataclasses import dataclass +import gradio as gr +from modules import shared, images, ui_common, ui_sections, sd_models, call_queue, generation_parameters_copypaste +from modules.video_models import hunyuan + + +@dataclass +class Model(): + name: str + repo: str + dit: str + + +MODELS = { + 'None': [], + 'Hunyuan Video': [ + Model('None', None, None), + Model('Hunyuan Video T2V', 'hunyuanvideo-community/HunyuanVideo', None), + Model('Hunyuan Video I2V', 'hunyuanvideo-community/HunyuanVideo', 'hunyuanvideo-community/HunyuanVideo-I2V'), # https://github.com/huggingface/diffusers/pull/10983 + Model('SkyReels Hunyuan T2V', 'hunyuanvideo-community/HunyuanVideo', 'Skywork/SkyReels-V1-Hunyuan-T2V'), # https://github.com/huggingface/diffusers/pull/10837 + Model('SkyReels Hunyuan I2V', 'hunyuanvideo-community/HunyuanVideo', 'Skywork/SkyReels-V1-Hunyuan-I2V'), + Model('Fast Hunyuan T2V', 'hunyuanvideo-community/HunyuanVideo', 'hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt'), # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 + ] +} + + +def engine_change(engine): + models = [model.name for model in MODELS.get(engine, [])] + return gr.update(choices=models, value=models[0] if len(models) > 0 else None) + + +def model_change(engine, model): + models = [model.name for model in MODELS.get(engine, [])] + selected = [m for m in MODELS[engine] if m.name == model][0] if len(models) > 0 else None + if selected: + if 'None' in selected.name: + sd_models.unload_model_weights() + msg = 'Video model unloaded' + elif 'Hunyuan' in selected.name: + msg = hunyuan.load(selected) + elif model != 'None': + msg = f'Video model not found: engine={engine} model={model}' + shared.log.error(msg) + else: + sd_models.unload_model_weights() + msg = 'Video model unloaded' + return [msg, gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False)] + + +def run_video(*args): + engine, model = args[2], args[3] + models = [model.name for model in MODELS.get(engine, [])] + selected = [m for m in MODELS[engine] if m.name == model][0] if len(models) > 0 else None + if selected and 'Hunyuan' in selected.name: + return hunyuan.generate(*args) + shared.log.error(f'Video model not found: args={args}') + return [], '', '', f'Video model not found: engine={engine} model={model}' + + +def create_ui(): + shared.log.debug('UI initialize: txt2img') + with gr.Blocks(analytics_enabled=False) as _video_interface: + prompt, styles, _negative, generate, _reprocess, paste, _networks, _token_counter, _token_button, _token_counter_negative, _token_button_negative = ui_sections.create_toprow(is_img2img=False, id_part="video", negative_visible=False, reprocess_visible=False) + prompt_image = gr.File(label="", elem_id="video_prompt_image", file_count="single", type="binary", visible=False) + prompt_image.change(fn=images.image_data, inputs=[prompt_image], outputs=[prompt, prompt_image]) + + with gr.Row(elem_id="video_interface", equal_height=False): + with gr.Column(variant='compact', elem_id="video_settings", scale=1): + + with gr.Row(): + engine = gr.Dropdown(label='Engine', choices=list(MODELS), value='None') + model = gr.Dropdown(label='Model', choices=[''], value=None) + with gr.Row(): + width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) + with gr.Row(): + frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15) + with gr.Row(): + with gr.Group(visible=False) as image_group: + gr.HTML("
  Init image") + image = gr.Image(elem_id="video_image", show_label=False, source="upload", interactive=True, type="pil", tool="select", image_mode="RGB", height=512) + with gr.Row(): + save_frames = gr.Checkbox(label='Save image frames', value=False) + with gr.Row(): + cc, duration, loop, pad, interpolate = ui_sections.create_video_inputs(tab='video') + override_settings = ui_common.create_override_inputs('video') + + # output panel with gallery + gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) + + # handle engine and model change + engine.change(fn=engine_change, inputs=[engine], outputs=[model]) + model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group]) + # handle restore fields + paste_fields = [ + (prompt, "Prompt"), + # main + (width, "Size-1"), + (height, "Size-2"), + (frames, "Frames"), + ] + generation_parameters_copypaste.add_paste_fields("txt2img", None, paste_fields, override_settings) + bindings = generation_parameters_copypaste.ParamBinding(paste_button=paste, tabname="video", source_text_component=prompt, source_image_component=None) + generation_parameters_copypaste.register_paste_params_button(bindings) + # hidden fields + task_id = gr.Textbox(visible=False, value='') + ui_state = gr.Textbox(visible=False, value='') + # generate args + video_args = [ + task_id, ui_state, + engine, model, + prompt, styles, + width, height, + frames, + image, + save_frames, + cc, duration, loop, pad, interpolate, + ] + # generate function + video_dict = dict( + fn=call_queue.wrap_gradio_gpu_call(run_video, extra_outputs=[None, '', ''], name='Video'), + _js="submit_video", + inputs=video_args, + outputs=[gallery, gen_info, html_info, html_log], + show_progress=False, + ) + prompt.submit(**video_dict) + generate.click(**video_dict) diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py new file mode 100644 index 000000000..6fed27cb3 --- /dev/null +++ b/modules/video_models/hunyuan.py @@ -0,0 +1,13 @@ +from modules import shared + + +def load(selected): + msg = f'Video load: model="{selected.name}" repo="{selected.repo}" dit="{selected.dit}"' + shared.log.info(msg) + return msg + + +def generate(*args, **kwargs): + # TODO hunyuanvideo: check if loaded + shared.log.debug(f'Video generate: args={args} kwargs={kwargs}') + return [], '', '', 'TBD' From c5c3be04fb757e0bc17e62869c301a6a99ab2351 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 18 Mar 2025 18:40:08 -0400 Subject: [PATCH 021/122] video tab first prototype Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 26 ++- javascript/black-teal-reimagined.css | 2 + javascript/extraNetworks.js | 6 +- javascript/progressBar.js | 4 +- javascript/sdnext.css | 4 +- javascript/ui.js | 5 + modules/lora/network.py | 5 + modules/modeldata.py | 2 +- modules/modelloader.py | 3 + modules/processing.py | 2 +- modules/processing_class.py | 12 +- modules/processing_info.py | 2 +- modules/ui.py | 2 +- modules/ui_sections.py | 10 +- modules/ui_video.py | 99 ++++++----- modules/video_models/hunyuan.py | 257 ++++++++++++++++++++++++++- 16 files changed, 368 insertions(+), 73 deletions(-) diff --git a/cli/api-grid.py b/cli/api-grid.py index bd45ea066..21eb685a4 100755 --- a/cli/api-grid.py +++ b/cli/api-grid.py @@ -45,12 +45,15 @@ options = Options() def post(): - req = requests.post(f'{server.url}{server.api}', - json=vars(options), - timeout=300, - verify=False, - auth=requests.auth.HTTPBasicAuth(server.user, server.password) if (server.user is not None) and (server.password is not None) else None) - return { 'error': req.status_code, 'reason': req.reason, 'url': req.url } if req.status_code != 200 else req.json() + try: + req = requests.post(f'{server.url}{server.api}', + json=vars(options), + timeout=300, + verify=False, + auth=requests.auth.HTTPBasicAuth(server.user, server.password) if (server.user is not None) and (server.password is not None) else None) + return { 'error': req.status_code, 'reason': req.reason, 'url': req.url } if req.status_code != 200 else req.json() + except Exception as e: + return { 'error': 0, 'reason': str(e), 'url': server.url } def generate(ts: float, x: int, y: int): # pylint: disable=redefined-outer-name @@ -105,8 +108,15 @@ def grid(x_file: str, y_file: str): setattr(options, param[0].strip(), param[1].strip()) log.info(server) - x = open(x_file, encoding='utf8').read().splitlines() if x_file is not None else [] - y = open(y_file, encoding='utf8').read().splitlines() if y_file is not None else [] + os.makedirs(server.folder, exist_ok=True) + try: + x = open(x_file, encoding='utf8').read().splitlines() if x_file is not None else [] + y = open(y_file, encoding='utf8').read().splitlines() if y_file is not None else [] + except Exception as e: + log.error(f'read file: x={x_file} y={y_file} {e}') + return + x = [line for line in x if ':' in line] + y = [line for line in y if ':' in line] t0 = time.time() log.info(f'grid: x={len(x)} y={len(y)} prefix={round(t0)}') vertical = [] diff --git a/javascript/black-teal-reimagined.css b/javascript/black-teal-reimagined.css index be6176ac4..ba8957c3e 100644 --- a/javascript/black-teal-reimagined.css +++ b/javascript/black-teal-reimagined.css @@ -816,6 +816,8 @@ svg.feather.feather-image, #txt2img_extra_search, #img2img_description, #img2img_extra_search, +#video_description, +#video_extra_search, #control_description, #control_extra_search { margin-top: 50px; diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index a9f055a32..b9f045a6b 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -9,6 +9,7 @@ const getENActiveTab = () => { if (gradioApp().getElementById('tab_txt2img').style.display === 'block') tabName = 'txt2img'; else if (gradioApp().getElementById('tab_img2img').style.display === 'block') tabName = 'img2img'; else if (gradioApp().getElementById('tab_control').style.display === 'block') tabName = 'control'; + else if (gradioApp().getElementById('tab_video').style.display === 'block') tabName = 'video'; // log('getENActiveTab', tabName); return tabName; }; @@ -491,7 +492,7 @@ function setupExtraNetworksForTab(tabname) { } async function showNetworks() { - for (const tabname of ['txt2img', 'img2img', 'control']) { + for (const tabname of ['txt2img', 'img2img', 'control', 'video']) { if (window.opts.extra_networks_show) gradioApp().getElementById(`${tabname}_extra_networks_btn`).click(); } log('showNetworks'); @@ -501,6 +502,7 @@ async function setupExtraNetworks() { setupExtraNetworksForTab('txt2img'); setupExtraNetworksForTab('img2img'); setupExtraNetworksForTab('control'); + setupExtraNetworksForTab('video'); function registerPrompt(tabname, id) { const textarea = gradioApp().querySelector(`#${id} > label > textarea`); @@ -515,6 +517,8 @@ async function setupExtraNetworks() { registerPrompt('img2img', 'img2img_neg_prompt'); registerPrompt('control', 'control_prompt'); registerPrompt('control', 'control_neg_prompt'); + registerPrompt('video', 'video_prompt'); + registerPrompt('video', 'video_neg_prompt'); log('initNetworks', window.opts.extra_networks_card_size); document.documentElement.style.setProperty('--card-size', `${window.opts.extra_networks_card_size}px`); } diff --git a/javascript/progressBar.js b/javascript/progressBar.js index 0bac99d6f..cfc9a9039 100644 --- a/javascript/progressBar.js +++ b/javascript/progressBar.js @@ -14,12 +14,14 @@ function checkPaused(state) { lastState.paused = state ? !state : !lastState.paused; const t_el = document.getElementById('txt2img_pause'); const i_el = document.getElementById('img2img_pause'); + const v_el = document.getElementById('video_pause'); if (t_el) t_el.innerText = lastState.paused ? 'Resume' : 'Pause'; if (i_el) i_el.innerText = lastState.paused ? 'Resume' : 'Pause'; + if (v_el) v_el.innerText = lastState.paused ? 'Resume' : 'Pause'; } function setProgress(res) { - const elements = ['txt2img_generate', 'img2img_generate', 'extras_generate', 'control_generate']; + const elements = ['txt2img_generate', 'img2img_generate', 'extras_generate', 'control_generate', 'video_generate']; const progress = res?.progress || 0; const job = res?.job || ''; let perc = ''; diff --git a/javascript/sdnext.css b/javascript/sdnext.css index d9aa365e7..ec5e4fa3d 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -113,6 +113,7 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- min-width: unset; display: block !important; } #txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt, #video_prompt, #video_neg_prompt { display: contents; } #txt2img_actions_column, #img2img_actions_column, #control_actions, #video_actions { flex-flow: wrap; justify-content: space-between; } +#txt2img_seed, #img2img_seed, #control_seed, #video_seed { min-width: 90px !important } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } .interrogate:hover { background: var(--button-primary-background-fill-hover) !important; } @@ -399,7 +400,8 @@ div:has(>#tab-gallery-folders) { flex-grow: 0 !important; background-color: var( #img2img_actions_column { display: flex; min-width: fit-content !important; flex-direction: row;justify-content: space-evenly; align-items: center;} #txt2img_generate_box, #img2img_generate_box, #txt2img_enqueue_wrapper,#img2img_enqueue_wrapper {display: flex;flex-direction: column;height: 4em !important;align-items: stretch;justify-content: space-evenly;} #img2img_interface, #img2img_results, #img2img_footer p { text-wrap: wrap; min-width: 100% !important; max-width: 100% !important;} /* maintain single column for from image operations on larger mobile devices */ - #txt2img_sampler, #txt2img_batch, #txt2img_seed_group, #txt2img_advanced, #txt2img_second_pass, #img2img_sampling_group, #img2img_resize_group, #img2img_batch_group, #img2img_seed_group, #img2img_denoise_group, #img2img_advanced_group { width: 100% !important; } /* fix from text/image UI elements to prevent them from moving around within the UI */ + #txt2img_sampler, #txt2img_batch, #txt2img_seed_group, #txt2img_advanced, #txt2img_second_pass, #img2img_sampling_group, #img2img_resize_group, #img2img_batch_group, #img2img_seed_group, #img2img_denoise_group, #img2img_advanced_group { width: 100% !important; } /* fix from text/image UI + elements to prevent them from moving around within the UI */ #img2img_resize_group .gradio-radio>div { display: flex; flex-direction: column; width: unset !important; } #inpaint_controls div { display:flex;flex-direction: row;} #inpaint_controls .gradio-radio>div { display: flex; flex-direction: column !important; } diff --git a/javascript/ui.js b/javascript/ui.js index badec0f01..9a5f2b794 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -155,6 +155,11 @@ function switch_to_control(...args) { return Array.from(arguments); } +function switch_to_video(...args) { + switchToTab('Video'); + return Array.from(arguments); +} + function switch_to_caption(...args) { switchToTab('Caption'); return Array.from(arguments); diff --git a/modules/lora/network.py b/modules/lora/network.py index 97feb76f1..c4768d9ad 100644 --- a/modules/lora/network.py +++ b/modules/lora/network.py @@ -17,6 +17,7 @@ class SdVersion(enum.Enum): SDXL = 4 SC = 5 F1 = 6 + HV = 7 class NetworkOnDisk: @@ -56,6 +57,8 @@ class NetworkOnDisk: return 'sd3' if base.startswith("flux"): return 'f1' + if base.startswith("hunyuan_video"): + return 'hv' if arch.startswith("stable-diffusion-v1"): return 'sd1' @@ -65,6 +68,8 @@ class NetworkOnDisk: return 'sc' if arch.startswith("flux"): return 'f1' + if arch.startswith("hunyuan-video"): + return 'hv' if "v1-5" in str(self.metadata.get('ss_sd_model_name', "")): return 'sd1' diff --git a/modules/modeldata.py b/modules/modeldata.py index 012372b3b..078a5b372 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -45,7 +45,7 @@ def get_model_type(pipe): model_type = 'cogvideox' elif "Sana" in name: model_type = 'sana' - elif 'HunyuanVideoPipeline' in name: + elif 'HunyuanVideoPipeline' in name or 'HunyuanSkyreels' in name: model_type = 'hunyuanvideo' else: model_type = name diff --git a/modules/modelloader.py b/modules/modelloader.py index f3621fdc3..2a1d6745a 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -64,6 +64,9 @@ def download_civit_meta(model_path: str, model_id): def download_civit_preview(model_path: str, preview_url: str): ext = os.path.splitext(preview_url)[1] preview_file = os.path.splitext(model_path)[0] + ext + if preview_file.endswith('.mp4'): + shared.log.warning(f'CivitAI download: url="{preview_url}" skip video') + return '' if os.path.exists(preview_file): return '' res = f'CivitAI download: url={preview_url} file="{preview_file}"' diff --git a/modules/processing.py b/modules/processing.py index 3a49f99ba..67d86021b 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -6,7 +6,7 @@ import numpy as np from PIL import Image, ImageOps from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_models, sd_checkpoint, sd_vae, processing_helpers, timer, face_restoration, token_merge from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet -from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl # pylint: disable=unused-import +from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl, StableDiffusionProcessingVideo # pylint: disable=unused-import from modules.processing_info import create_infotext from modules.modeldata import model_data from modules import pag diff --git a/modules/processing_class.py b/modules/processing_class.py index 9b61c4539..a0d2104b5 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -360,6 +360,15 @@ class StableDiffusionProcessing: self.scripts = None +class StableDiffusionProcessingVideo(StableDiffusionProcessing): + def __init__(self, **kwargs): + self.prompt_template: str = None + self.frames: int = 1 + self.scheduler_shift: float = 0.0 + self.vae_tile_frames: int = 0 + debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access + super().__init__(**kwargs) + class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing): def __init__(self, **kwargs): debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access @@ -592,9 +601,6 @@ class StableDiffusionProcessingControl(StableDiffusionProcessingImg2Img): self.hr_upscale_to_x, self.hr_upscale_to_y = 8 * int(self.width * scale / 8), 8 * int(self.height * scale / 8) else: self.hr_upscale_to_x, self.hr_upscale_to_y = self.hr_resize_x, self.hr_resize_y - # hypertile_set(self, hr=True) - # shared.state.job_count = 2 * self.n_iter - # shared.log.debug(f'Control refine: upscaler="{self.hr_upscaler}" scale={scale} fixed={not use_scale} size={self.hr_upscale_to_x}x{self.hr_upscale_to_y}') def switch_class(p: StableDiffusionProcessing, new_class: type, dct: dict = None): diff --git a/modules/processing_info.py b/modules/processing_info.py index 5677a538c..3b0e1bd24 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -50,7 +50,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No "Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None, "Sampler": p.sampler_name if p.sampler_name != 'Default' else None, "Seed": all_seeds[index], - "Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", + "Seed resize from": None if p.seed_resize_from_w <= 0 or p.seed_resize_from_h <= 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", "CFG scale": p.cfg_scale if p.cfg_scale > 1.0 else None, "CFG rescale": p.diffusers_guidance_rescale if p.diffusers_guidance_rescale > 0 else None, "CFG end": p.cfg_end if p.cfg_end < 1.0 else None, diff --git a/modules/ui.py b/modules/ui.py index 371a26945..91fcf3cf3 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -147,7 +147,7 @@ def create_ui(startup_timer = None): timer.startup.record("ui-control") with gr.Blocks(analytics_enabled=False) as video_interface: - if shared.native and shared.cmd_opts.experimental: + if shared.native: from modules import ui_video ui_video.create_ui() timer.startup.record("ui-video") diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 20b1cfc50..d2e73bdcc 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -121,18 +121,18 @@ def create_batch_inputs(tab, accordion=True): return batch_count, batch_size -def create_seed_inputs(tab, reuse_visible=True): - with gr.Accordion(open=False, label="Seed", elem_id=f"{tab}_seed_group", elem_classes=["small-accordion"]): +def create_seed_inputs(tab, reuse_visible=True, accordion=True, subseed_visible=True, seed_resize_visible=False): + with gr.Accordion(open=False, label="Seed", elem_id=f"{tab}_seed_group", elem_classes=["small-accordion"]) if accordion else gr.Group(): with gr.Row(elem_id=f"{tab}_seed_row", variant="compact"): seed = gr.Number(label='Initial seed', value=-1, elem_id=f"{tab}_seed", container=True) random_seed = ToolButton(ui_symbols.random, elem_id=f"{tab}_random_seed", label='Random seed') reuse_seed = ToolButton(ui_symbols.reuse, elem_id=f"{tab}_reuse_seed", label='Reuse seed', visible=reuse_visible) - with gr.Row(elem_id=f"{tab}_subseed_row", variant="compact", visible=True): + with gr.Row(elem_id=f"{tab}_subseed_row", variant="compact", visible=subseed_visible): subseed = gr.Number(label='Variation', value=-1, elem_id=f"{tab}_subseed", container=True) random_subseed = ToolButton(ui_symbols.random, elem_id=f"{tab}_random_subseed") reuse_subseed = ToolButton(ui_symbols.reuse, elem_id=f"{tab}_reuse_subseed", visible=reuse_visible) subseed_strength = gr.Slider(label='Variation strength', value=0.0, minimum=0, maximum=1, step=0.01, elem_id=f"{tab}_subseed_strength") - with gr.Row(visible=False): + with gr.Row(visible=seed_resize_visible): seed_resize_from_w = gr.Slider(minimum=0, maximum=4096, step=8, label="Resize seed from width", value=0, elem_id=f"{tab}_seed_resize_from_w") seed_resize_from_h = gr.Slider(minimum=0, maximum=4096, step=8, label="Resize seed from height", value=0, elem_id=f"{tab}_seed_resize_from_h") random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) @@ -150,7 +150,7 @@ def create_video_inputs(tab:str): ] with gr.Column(): video_codecs = ['None', 'GIF', 'PNG', 'MP4/MP4V', 'MP4/AVC1', 'MP4/JVT3', 'MKV/H264', 'AVI/DIVX', 'AVI/RGBA', 'MJPEG/MJPG', 'MPG/MPG1', 'AVR/AVR1'] - video_type = gr.Dropdown(label='Video type', choices=video_codecs, value='None', elem_id=f"{tab}_video_type") + video_type = gr.Dropdown(label='Save video', choices=video_codecs, value='None', elem_id=f"{tab}_video_type") with gr.Column(): video_duration = gr.Slider(label='Duration', minimum=0.25, maximum=300, step=0.25, value=2, visible=False, elem_id=f"{tab}_video_duration") video_loop = gr.Checkbox(label='Loop', value=True, visible=False, elem_id=f"{tab}_video_loop") diff --git a/modules/ui_video.py b/modules/ui_video.py index 130aebd0d..97529e188 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -1,40 +1,21 @@ -# TODO hunyuanvideo: seed, scheduler, scheduler_shift, guidance_scale=1.0, true_cfg_scale=6.0, num_inference_steps=30, prompt_template, vae, offloading +# TODO hunyuanvideo: prompt_template, lora +# TODO hunyuanvideo: teacache, pab, fastercache, paraattention, perflow # TODO modernui video tab -from dataclasses import dataclass import gradio as gr -from modules import shared, images, ui_common, ui_sections, sd_models, call_queue, generation_parameters_copypaste +from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste +from modules.ui_components import ToolButton from modules.video_models import hunyuan -@dataclass -class Model(): - name: str - repo: str - dit: str - - -MODELS = { - 'None': [], - 'Hunyuan Video': [ - Model('None', None, None), - Model('Hunyuan Video T2V', 'hunyuanvideo-community/HunyuanVideo', None), - Model('Hunyuan Video I2V', 'hunyuanvideo-community/HunyuanVideo', 'hunyuanvideo-community/HunyuanVideo-I2V'), # https://github.com/huggingface/diffusers/pull/10983 - Model('SkyReels Hunyuan T2V', 'hunyuanvideo-community/HunyuanVideo', 'Skywork/SkyReels-V1-Hunyuan-T2V'), # https://github.com/huggingface/diffusers/pull/10837 - Model('SkyReels Hunyuan I2V', 'hunyuanvideo-community/HunyuanVideo', 'Skywork/SkyReels-V1-Hunyuan-I2V'), - Model('Fast Hunyuan T2V', 'hunyuanvideo-community/HunyuanVideo', 'hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt'), # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 - ] -} - - def engine_change(engine): - models = [model.name for model in MODELS.get(engine, [])] - return gr.update(choices=models, value=models[0] if len(models) > 0 else None) + found = [model.name for model in hunyuan.models.get(engine, [])] + return gr.update(choices=found, value=found[0] if len(found) > 0 else None) def model_change(engine, model): - models = [model.name for model in MODELS.get(engine, [])] - selected = [m for m in MODELS[engine] if m.name == model][0] if len(models) > 0 else None + found = [model.name for model in hunyuan.models.get(engine, [])] + selected = [m for m in hunyuan.models[engine] if m.name == model][0] if len(found) > 0 else None if selected: if 'None' in selected.name: sd_models.unload_model_weights() @@ -52,8 +33,8 @@ def model_change(engine, model): def run_video(*args): engine, model = args[2], args[3] - models = [model.name for model in MODELS.get(engine, [])] - selected = [m for m in MODELS[engine] if m.name == model][0] if len(models) > 0 else None + found = [model.name for model in hunyuan.models.get(engine, [])] + selected = [m for m in hunyuan.models[engine] if m.name == model][0] if len(found) > 0 else None if selected and 'Hunyuan' in selected.name: return hunyuan.generate(*args) shared.log.error(f'Video model not found: args={args}') @@ -61,38 +42,60 @@ def run_video(*args): def create_ui(): - shared.log.debug('UI initialize: txt2img') + shared.log.debug('UI initialize: video') with gr.Blocks(analytics_enabled=False) as _video_interface: - prompt, styles, _negative, generate, _reprocess, paste, _networks, _token_counter, _token_button, _token_counter_negative, _token_button_negative = ui_sections.create_toprow(is_img2img=False, id_part="video", negative_visible=False, reprocess_visible=False) + prompt, styles, negative, generate, _reprocess, paste, networks_button, _token_counter, _token_button, _token_counter_negative, _token_button_negative = ui_sections.create_toprow(is_img2img=False, id_part="video", negative_visible=True, reprocess_visible=False) prompt_image = gr.File(label="", elem_id="video_prompt_image", file_count="single", type="binary", visible=False) prompt_image.change(fn=images.image_data, inputs=[prompt_image], outputs=[prompt, prompt_image]) + with gr.Row(variant='compact', elem_id="video_extra_networks", elem_classes=["extra_networks_root"], visible=False) as extra_networks_ui: + from modules import ui_extra_networks + extra_networks_ui = ui_extra_networks.create_ui(extra_networks_ui, networks_button, 'video', skip_indexing=shared.opts.extra_network_skip_indexing) + timer.startup.record('ui-networks') + with gr.Row(elem_id="video_interface", equal_height=False): with gr.Column(variant='compact', elem_id="video_settings", scale=1): with gr.Row(): - engine = gr.Dropdown(label='Engine', choices=list(MODELS), value='None') - model = gr.Dropdown(label='Model', choices=[''], value=None) + engine = gr.Dropdown(label='Engine', choices=list(hunyuan.models), value='None', elem_id="video_engine") + model = gr.Dropdown(label='Model', choices=[''], value=None, elem_id="video_model") with gr.Row(): width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) with gr.Row(): - frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15) + frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15, elem_id="video_frames") + seed = gr.Number(label='Initial seed', value=-1, elem_id="video_seed", container=True) + random_seed = ToolButton(ui_symbols.random, elem_id="video_random_seed", label='Random seed') + reuse_seed = ToolButton(ui_symbols.reuse, elem_id="video_reuse_seed", label='Reuse seed') + steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") with gr.Row(): - with gr.Group(visible=False) as image_group: + sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") + with gr.Row(): + guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") + guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") + with gr.Row(): + vae_type = gr.Dropdown(label='VAE decode', choices=['Default', 'Tiny', 'Remote'], value='Default', elem_id="video_vae_type") + vae_tile_frames = gr.Slider(label='Tile frames', minimum=1, maximum=64, step=1, value=16, elem_id="video_vae_tile_frames") + with gr.Row(): + with gr.Group(visible=False, elem_id='video_init_image') as image_group: gr.HTML("
  Init image") - image = gr.Image(elem_id="video_image", show_label=False, source="upload", interactive=True, type="pil", tool="select", image_mode="RGB", height=512) + init_image = gr.Image(elem_id="video_image", show_label=False, source="upload", interactive=True, type="pil", tool="select", image_mode="RGB", height=512) with gr.Row(): - save_frames = gr.Checkbox(label='Save image frames', value=False) + save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") with gr.Row(): - cc, duration, loop, pad, interpolate = ui_sections.create_video_inputs(tab='video') + video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video') override_settings = ui_common.create_override_inputs('video') # output panel with gallery gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) + # connect reuse seed button + ui_common.connect_reuse_seed(seed, reuse_seed, gen_info, is_subseed=False) + random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) + # handle engine and model change + engine.change(fn=engine_change, inputs=[engine], outputs=[model]) + model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group]) + # setup extra networks + ui_extra_networks.setup_ui(extra_networks_ui, gallery) - # handle engine and model change - engine.change(fn=engine_change, inputs=[engine], outputs=[model]) - model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group]) # handle restore fields paste_fields = [ (prompt, "Prompt"), @@ -101,7 +104,7 @@ def create_ui(): (height, "Size-2"), (frames, "Frames"), ] - generation_parameters_copypaste.add_paste_fields("txt2img", None, paste_fields, override_settings) + generation_parameters_copypaste.add_paste_fields("video", None, paste_fields, override_settings) bindings = generation_parameters_copypaste.ParamBinding(paste_button=paste, tabname="video", source_text_component=prompt, source_image_component=None) generation_parameters_copypaste.register_paste_params_button(bindings) # hidden fields @@ -111,12 +114,18 @@ def create_ui(): video_args = [ task_id, ui_state, engine, model, - prompt, styles, + prompt, negative, styles, width, height, frames, - image, + steps, sampler_index, + sampler_shift, + seed, + guidance_scale, guidance_true, + init_image, + vae_type, vae_tile_frames, save_frames, - cc, duration, loop, pad, interpolate, + video_type, video_duration, video_loop, video_pad, video_interpolate, + override_settings, ] # generate function video_dict = dict( diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py index 6fed27cb3..051ecde85 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/hunyuan.py @@ -1,13 +1,260 @@ -from modules import shared +from dataclasses import dataclass +import os +import time +import torch +import transformers +import diffusers +from modules import shared, sd_models, sd_checkpoint, sd_samplers, processing, model_quant, devices, images, timer, ui_common + + +@dataclass +class Model(): + name: str + repo: str + dit: str + subfolder: str + +models = { + 'None': [], + 'Hunyuan Video': [ + Model(name='None', repo=None, dit=None, subfolder=None), + Model(name='Hunyuan Video T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='hunyuanvideo-community/HunyuanVideo', subfolder='transformer'), + Model(name='Hunyuan Video I2V', repo='hunyuanvideo-community/HunyuanVideo-I2V', dit='hunyuanvideo-community/HunyuanVideo-I2V', subfolder='transformer'), # https://github.com/huggingface/diffusers/pull/10983 + Model(name='SkyReels Hunyuan T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='Skywork/SkyReels-V1-Hunyuan-T2V', subfolder=None), # https://github.com/huggingface/diffusers/pull/10837 + Model(name='SkyReels Hunyuan I2V', repo='hunyuanvideo-community/HunyuanVideo', dit='Skywork/SkyReels-V1-Hunyuan-I2V', subfolder=None), + Model(name='Fast Hunyuan T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='FastVideo/FastHunyuan-diffusers', subfolder='transformer'), # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 + ] +} +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None +loaded_model = None +prompt_template = { + "template": ( + "<|start_header_id|>system<|end_header_id|>" + "\nDescribe the video by detailing the following aspects: \n" + "1. The main content and theme of the video.\n" + "2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects.\n" + "3. Actions, events, behaviors, temporal relationships, and physical movement changes of the objects.\n" + "4. Background environment, light, style and atmosphere.\n" + "5. Camera angles, movements, and transitions used in the video.\n" + "<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>" + ), + "crop_start": 95, +} + + +def hijack_decode(*args, **kwargs): + t0 = time.time() + vae: diffusers.AutoencoderKLHunyuanVideo = shared.sd_model.vae + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) + res = shared.sd_model.vae.orig_decode(*args, **kwargs) + t1 = time.time() + timer.process.add('vae', t1-t0) + shared.log.debug(f'Video: vae={vae.__class__.__name__} tile={vae.tile_sample_min_width}:{vae.tile_sample_min_height}:{vae.tile_sample_min_num_frames} stride={vae.tile_sample_stride_width}:{vae.tile_sample_stride_height}:{vae.tile_sample_stride_num_frames} time={t1-t0:.2f}') + return res + + +def hijack_encode_prompt(*args, **kwargs): + t0 = time.time() + res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + t1 = time.time() + timer.process.add('te', t1-t0) + shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + return res def load(selected): - msg = f'Video load: model="{selected.name}" repo="{selected.repo}" dit="{selected.dit}"' + if selected is None: + return + global loaded_model # pylint: disable=global-statement + if loaded_model == selected.name: + return + sd_models.unload_model_weights() + t0 = time.time() + + quant_args = model_quant.create_config(module='Model') + cls = diffusers.HunyuanVideoTransformer3DModel + try: + debug(f'Video load: module=transofrmer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} quant={quant_args is not None}') + transformer = cls.from_pretrained( + pretrained_model_name_or_path=selected.dit, + subfolder=selected.subfolder, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=transformer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} {e}') + + quant_args = model_quant.create_config(module='Text Encoder') + if 'I2V' in selected.repo: + cls = transformers.LlavaForConditionalGeneration + else: + cls = transformers.LlamaModel + try: + debug(f'Video load: module=te repo="{selected.repo}" cls={cls.__name__} quant={quant_args is not None}') + text_encoder = cls.from_pretrained( + pretrained_model_name_or_path=selected.repo, + subfolder="text_encoder", + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=te repo="{selected.repo}" cls={cls.__name__} {e}') + + cls = transformers.CLIPTextModel + try: + debug(f'Video load: module=clip repo="{selected.repo}" cls={cls.__name__} quant=False') + text_encoder_2 = transformers.CLIPTextModel.from_pretrained( + pretrained_model_name_or_path=selected.repo, + subfolder="text_encoder_2", + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + ) + except Exception as e: + shared.log.error(f'video load: module=clip repo="{selected.repo}" cls={cls.__name__} {e}') + + cls = diffusers.AutoencoderKLHunyuanVideo + try: + debug(f'Video load: module=vae repo="{selected.repo}" cls={cls.__name__} quant=False') + vae = diffusers.AutoencoderKLHunyuanVideo.from_pretrained( + pretrained_model_name_or_path=selected.repo, + subfolder="vae", + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + ) + except Exception as e: + shared.log.error(f'video load: module=vae repo="{selected.repo}" cls={cls.__name__} {e}') + + if selected.name == 'Hunyuan Video I2V': + cls = diffusers.HunyuanVideoImageToVideoPipeline + elif selected.name == 'SkyReels Hunyuan I2V': + cls = diffusers.HunyuanSkyreelsImageToVideoPipeline + else: + cls = diffusers.HunyuanVideoPipeline + try: + debug(f'Video load: module=pipe repo="{selected.repo}" cls={cls.__name__} quant=False') + shared.sd_model = cls.from_pretrained( + pretrained_model_name_or_path=selected.repo, + transformer=transformer, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + vae=vae, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + ) + except Exception as e: + shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={cls.__name__} {e}') + + t1 = time.time() + sd_models.set_diffuser_options(shared.sd_model) + shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) + shared.sd_model.sd_model_hash = None + shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode + shared.sd_model.vae.decode = hijack_decode + shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt + shared.sd_model.encode_prompt = hijack_encode_prompt + shared.sd_model.vae.enable_slicing() + loaded_model = selected.name + msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' shared.log.info(msg) return msg def generate(*args, **kwargs): - # TODO hunyuanvideo: check if loaded - shared.log.debug(f'Video generate: args={args} kwargs={kwargs}') - return [], '', '', 'TBD' + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + shared.log.error('Video: model not selected') + return [], '', '', 'Video model not selected' + if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: + found = [model.name for model in models.get(engine, [])] + selected = [m for m in models[engine] if m.name == model][0] if len(found) > 0 else None + load(selected) + if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: + shared.log.error('Video: model not loaded') + return [], '', '', 'Video model not loaded' + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=16 * int(width // 16), + height=16 * int(height // 16), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + p.scripts = None + p.script_args = args + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + shared.log.error('Video: init image not set') + return [], '', '', 'Error: init image not set' + p.task_args['image'] = init_image + + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # handle sampler and seed + if p.sampler_name != 'Default': + shared.sd_model.scheduler = sd_samplers.create_sampler(p.sampler_name, shared.sd_model) + p.sampler_name = 'Default' # avoid double creation + if hasattr(shared.sd_model.scheduler, '_shift') and sampler_shift > 0: + shared.sd_model.scheduler._shift = sampler_shift # pylint: disable=protected-access + + # handle vae + if vae_tile_frames > p.frames: + shared.sd_model.vae.tile_sample_min_num_frames = vae_tile_frames + shared.sd_model.vae.use_framewise_decoding = True + shared.sd_model.vae.enable_tiling() + else: + shared.sd_model.vae.use_framewise_decoding = False + shared.sd_model.vae.disable_tiling() + + # set args + processing.fix_seed(p) + p.prompt = shared.prompt_styles.apply_styles_to_prompt(prompt, p.styles) + p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(negative, p.styles) + p.task_args['width'] = p.width + p.task_args['height'] = p.height + p.task_args['num_inference_steps'] = p.steps + p.task_args['num_frames'] = p.frames + p.task_args['generator'] = torch.manual_seed(p.seed) + p.task_args['guidance_scale'] = p.cfg_scale + p.task_args['true_cfg_scale'] = p.diffusers_guidance_rescale + p.task_args['prompt_template'] = prompt_template + p.task_args['output_type'] = 'pil' + p.task_args['prompt'] = p.prompt + p.task_args['negative_prompt'] = p.negative_prompt + p.ops.append('video') + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + t0 = time.time() + processed = processing.process_images(p) + t1 = time.time() + shared.state.disable_preview = False + + p.close() + if processed is None or len(processed.images) == 0: + return [], '', '', 'Error: processing failed' + shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') + if video_type != 'None': + images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + + generation_info_js = processed.js() if processed is not None else '' + return processed.images, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) From 9ac936ab140d048455655c52f925cd5f9d0757e6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 18 Mar 2025 19:00:30 -0400 Subject: [PATCH 022/122] update api-grid Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 54 +++++++++++++++++++++++++++++++++---------------- 1 file changed, 37 insertions(+), 17 deletions(-) diff --git a/cli/api-grid.py b/cli/api-grid.py index 21eb685a4..6b0b7d5d2 100755 --- a/cli/api-grid.py +++ b/cli/api-grid.py @@ -18,6 +18,7 @@ class Options: # set default parameters here negative_prompt: str = '' seed: int = -1 steps: int = 20 + cfg_scale: float = 6.0 sampler_name: str = "Default" width: int = 1024 height: int = 1024 @@ -32,6 +33,7 @@ class Server: # set server and save options here or use command line arguments user: str = None password: str = None folder: str = '/tmp' + name: str = str(round(time.time())) images: bool = False grids: bool = False labels: bool = False @@ -56,7 +58,7 @@ def post(): return { 'error': 0, 'reason': str(e), 'url': server.url } -def generate(ts: float, x: int, y: int): # pylint: disable=redefined-outer-name +def generate(x: int, y: int): # pylint: disable=redefined-outer-name t0 = time.time() log.info(f'x={x} y={y} {options}') data = post() @@ -68,7 +70,7 @@ def generate(ts: float, x: int, y: int): # pylint: disable=redefined-outer-name image = Image.open(io.BytesIO(base64.b64decode(b64))) images.append(image) info = data['info'] - fn = os.path.join(server.folder, f'{round(ts)}-{x}-{y}.jpg') if server.images else None + fn = os.path.join(server.folder, f'{server.name}-{x}-{y}.jpg') if server.images else None log.info(f'image: time={t1-t0:.2f} size={image.size} fn="{fn}" info="{info}"') if fn is not None: image.save(fn) @@ -118,7 +120,7 @@ def grid(x_file: str, y_file: str): x = [line for line in x if ':' in line] y = [line for line in y if ':' in line] t0 = time.time() - log.info(f'grid: x={len(x)} y={len(y)} prefix={round(t0)}') + log.info(f'grid: x={len(x)} y={len(y)} prefix={server.name}') vertical = [] Image.MAX_IMAGE_PIXELS = None for j in range(max(1, len(y))): @@ -129,7 +131,7 @@ def grid(x_file: str, y_file: str): set_param(x[i]) if len(y) > i: set_param(y[i]) - images = generate(t0, i, j) + images = generate(i, j) if images is not None and len(images) > 0: horizontal.extend(images) labels.append(f'{x[i] if len(x) > i else ""}\n{y[j] if len(y) > j else ""}') @@ -145,7 +147,7 @@ def grid(x_file: str, y_file: str): log.warning('grid: empty grid') return merged = merge(vertical, horizontal=False) - fn = os.path.join(server.folder, f'{round(t0)}.jpg') + fn = os.path.join(server.folder, f'{server.name}.jpg') merged.save(fn) log.info(f'grid: size={merged.size} fn="{fn}"') t1 = time.time() @@ -155,22 +157,40 @@ def grid(x_file: str, y_file: str): if __name__ == "__main__": log.info(__file__) parser = argparse.ArgumentParser(description = 'api-txt2img') - parser.add_argument('--x', required=False, default=None, help='file to use for x-axis values') - parser.add_argument('--y', required=False, default=None, help='file to use for y-axis values') - parser.add_argument('--folder', required=False, default='/tmp', help='folder to use for saving images') - parser.add_argument('--image', required=False, default=False, help='save individual images') - parser.add_argument('--grid', required=False, default=True, help='save image grids') - parser.add_argument('--labels', required=False, default=True, help='draw image labels') - parser.add_argument('--url', required=False, default='http://127.0.0.1:7860', help='server url') - parser.add_argument('--user', required=False, default=None, help='server user') - parser.add_argument('--password', required=False, default=None, help='server password') + parser.add_argument('--x', type=str, required=False, default=None, help='file to use for x-axis values') + parser.add_argument('--y', type=str, required=False, default=None, help='file to use for y-axis values') + parser.add_argument('--folder', type=str, required=False, default='/tmp', help='folder to use for saving images') + parser.add_argument('--name', type=str, required=False, default=str(round(time.time())), help='name prefix to use for saving images and grids') + parser.add_argument('--image', type=bool, required=False, default=False, help='save individual images') + parser.add_argument('--grid', type=bool, required=False, default=True, help='save image grids') + parser.add_argument('--labels', type=bool, required=False, default=True, help='draw image labels') + parser.add_argument('--url', type=str, required=False, default='http://127.0.0.1:7860', help='server url') + parser.add_argument('--user', type=str, required=False, default=None, help='server user') + parser.add_argument('--password', type=str, required=False, default=None, help='server password') + parser.add_argument('--prompt', type=str, required=False, default='', help='generate prompt') + parser.add_argument('--negative', type=str, required=False, default='', help='generate negative prompt') + parser.add_argument('--sampler', type=str, required=False, default='Default', help='generate sampler') + parser.add_argument('--width', type=int, required=False, default=1024, help='generate width') + parser.add_argument('--height', type=int, required=False, default=1024, help='generate height') + parser.add_argument('--steps', type=int, required=False, default=20, help='generate steps') + parser.add_argument('--cfg', type=float, required=False, default=6.0, help='generate guidance scale') + parser.add_argument('--seed', type=int, required=False, default=-1, help='generate seed') args = parser.parse_args() log.info(args) server.folder = args.folder - server.images = args.image - server.grids = args.grid - server.labels = args.labels + server.name = args.name + server.images = bool(args.image) + server.grids = bool(args.grid) + server.labels = bool(args.labels) server.url = args.url server.user = args.user server.password = args.password + options.prompt = args.prompt + options.negative_prompt = args.negative + options.width = int(args.width) + options.height = int(args.height) + options.sampler_name = args.sampler + options.seed = int(args.seed) + options.steps = int(args.steps) + options.cfg_scale = float(args.cfg) grid(args.x, args.y) From a88f919ceec081d8a2443f6919157ad824e310be Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 18 Mar 2025 19:18:46 -0400 Subject: [PATCH 023/122] continuing on video tab Signed-off-by: Vladimir Mandic --- modules/ui_video.py | 14 ++++++++++---- modules/video_models/hunyuan.py | 14 ++++++++------ 2 files changed, 18 insertions(+), 10 deletions(-) diff --git a/modules/ui_video.py b/modules/ui_video.py index 97529e188..a54296542 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -38,7 +38,7 @@ def run_video(*args): if selected and 'Hunyuan' in selected.name: return hunyuan.generate(*args) shared.log.error(f'Video model not found: args={args}') - return [], '', '', f'Video model not found: engine={engine} model={model}' + return [], None, '', '', f'Video model not found: engine={engine} model={model}' def create_ui(): @@ -85,8 +85,14 @@ def create_ui(): video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video') override_settings = ui_common.create_override_inputs('video') - # output panel with gallery - gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) + # output panel with gallery and video tabs + with gr.Column(elem_id='video-output-column', scale=3) as _column_output: + with gr.Tabs(elem_classes=['video-output-tabs'], elem_id='video-output-tabs'): + with gr.Tab('Frames', id='out-gallery'): + gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) + with gr.Tab('Video', id='out-video'): + video = gr.Video(label="Output", show_label=False, elem_id='control_output_video', elem_classes=['control-image']) + # connect reuse seed button ui_common.connect_reuse_seed(seed, reuse_seed, gen_info, is_subseed=False) random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) @@ -132,7 +138,7 @@ def create_ui(): fn=call_queue.wrap_gradio_gpu_call(run_video, extra_outputs=[None, '', ''], name='Video'), _js="submit_video", inputs=video_args, - outputs=[gallery, gen_info, html_info, html_log], + outputs=[gallery, video, gen_info, html_info, html_log], show_progress=False, ) prompt.submit(**video_dict) diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py index 051ecde85..dd83ce2d6 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/hunyuan.py @@ -166,14 +166,14 @@ def generate(*args, **kwargs): task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': shared.log.error('Video: model not selected') - return [], '', '', 'Video model not selected' + return [], None, '', '', 'Video model not selected' if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: found = [model.name for model in models.get(engine, [])] selected = [m for m in models[engine] if m.name == model][0] if len(found) > 0 else None load(selected) if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: shared.log.error('Video: model not loaded') - return [], '', '', 'Video model not loaded' + return [], None, '', '', 'Video model not loaded' debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') p = processing.StableDiffusionProcessingVideo( @@ -201,7 +201,7 @@ def generate(*args, **kwargs): if 'I2V' in model: if init_image is None: shared.log.error('Video: init image not set') - return [], '', '', 'Error: init image not set' + return [], None, '', '', 'Error: init image not set' p.task_args['image'] = init_image shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) @@ -251,10 +251,12 @@ def generate(*args, **kwargs): p.close() if processed is None or len(processed.images) == 0: - return [], '', '', 'Error: processing failed' + return [], None, '', '', 'Error: processing failed' shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + else: + video_file = None generation_info_js = processed.js() if processed is not None else '' - return processed.images, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) From 4e154ee28fd8a75f861b2b9b42367b8ee5c21e18 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 18 Mar 2025 19:21:13 -0400 Subject: [PATCH 024/122] cleanup Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 2 +- wiki | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/cli/api-grid.py b/cli/api-grid.py index 6b0b7d5d2..98087c1af 100755 --- a/cli/api-grid.py +++ b/cli/api-grid.py @@ -156,7 +156,7 @@ def grid(x_file: str, y_file: str): if __name__ == "__main__": log.info(__file__) - parser = argparse.ArgumentParser(description = 'api-txt2img') + parser = argparse.ArgumentParser(description = 'api-grid') parser.add_argument('--x', type=str, required=False, default=None, help='file to use for x-axis values') parser.add_argument('--y', type=str, required=False, default=None, help='file to use for y-axis values') parser.add_argument('--folder', type=str, required=False, default='/tmp', help='folder to use for saving images') diff --git a/wiki b/wiki index 37bc610f4..6593476c0 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 37bc610f46affcdf417cfea61324acd505cad723 +Subproject commit 6593476c0c61b4fed8594a12ff56601a440b48b0 From 1569dc6d07ce03f6fcc409275e2ce909313c809b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 08:37:44 -0400 Subject: [PATCH 025/122] fix api-grid Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 12 +++++++----- modules/ui_video.py | 4 ++-- modules/video_models/hunyuan.py | 22 ++++++++++++++++------ 3 files changed, 25 insertions(+), 13 deletions(-) diff --git a/cli/api-grid.py b/cli/api-grid.py index 98087c1af..0d6e53314 100755 --- a/cli/api-grid.py +++ b/cli/api-grid.py @@ -102,12 +102,14 @@ def merge(images: list[Image.Image], horizontal: bool, labels: list[str] = None) def grid(x_file: str, y_file: str): def set_param(line): param = line.split(':', maxsplit=1) - if param[0] == 'prompt': - options.prompt += f'{param[1]} ' # prompt is appended so its not overwritten - elif param[0] == 'lora': - options.prompt += f' ' # lora is appended to prompt + k = param[0].strip() + v = param[1].strip() if len(param) > 1 else '' + if k == 'prompt': + options.prompt += f'{v} ' # prompt is appended so its not overwritten + elif k == 'lora': + options.prompt += f' ' # lora is appended to prompt else: - setattr(options, param[0].strip(), param[1].strip()) + setattr(options, k, v) log.info(server) os.makedirs(server.folder, exist_ok=True) diff --git a/modules/ui_video.py b/modules/ui_video.py index a54296542..97b9d2ec7 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -78,7 +78,7 @@ def create_ui(): with gr.Row(): with gr.Group(visible=False, elem_id='video_init_image') as image_group: gr.HTML("
  Init image") - init_image = gr.Image(elem_id="video_image", show_label=False, source="upload", interactive=True, type="pil", tool="select", image_mode="RGB", height=512) + init_image = gr.Image(elem_id="video_image", show_label=False, type="pil", image_mode="RGB", height=512) with gr.Row(): save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") with gr.Row(): @@ -91,7 +91,7 @@ def create_ui(): with gr.Tab('Frames', id='out-gallery'): gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) with gr.Tab('Video', id='out-video'): - video = gr.Video(label="Output", show_label=False, elem_id='control_output_video', elem_classes=['control-image']) + video = gr.Video(label="Output", show_label=False, elem_id='control_output_video', elem_classes=['control-image'], height=512, autoplay=False) # connect reuse seed button ui_common.connect_reuse_seed(seed, reuse_seed, gen_info, is_subseed=False) diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py index dd83ce2d6..c95316329 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/hunyuan.py @@ -4,7 +4,7 @@ import time import torch import transformers import diffusers -from modules import shared, sd_models, sd_checkpoint, sd_samplers, processing, model_quant, devices, images, timer, ui_common +from modules import shared, errors, sd_models, sd_checkpoint, sd_samplers, processing, model_quant, devices, images, timer, ui_common @dataclass @@ -49,7 +49,7 @@ def hijack_decode(*args, **kwargs): res = shared.sd_model.vae.orig_decode(*args, **kwargs) t1 = time.time() timer.process.add('vae', t1-t0) - shared.log.debug(f'Video: vae={vae.__class__.__name__} tile={vae.tile_sample_min_width}:{vae.tile_sample_min_height}:{vae.tile_sample_min_num_frames} stride={vae.tile_sample_stride_width}:{vae.tile_sample_stride_height}:{vae.tile_sample_stride_num_frames} time={t1-t0:.2f}') + debug(f'Video: vae={vae.__class__.__name__} tile={vae.tile_sample_min_width}:{vae.tile_sample_min_height}:{vae.tile_sample_min_num_frames} stride={vae.tile_sample_stride_width}:{vae.tile_sample_stride_height}:{vae.tile_sample_stride_num_frames} time={t1-t0:.2f}') return res @@ -58,7 +58,7 @@ def hijack_encode_prompt(*args, **kwargs): res = shared.sd_model.orig_encode_prompt(*args, **kwargs) t1 = time.time() timer.process.add('te', t1-t0) - shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) return res @@ -98,6 +98,8 @@ def load(selected): subfolder="text_encoder", cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, + # torch_dtype='auto', # special case as text and vision nested models have different dtypes + # attn_implementation="flash_attention_2", # testing different attention types **quant_args ) except Exception as e: @@ -203,6 +205,7 @@ def generate(*args, **kwargs): shared.log.error('Video: init image not set') return [], None, '', '', 'Error: init image not set' p.task_args['image'] = init_image + # p.task_args['image'] = init_image.resize((336, 336), Image.Resampling.LANCZOS) shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) devices.torch_gc(force=True) @@ -234,7 +237,7 @@ def generate(*args, **kwargs): p.task_args['generator'] = torch.manual_seed(p.seed) p.task_args['guidance_scale'] = p.cfg_scale p.task_args['true_cfg_scale'] = p.diffusers_guidance_rescale - p.task_args['prompt_template'] = prompt_template + # p.task_args['prompt_template'] = prompt_template # t2v and i2v have different templates p.task_args['output_type'] = 'pil' p.task_args['prompt'] = p.prompt p.task_args['negative_prompt'] = p.negative_prompt @@ -245,13 +248,20 @@ def generate(*args, **kwargs): shared.state.disable_preview = True shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') t0 = time.time() - processed = processing.process_images(p) + try: + processed = processing.process_images(p) + except Exception as e: + shared.log.error(f'Video: exception={e}') + errors.display(e, 'video') + processed = None + shared.state.disable_preview = False + return [], None, '', '', str(e) t1 = time.time() shared.state.disable_preview = False p.close() if processed is None or len(processed.images) == 0: - return [], None, '', '', 'Error: processing failed' + return [], None, '', '', 'Video: processing failed' shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') if video_type != 'None': video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) From 1e0f512ccbde9988459ecff173c80aaf51a0b4fa Mon Sep 17 00:00:00 2001 From: Disty0 Date: Wed, 19 Mar 2025 15:42:36 +0300 Subject: [PATCH 026/122] ROCm disable FP16 for gfx1102 --- CHANGELOG.md | 5 ++++- modules/devices.py | 15 ++++++++++++--- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index ca168767f..7a12b6492 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -51,12 +51,15 @@ Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP mode - update `diffusers` and other requirements - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - **IPEX** - - add `--upgrade` to torch_command when using `--use-nightly` for *ipex* and *rocm* + - add `--upgrade` to torch_command when using `--use-nightly` - add xpu to profiler - fix untyped_storage, torch.eye and torch.cuda.device ops - fix torch 2.7 compatibility - fix performance with balanced offload - fix triton and torch.compile +- **ROCm** + - add `--upgrade` to torch_command when using `--use-nightly` + - disable fp16 for gfx1102 (rx 7600 and rx 7500 series) gpus - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/devices.py b/modules/devices.py index caa544652..bd84ff7d4 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -319,9 +319,18 @@ def test_fp16(): global fp16_ok # pylint: disable=global-statement if fp16_ok is not None: return fp16_ok - if sys.platform == "darwin" or backend == 'openvino': # override - fp16_ok = False - return fp16_ok + if opts.cuda_dtype != 'FP16': # don't override if the user sets it + if sys.platform == "darwin" or backend == 'openvino': # override + fp16_ok = False + return fp16_ok + elif backend == 'rocm': + # gfx1102 (RX 7600, 7500, 7650 and 7700S) causes segfaults with fp16 + # agent can be overriden to gfx1100 to get gfx1102 working with ROCm so check the gpu name as well + agent = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000") + agent_name = getattr(torch.cuda.get_device_properties(device), "name", "AMD Radeon RX 0000") + if agent == "gfx1102" or (agent == "gfx1100" and any(i in agent_name for i in ("7600", "7500", "7650", "7700S"))): + fp16_ok = False + return fp16_ok try: x = torch.tensor([[1.5,.0,.0,.0]]).to(device=device, dtype=torch.float16) layerNorm = torch.nn.LayerNorm(4, eps=0.00001, elementwise_affine=True, dtype=torch.float16, device=device) From 9c0846ba4fab91180439c87a8e1d50795fa4d196 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 09:00:06 -0400 Subject: [PATCH 027/122] cleanup Signed-off-by: Vladimir Mandic --- cli/api-grid.py | 4 ++-- wiki | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/cli/api-grid.py b/cli/api-grid.py index 0d6e53314..e41276bb9 100755 --- a/cli/api-grid.py +++ b/cli/api-grid.py @@ -131,8 +131,8 @@ def grid(x_file: str, y_file: str): for i in range(max(1, len(x))): if len(x) > i: set_param(x[i]) - if len(y) > i: - set_param(y[i]) + if len(y) > j: + set_param(y[j]) images = generate(i, j) if images is not None and len(images) > 0: horizontal.extend(images) diff --git a/wiki b/wiki index 6593476c0..d50882dcb 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 6593476c0c61b4fed8594a12ff56601a440b48b0 +Subproject commit d50882dcb83a1441591b0d491efd143e12c1930a From d5cfd61e50c2ea3984186ab50a0b1ddbdc24f89f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 10:49:45 -0400 Subject: [PATCH 028/122] fix sd35 with batch Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 4 +++- modules/processing_vae.py | 3 --- modules/video_models/hunyuan.py | 17 ++++++++++++----- 3 files changed, 15 insertions(+), 9 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 7a12b6492..67745eacf 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,8 +3,9 @@ ## Update for 2025-03-17 ### TODO - - Gemma3 requires `git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` + - Gemma3 requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - Remote VAE encode for SD15 and Flux.1: + - HunyuanVideo-I2V: ### Highlights for 2025-03-17 @@ -50,6 +51,7 @@ Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP mode - add quantization support to **CogView-3Plus** - update `diffusers` and other requirements - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion + - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` - add xpu to profiler diff --git a/modules/processing_vae.py b/modules/processing_vae.py index be819d658..348bbc0c5 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -293,9 +293,6 @@ def vae_decode(latents, model, output_type='np', vae_type='Full', width=None, he latents = model._unpack_latents(latents, height, width, model.vae_scale_factor) # pylint: disable=protected-access if len(latents.shape) == 3: # lost a batch dim in hires latents = latents.unsqueeze(0) - if latents.shape[0] == 4 and latents.shape[1] != 4: # likely animatediff latent - latents = latents.permute(1, 0, 2, 3) - if latents.shape[-1] <= 4: # not a latent, likely an image decoded = latents.float().cpu().numpy() elif vae_type == 'Full' and hasattr(model, "vae"): diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py index c95316329..2b83e6b63 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/hunyuan.py @@ -63,6 +63,12 @@ def hijack_encode_prompt(*args, **kwargs): return res +def get_quant(args): + if args is not None and "quantization_config" in args: + return args['quantization_config'].__class__.__name__ + return None + + def load(selected): if selected is None: return @@ -75,7 +81,7 @@ def load(selected): quant_args = model_quant.create_config(module='Model') cls = diffusers.HunyuanVideoTransformer3DModel try: - debug(f'Video load: module=transofrmer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} quant={quant_args is not None}') + debug(f'Video load: module=transofrmer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} quant={get_quant(quant_args)}') transformer = cls.from_pretrained( pretrained_model_name_or_path=selected.dit, subfolder=selected.subfolder, @@ -92,7 +98,7 @@ def load(selected): else: cls = transformers.LlamaModel try: - debug(f'Video load: module=te repo="{selected.repo}" cls={cls.__name__} quant={quant_args is not None}') + debug(f'Video load: module=te repo="{selected.repo}" cls={cls.__name__} quant={get_quant(quant_args)}') text_encoder = cls.from_pretrained( pretrained_model_name_or_path=selected.repo, subfolder="text_encoder", @@ -107,7 +113,7 @@ def load(selected): cls = transformers.CLIPTextModel try: - debug(f'Video load: module=clip repo="{selected.repo}" cls={cls.__name__} quant=False') + debug(f'Video load: module=clip repo="{selected.repo}" cls={cls.__name__} quant=None') text_encoder_2 = transformers.CLIPTextModel.from_pretrained( pretrained_model_name_or_path=selected.repo, subfolder="text_encoder_2", @@ -119,7 +125,7 @@ def load(selected): cls = diffusers.AutoencoderKLHunyuanVideo try: - debug(f'Video load: module=vae repo="{selected.repo}" cls={cls.__name__} quant=False') + debug(f'Video load: module=vae repo="{selected.repo}" cls={cls.__name__} quant=None') vae = diffusers.AutoencoderKLHunyuanVideo.from_pretrained( pretrained_model_name_or_path=selected.repo, subfolder="vae", @@ -136,7 +142,7 @@ def load(selected): else: cls = diffusers.HunyuanVideoPipeline try: - debug(f'Video load: module=pipe repo="{selected.repo}" cls={cls.__name__} quant=False') + debug(f'Video load: module=pipe repo="{selected.repo}" cls={cls.__name__} quant=None') shared.sd_model = cls.from_pretrained( pretrained_model_name_or_path=selected.repo, transformer=transformer, @@ -205,6 +211,7 @@ def generate(*args, **kwargs): shared.log.error('Video: init image not set') return [], None, '', '', 'Error: init image not set' p.task_args['image'] = init_image + # from PIL import Image # p.task_args['image'] = init_image.resize((336, 336), Image.Resampling.LANCZOS) shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) From 0c890b50e0f9e27bb2307d4cffec46a3fc16c4f2 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Thu, 20 Mar 2025 23:03:23 +0900 Subject: [PATCH 029/122] proper zluda detection --- modules/devices.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/modules/devices.py b/modules/devices.py index bd84ff7d4..b2a488d7b 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -52,7 +52,8 @@ def has_zluda() -> bool: return False try: dev = torch.device("cuda") - return torch.cuda.get_device_name(dev).endswith("[ZLUDA]") + cc = torch.cuda.get_device_capability(dev) + return cc == (8, 8) except Exception: return False From 878fb9f75b41daa444f010467d9c052b8bb02767 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Thu, 20 Mar 2025 23:03:43 +0900 Subject: [PATCH 030/122] zluda 3.9.1 --- modules/zluda_installer.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index dc3d85e61..4a0b9ffa4 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -45,10 +45,10 @@ def set_default_agent(agent: rocm.Agent): global hipBLASLt_available, hipBLASLt_enabled # pylint: disable=global-statement hipBLASLt_available = is_nightly and os.path.exists(rocm.blaslt_tensile_libpath) - hipBLASLt_enabled = hipBLASLt_available and ((not os.path.exists(path) and nightly) or os.path.exists(os.path.join(path, 'cublasLt.dll'))) + hipBLASLt_enabled = hipBLASLt_available and os.path.exists(os.path.join(rocm.path, "bin", "hipblaslt.dll")) global MIOpen_available # pylint: disable=global-statement - MIOpen_available = is_nightly and (skip_arch_test or agent.gfx_version in (0x908, 0x90a, 0x940, 0x941, 0x942, 0x1030, 0x1100, 0x1101, 0x1102, 0x1150,)) + MIOpen_available = is_nightly and os.path.exists(os.path.join(rocm.path, "bin", "MIOpen.dll")) def is_reinstall_needed() -> bool: # ZLUDA<3.8.7 @@ -60,7 +60,7 @@ def install() -> None: return platform = "windows" - commit = os.environ.get("ZLUDA_HASH", "4d14bf95d4c500863e240a0b1fa82793d0da789b") + commit = os.environ.get("ZLUDA_HASH", "ae0540beb129ffd140226ce956b386619b38f84c") if nightly: platform = "nightly-" + platform urllib.request.urlretrieve(f'https://github.com/lshqqytiger/ZLUDA/releases/download/rel.{commit}/ZLUDA-{platform}-rocm{rocm.version[0]}-amd64.zip', '_zluda') From 135631bdea3fabc11d023753e3c684e1822bf906 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Thu, 20 Mar 2025 23:03:50 +0900 Subject: [PATCH 031/122] zluda enable triton experimentally --- modules/zluda_hijacks.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index 872127cf1..9f622be94 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -9,6 +9,36 @@ def topk(input: torch.Tensor, *args, **kwargs): # pylint: disable=redefined-buil return torch.return_types.topk((values.to(device), indices.to(device),)) +class DeviceProperties: + PROPERTIES_OVERRIDE = {"regs_per_multiprocessor": 65535} + internal: torch._C._CudaDeviceProperties + + def __init__(self, props: torch._C._CudaDeviceProperties): + self.internal = props + + def __getattr__(self, name): + if name in DeviceProperties.PROPERTIES_OVERRIDE: + return DeviceProperties.PROPERTIES_OVERRIDE[name] + return getattr(self.internal, name) + + +__get_device_properties = torch.cuda._get_device_properties # pylint: disable=protected-access +def torch_cuda__get_device_properties(device): + return DeviceProperties(__get_device_properties(device)) + + def do_hijack(): torch.version.hip = rocm.version torch.topk = topk + + torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access + try: + import triton + _get_device_properties = triton.runtime.driver.active.utils.get_device_properties + def triton_runtime_driver_active_utils_get_device_properties(device): + props = _get_device_properties(device) + props["mem_bus_width"] = 384 + return props + triton.runtime.driver.active.utils.get_device_properties = triton_runtime_driver_active_utils_get_device_properties + except Exception: + pass From 9bf68389628065fb7ca8cbc13a4fe80f40914894 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 14:15:54 -0400 Subject: [PATCH 032/122] update video tab Signed-off-by: Vladimir Mandic --- TODO.md | 8 + modules/onnx_impl/execution_providers.py | 6 +- modules/processing_diffusers.py | 4 +- modules/ui_caption.py | 8 +- modules/ui_video.py | 33 ++-- modules/video_models/hunyuan.py | 225 +++-------------------- modules/video_models/ltx.py | 0 modules/video_models/models_def.py | 67 +++++++ modules/video_models/video_utils.py | 109 +++++++++++ scripts/ltxvideo.py | 2 +- 10 files changed, 237 insertions(+), 225 deletions(-) create mode 100644 modules/video_models/ltx.py create mode 100644 modules/video_models/models_def.py create mode 100644 modules/video_models/video_utils.py diff --git a/TODO.md b/TODO.md index 0575e97c7..f262d8b78 100644 --- a/TODO.md +++ b/TODO.md @@ -4,6 +4,14 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Current +- Video tab: + - remote vae + - tiny vae + - lora + - accelerators: teacache, pab, fastercache, paraattention, perflow + - modernui tab +- Detailer daemon: https://github.com/muerrilla/sd-webui-detail-daemon/blob/main/scripts/detail_daemon.py + ## Future Candidates - Redesign postprocessing diff --git a/modules/onnx_impl/execution_providers.py b/modules/onnx_impl/execution_providers.py index e38199d0f..25e0056b3 100644 --- a/modules/onnx_impl/execution_providers.py +++ b/modules/onnx_impl/execution_providers.py @@ -107,10 +107,14 @@ def install_execution_provider(ep: ExecutionProvider): elif ep == ExecutionProvider.OpenVINO: packages.append("openvino") packages.append("onnxruntime-openvino") + log.info(f'ONNX install: {packages}') for package in packages: res += install(package) res += '
' res += 'Server restart required' log.info("Server restart required") - importlib.reload(ort) + try: + importlib.reload(ort) + except Exception: + pass return res diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index a5f6b87c9..366a62995 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -71,8 +71,10 @@ def process_base(p: processing.StableDiffusionProcessing): eta=shared.opts.scheduler_eta, guidance_scale=p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, + true_cfg_scale=p.diffusers_guidance_rescale, denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None, denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None, + num_frames=getattr(p, 'frames', None), output_type='latent', clip_skip=p.clip_skip, desc='Base', @@ -364,7 +366,7 @@ def process_decode(p: processing.StableDiffusionProcessing, output): else: width = getattr(p, 'width', 0) height = getattr(p, 'height', 0) - frames = p.task_args.get('num_frames', None) + frames = p.task_args.get('num_frames', None) or getattr(p, 'frames', None) if isinstance(output.images, list): results = [] for i in range(len(output.images)): diff --git a/modules/ui_caption.py b/modules/ui_caption.py index 474427e9d..ef8d2d6e2 100644 --- a/modules/ui_caption.py +++ b/modules/ui_caption.py @@ -45,9 +45,9 @@ def create_ui(): vlm_model = gr.Dropdown(list(vqa.vlm_models), value=list(vqa.vlm_models)[0], label='VLM Model', elem_id='vlm_model') with gr.Accordion(label='Advanced options', open=False, visible=True): with gr.Row(): - vlm_max_tokens = gr.Slider(label='Max tokens', value=shared.opts.interrogate_vlm_max_length, minimum=16, maximum=4096, step=1, elem_id='vlm_max_tokens') - vlm_num_beams = gr.Slider(label='Num beams', value=shared.opts.interrogate_vlm_num_beams, minimum=1, maximum=16, step=1, elem_id='vlm_num_beams') - vlm_temperature = gr.Slider(label='Temperature', value=shared.opts.interrogate_vlm_temperature, minimum=0.1, maximum=1.0, step=0.01, elem_id='vlm_temperature') + vlm_max_tokens = gr.Slider(label='VLM max tokens', value=shared.opts.interrogate_vlm_max_length, minimum=16, maximum=4096, step=1, elem_id='vlm_max_tokens') + vlm_num_beams = gr.Slider(label='VLM num beams', value=shared.opts.interrogate_vlm_num_beams, minimum=1, maximum=16, step=1, elem_id='vlm_num_beams') + vlm_temperature = gr.Slider(label='VLM Temperature', value=shared.opts.interrogate_vlm_temperature, minimum=0.1, maximum=1.0, step=0.01, elem_id='vlm_temperature') with gr.Row(): vlm_top_k = gr.Slider(label='Top-K', value=shared.opts.interrogate_vlm_top_k, minimum=0, maximum=99, step=1, elem_id='vlm_top_k') vlm_top_p = gr.Slider(label='Top-P', value=shared.opts.interrogate_vlm_top_p, minimum=0.0, maximum=1.0, step=0.01, elem_id='vlm_top_p') @@ -90,7 +90,7 @@ def create_ui(): clip_max_flavors = gr.Slider(label='Max flavors', value=shared.opts.interrogate_clip_max_flavors, minimum=1, maximum=64, step=1, elem_id='clip_max_flavors') clip_flavor_count = gr.Slider(label='Intermediates', value=shared.opts.interrogate_clip_flavor_count, minimum=256, maximum=4096, step=8, elem_id='clip_flavor_intermediate_count') with gr.Row(): - clip_num_beams = gr.Slider(label='Num beams', value=shared.opts.interrogate_clip_num_beams, minimum=1, maximum=16, step=1, elem_id='clip_num_beams') + clip_num_beams = gr.Slider(label='CLiP num beams', value=shared.opts.interrogate_clip_num_beams, minimum=1, maximum=16, step=1, elem_id='clip_num_beams') clip_min_length.change(fn=update_clip_params, inputs=[clip_min_length, clip_max_length, clip_chunk_size, clip_min_flavors, clip_max_flavors, clip_flavor_count, clip_num_beams], outputs=[]) clip_max_length.change(fn=update_clip_params, inputs=[clip_min_length, clip_max_length, clip_chunk_size, clip_min_flavors, clip_max_flavors, clip_flavor_count, clip_num_beams], outputs=[]) clip_chunk_size.change(fn=update_clip_params, inputs=[clip_min_length, clip_max_length, clip_chunk_size, clip_min_flavors, clip_max_flavors, clip_flavor_count, clip_num_beams], outputs=[]) diff --git a/modules/ui_video.py b/modules/ui_video.py index 97b9d2ec7..2c5c324a5 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -1,30 +1,23 @@ -# TODO hunyuanvideo: prompt_template, lora -# TODO hunyuanvideo: teacache, pab, fastercache, paraattention, perflow -# TODO modernui video tab - import gradio as gr from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste from modules.ui_components import ToolButton -from modules.video_models import hunyuan +from modules.video_models import models_def, video_utils, hunyuan, ltx def engine_change(engine): - found = [model.name for model in hunyuan.models.get(engine, [])] + found = [model.name for model in models_def.models.get(engine, [])] return gr.update(choices=found, value=found[0] if len(found) > 0 else None) def model_change(engine, model): - found = [model.name for model in hunyuan.models.get(engine, [])] - selected = [m for m in hunyuan.models[engine] if m.name == model][0] if len(found) > 0 else None + found = [model.name for model in models_def.models.get(engine, [])] + selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if selected: if 'None' in selected.name: sd_models.unload_model_weights() msg = 'Video model unloaded' - elif 'Hunyuan' in selected.name: - msg = hunyuan.load(selected) - elif model != 'None': - msg = f'Video model not found: engine={engine} model={model}' - shared.log.error(msg) + else: + msg = video_utils.load_model(selected) else: sd_models.unload_model_weights() msg = 'Video model unloaded' @@ -33,10 +26,13 @@ def model_change(engine, model): def run_video(*args): engine, model = args[2], args[3] - found = [model.name for model in hunyuan.models.get(engine, [])] - selected = [m for m in hunyuan.models[engine] if m.name == model][0] if len(found) > 0 else None + found = [model.name for model in models_def.models.get(engine, [])] + selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if selected and 'Hunyuan' in selected.name: return hunyuan.generate(*args) + elif selected and 'LTX' in selected.name: + pass + # return ltx.generate(*args) shared.log.error(f'Video model not found: args={args}') return [], None, '', '', f'Video model not found: engine={engine} model={model}' @@ -57,7 +53,7 @@ def create_ui(): with gr.Column(variant='compact', elem_id="video_settings", scale=1): with gr.Row(): - engine = gr.Dropdown(label='Engine', choices=list(hunyuan.models), value='None', elem_id="video_engine") + engine = gr.Dropdown(label='Engine', choices=list(models_def.models), value='None', elem_id="video_engine") model = gr.Dropdown(label='Model', choices=[''], value=None, elem_id="video_model") with gr.Row(): width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) @@ -69,6 +65,7 @@ def create_ui(): steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") with gr.Row(): sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") + dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift") with gr.Row(): guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") @@ -86,7 +83,7 @@ def create_ui(): override_settings = ui_common.create_override_inputs('video') # output panel with gallery and video tabs - with gr.Column(elem_id='video-output-column', scale=3) as _column_output: + with gr.Column(elem_id='video-output-column', scale=2) as _column_output: with gr.Tabs(elem_classes=['video-output-tabs'], elem_id='video-output-tabs'): with gr.Tab('Frames', id='out-gallery'): gallery, gen_info, html_info, _html_info_formatted, html_log = ui_common.create_output_panel("video", prompt=prompt, preview=False, transfer=False, scale=2) @@ -124,7 +121,7 @@ def create_ui(): width, height, frames, steps, sampler_index, - sampler_shift, + sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, diff --git a/modules/video_models/hunyuan.py b/modules/video_models/hunyuan.py index 2b83e6b63..e42139189 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/hunyuan.py @@ -1,187 +1,22 @@ -from dataclasses import dataclass import os import time -import torch -import transformers -import diffusers -from modules import shared, errors, sd_models, sd_checkpoint, sd_samplers, processing, model_quant, devices, images, timer, ui_common +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils -@dataclass -class Model(): - name: str - repo: str - dit: str - subfolder: str - -models = { - 'None': [], - 'Hunyuan Video': [ - Model(name='None', repo=None, dit=None, subfolder=None), - Model(name='Hunyuan Video T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='hunyuanvideo-community/HunyuanVideo', subfolder='transformer'), - Model(name='Hunyuan Video I2V', repo='hunyuanvideo-community/HunyuanVideo-I2V', dit='hunyuanvideo-community/HunyuanVideo-I2V', subfolder='transformer'), # https://github.com/huggingface/diffusers/pull/10983 - Model(name='SkyReels Hunyuan T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='Skywork/SkyReels-V1-Hunyuan-T2V', subfolder=None), # https://github.com/huggingface/diffusers/pull/10837 - Model(name='SkyReels Hunyuan I2V', repo='hunyuanvideo-community/HunyuanVideo', dit='Skywork/SkyReels-V1-Hunyuan-I2V', subfolder=None), - Model(name='Fast Hunyuan T2V', repo='hunyuanvideo-community/HunyuanVideo', dit='FastVideo/FastHunyuan-diffusers', subfolder='transformer'), # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 - ] -} debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None -loaded_model = None -prompt_template = { - "template": ( - "<|start_header_id|>system<|end_header_id|>" - "\nDescribe the video by detailing the following aspects: \n" - "1. The main content and theme of the video.\n" - "2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects.\n" - "3. Actions, events, behaviors, temporal relationships, and physical movement changes of the objects.\n" - "4. Background environment, light, style and atmosphere.\n" - "5. Camera angles, movements, and transitions used in the video.\n" - "<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>" - ), - "crop_start": 95, -} - - -def hijack_decode(*args, **kwargs): - t0 = time.time() - vae: diffusers.AutoencoderKLHunyuanVideo = shared.sd_model.vae - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - debug(f'Video: vae={vae.__class__.__name__} tile={vae.tile_sample_min_width}:{vae.tile_sample_min_height}:{vae.tile_sample_min_num_frames} stride={vae.tile_sample_stride_width}:{vae.tile_sample_stride_height}:{vae.tile_sample_stride_num_frames} time={t1-t0:.2f}') - return res - - -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - -def get_quant(args): - if args is not None and "quantization_config" in args: - return args['quantization_config'].__class__.__name__ - return None - - -def load(selected): - if selected is None: - return - global loaded_model # pylint: disable=global-statement - if loaded_model == selected.name: - return - sd_models.unload_model_weights() - t0 = time.time() - - quant_args = model_quant.create_config(module='Model') - cls = diffusers.HunyuanVideoTransformer3DModel - try: - debug(f'Video load: module=transofrmer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} quant={get_quant(quant_args)}') - transformer = cls.from_pretrained( - pretrained_model_name_or_path=selected.dit, - subfolder=selected.subfolder, - torch_dtype=devices.dtype, - cache_dir=shared.opts.hfcache_dir, - **quant_args - ) - except Exception as e: - shared.log.error(f'video load: module=transformer repo="{selected.dit}" subfolder="{selected.subfolder}" cls={cls.__name__} {e}') - - quant_args = model_quant.create_config(module='Text Encoder') - if 'I2V' in selected.repo: - cls = transformers.LlavaForConditionalGeneration - else: - cls = transformers.LlamaModel - try: - debug(f'Video load: module=te repo="{selected.repo}" cls={cls.__name__} quant={get_quant(quant_args)}') - text_encoder = cls.from_pretrained( - pretrained_model_name_or_path=selected.repo, - subfolder="text_encoder", - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - # torch_dtype='auto', # special case as text and vision nested models have different dtypes - # attn_implementation="flash_attention_2", # testing different attention types - **quant_args - ) - except Exception as e: - shared.log.error(f'video load: module=te repo="{selected.repo}" cls={cls.__name__} {e}') - - cls = transformers.CLIPTextModel - try: - debug(f'Video load: module=clip repo="{selected.repo}" cls={cls.__name__} quant=None') - text_encoder_2 = transformers.CLIPTextModel.from_pretrained( - pretrained_model_name_or_path=selected.repo, - subfolder="text_encoder_2", - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - ) - except Exception as e: - shared.log.error(f'video load: module=clip repo="{selected.repo}" cls={cls.__name__} {e}') - - cls = diffusers.AutoencoderKLHunyuanVideo - try: - debug(f'Video load: module=vae repo="{selected.repo}" cls={cls.__name__} quant=None') - vae = diffusers.AutoencoderKLHunyuanVideo.from_pretrained( - pretrained_model_name_or_path=selected.repo, - subfolder="vae", - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - ) - except Exception as e: - shared.log.error(f'video load: module=vae repo="{selected.repo}" cls={cls.__name__} {e}') - - if selected.name == 'Hunyuan Video I2V': - cls = diffusers.HunyuanVideoImageToVideoPipeline - elif selected.name == 'SkyReels Hunyuan I2V': - cls = diffusers.HunyuanSkyreelsImageToVideoPipeline - else: - cls = diffusers.HunyuanVideoPipeline - try: - debug(f'Video load: module=pipe repo="{selected.repo}" cls={cls.__name__} quant=None') - shared.sd_model = cls.from_pretrained( - pretrained_model_name_or_path=selected.repo, - transformer=transformer, - text_encoder=text_encoder, - text_encoder_2=text_encoder_2, - vae=vae, - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - ) - except Exception as e: - shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={cls.__name__} {e}') - - t1 = time.time() - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) - shared.sd_model.sd_model_hash = None - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.vae.decode = hijack_decode - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.encode_prompt = hijack_encode_prompt - shared.sd_model.vae.enable_slicing() - loaded_model = selected.name - msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' - shared.log.info(msg) - return msg def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': - shared.log.error('Video: model not selected') - return [], None, '', '', 'Video model not selected' + return video_utils.queue_err('model not selected') if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - found = [model.name for model in models.get(engine, [])] - selected = [m for m in models[engine] if m.name == model][0] if len(found) > 0 else None - load(selected) + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + video_utils.load_model(selected) if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - shared.log.error('Video: model not loaded') - return [], None, '', '', 'Video model not loaded' + return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') p = processing.StableDiffusionProcessingVideo( @@ -202,27 +37,24 @@ def generate(*args, **kwargs): override_settings=override_settings, ) p.scripts = None - p.script_args = args + p.script_args = None p.state = ui_state p.do_not_save_grid = True p.do_not_save_samples = not save_frames if 'I2V' in model: if init_image is None: - shared.log.error('Video: init image not set') - return [], None, '', '', 'Error: init image not set' - p.task_args['image'] = init_image - # from PIL import Image - # p.task_args['image'] = init_image.resize((336, 336), Image.Resampling.LANCZOS) + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + # cleanup memory shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) devices.torch_gc(force=True) # handle sampler and seed - if p.sampler_name != 'Default': - shared.sd_model.scheduler = sd_samplers.create_sampler(p.sampler_name, shared.sd_model) - p.sampler_name = 'Default' # avoid double creation - if hasattr(shared.sd_model.scheduler, '_shift') and sampler_shift > 0: - shared.sd_model.scheduler._shift = sampler_shift # pylint: disable=protected-access + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_shift'] = sampler_shift # handle vae if vae_tile_frames > p.frames: @@ -237,39 +69,32 @@ def generate(*args, **kwargs): processing.fix_seed(p) p.prompt = shared.prompt_styles.apply_styles_to_prompt(prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(negative, p.styles) - p.task_args['width'] = p.width - p.task_args['height'] = p.height - p.task_args['num_inference_steps'] = p.steps - p.task_args['num_frames'] = p.frames - p.task_args['generator'] = torch.manual_seed(p.seed) - p.task_args['guidance_scale'] = p.cfg_scale - p.task_args['true_cfg_scale'] = p.diffusers_guidance_rescale - # p.task_args['prompt_template'] = prompt_template # t2v and i2v have different templates - p.task_args['output_type'] = 'pil' p.task_args['prompt'] = p.prompt p.task_args['negative_prompt'] = p.negative_prompt + p.task_args['output_type'] = 'pil' p.ops.append('video') debug(f'Video: task_args={p.task_args}') # run processing shared.state.disable_preview = True shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None t0 = time.time() try: processed = processing.process_images(p) except Exception as e: - shared.log.error(f'Video: exception={e}') + err = str(e) errors.display(e, 'video') - processed = None - shared.state.disable_preview = False - return [], None, '', '', str(e) t1 = time.time() shared.state.disable_preview = False - + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift p.close() + if err: + return video_utils.queue_err(err) if processed is None or len(processed.images) == 0: - return [], None, '', '', 'Video: processing failed' - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') if video_type != 'None': video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) else: diff --git a/modules/video_models/ltx.py b/modules/video_models/ltx.py new file mode 100644 index 000000000..e69de29bb diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py new file mode 100644 index 000000000..aba1b76b9 --- /dev/null +++ b/modules/video_models/models_def.py @@ -0,0 +1,67 @@ +from dataclasses import dataclass +import diffusers +import transformers + + +@dataclass +class Model(): + name: str + repo: str = None + repo_cls: classmethod = None + dit: str = None + dit_cls: classmethod = None + dit_folder: str = 'transformer' + te: str = None + te_cls: classmethod = None + te_folder: str = 'text_encoder' + te_hijack: bool = True + vae_hijack: bool = True + + +models = { + 'None': [], + 'Hunyuan Video': [ + Model(name='None'), + Model(name='Hunyuan Video T2V', + repo='hunyuanvideo-community/HunyuanVideo', + repo_cls=diffusers.HunyuanVideoPipeline, + te_cls=transformers.LlamaModel, + dit_cls=diffusers.HunyuanVideoTransformer3DModel), + Model(name='Hunyuan Video I2V', # https://github.com/huggingface/diffusers/pull/10983 + repo='hunyuanvideo-community/HunyuanVideo-I2V', + repo_cls=diffusers.HunyuanVideoImageToVideoPipeline, + te_cls=transformers.LlavaForConditionalGeneration, + dit_cls=diffusers.HunyuanVideoTransformer3DModel), + Model(name='SkyReels Hunyuan T2V', # https://github.com/huggingface/diffusers/pull/10837 + repo='hunyuanvideo-community/HunyuanVideo', + repo_cls=diffusers.HunyuanVideoPipeline, + te_cls=transformers.LlamaModel, + dit='Skywork/SkyReels-V1-Hunyuan-T2V', + dit_folder=None, + dit_cls=diffusers.HunyuanVideoTransformer3DModel), + Model(name='SkyReels Hunyuan I2V', # https://github.com/huggingface/diffusers/pull/10837 + repo='hunyuanvideo-community/HunyuanVideo', + te_cls=transformers.LlamaModel, + dit='Skywork/SkyReels-V1-Hunyuan-I2V', + dit_folder=None, + dit_cls=diffusers.HunyuanVideoTransformer3DModel), + Model(name='Fast Hunyuan T2V', # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 + repo='hunyuanvideo-community/HunyuanVideo', + repo_cls=diffusers.HunyuanVideoPipeline, + te_cls=transformers.LlamaModel, + dit='FastVideo/FastHunyuan-diffusers', + dit_cls=diffusers.HunyuanVideoTransformer3DModel), + ], +} + +""" +'LTX Video': [ + Model(name='None'), + Model(name='LTXVideo 0.9.0 T2V', repo='a-r-r-o-w/LTX-Video-diffusers', subfolder='transformer'), + Model(name='LTXVideo 0.9.1 T2V', repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', subfolder='transformer'), + Model(name='LTXVideo 0.9.5 T2V', repo='Lightricks/LTX-Video-0.9.5'), # https://github.com/huggingface/diffusers/pull/10968 + Model(name='LTXVideo 0.9.0 I2V', repo='a-r-r-o-w/LTX-Video-diffusers', subfolder='transformer'), + Model(name='LTXVideo 0.9.1 I2V', repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', subfolder='transformer'), + Model(name='LTXVideo 0.9.5 I2V', repo='Lightricks/LTX-Video-0.9.5', subfolder='transformer'), +]," +""" diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py new file mode 100644 index 000000000..3313cab69 --- /dev/null +++ b/modules/video_models/video_utils.py @@ -0,0 +1,109 @@ +import os +import time +from modules import shared, timer, sd_models, sd_checkpoint, model_quant, devices +from modules.video_models import models_def + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def queue_err(msg): + shared.log.error(f'Video: {msg}') + return [], None, '', '', f'Error: {msg}' + + +def get_quant(args): + if args is not None and "quantization_config" in args: + return args['quantization_config'].__class__.__name__ + return None + + +def hijack_vae_decode(*args, **kwargs): + t0 = time.time() + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) + res = shared.sd_model.vae.orig_decode(*args, **kwargs) + t1 = time.time() + timer.process.add('vae', t1-t0) + debug(f'Video decode: vae={shared.sd_model.vae.__class__.__name__} time={t1-t0:.2f}') + return res + + +def hijack_encode_prompt(*args, **kwargs): + t0 = time.time() + res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + t1 = time.time() + timer.process.add('te', t1-t0) + debug(f'Video encode: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + return res + +loaded_model = None + + +def load_model(selected: models_def.Model): + if selected is None: + return + global loaded_model # pylint: disable=global-statement + if loaded_model == selected.name: + return + sd_models.unload_model_weights() + t0 = time.time() + + # text encoder + try: + quant_args = model_quant.create_config(module='Text Encoder') + debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={get_quant(quant_args)}') + text_encoder = selected.te_cls.from_pretrained( + pretrained_model_name_or_path=selected.te or selected.repo, + subfolder=selected.te_folder, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}') + text_encoder = None + + # transformer + try: + quant_args = model_quant.create_config(module='Model') + debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={get_quant(quant_args)}') + transformer = selected.dit_cls.from_pretrained( + pretrained_model_name_or_path=selected.dit or selected.repo, + subfolder=selected.dit_folder, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}') + transformer = None + + # model + try: + debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}') + shared.sd_model = selected.repo_cls.from_pretrained( + pretrained_model_name_or_path=selected.repo, + transformer=transformer, + text_encoder=text_encoder, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + ) + except Exception as e: + shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}') + + t1 = time.time() + sd_models.set_diffuser_options(shared.sd_model) + shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) + shared.sd_model.sd_model_hash = None + if selected.vae_hijack: + shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode + shared.sd_model.vae.decode = hijack_vae_decode + if selected.te_hijack: + shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt + shared.sd_model.encode_prompt = hijack_encode_prompt + shared.sd_model.vae.enable_slicing() + loaded_model = selected.name + msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' + shared.log.info(msg) + return msg diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py index 135256ef2..a1933673e 100644 --- a/scripts/ltxvideo.py +++ b/scripts/ltxvideo.py @@ -11,6 +11,7 @@ from modules.teacache.teacache_ltx import teacache_forward repos = { '0.9.0': 'a-r-r-o-w/LTX-Video-diffusers', '0.9.1': 'a-r-r-o-w/LTX-Video-0.9.1-diffusers', + '0.9.5': 'Lightricks/LTX-Video-0.9.5', 'custom': None, } @@ -31,7 +32,6 @@ def load_quants(kwargs, repo_id): def hijack_decode(*args, **kwargs): t0 = time.time() - # vae: diffusers.AutoencoderKLHunyuanVideo = shared.sd_model.vae shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) res = shared.sd_model.vae.orig_decode(*args, **kwargs) t1 = time.time() From 79a5391e47a81c31372b7a12e817a1011d18fcdf Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 18:14:23 -0400 Subject: [PATCH 033/122] video unified component loader Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + installer.py | 2 +- modules/ui_video.py | 33 ++++-- modules/video_models/ltx.py | 0 modules/video_models/models_def.py | 108 ++++++++++++++++-- modules/video_models/run_allegro.py | 91 +++++++++++++++ modules/video_models/run_cog.py | 91 +++++++++++++++ .../{hunyuan.py => run_hunyuan.py} | 39 +++---- modules/video_models/run_ltx.py | 91 +++++++++++++++ modules/video_models/run_mochi.py | 91 +++++++++++++++ modules/video_models/video_utils.py | 31 ++++- 11 files changed, 531 insertions(+), 47 deletions(-) delete mode 100644 modules/video_models/ltx.py create mode 100644 modules/video_models/run_allegro.py create mode 100644 modules/video_models/run_cog.py rename modules/video_models/{hunyuan.py => run_hunyuan.py} (77%) create mode 100644 modules/video_models/run_ltx.py create mode 100644 modules/video_models/run_mochi.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 67745eacf..4a6ba094e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,7 @@ - Gemma3 requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - Remote VAE encode for SD15 and Flux.1: - HunyuanVideo-I2V: + - LTXVideo condition input ### Highlights for 2025-03-17 diff --git a/installer.py b/installer.py index fb54f25e4..0411ef496 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git: return - sha = '82188cef0487837b8c70fc3f36ea63c05c85f341' # diffusers commit hash + sha = '56f740051dae2d410677292a5c9e5b66e60f87dc' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/ui_video.py b/modules/ui_video.py index 2c5c324a5..05a3bfc09 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -1,7 +1,7 @@ import gradio as gr from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste from modules.ui_components import ToolButton -from modules.video_models import models_def, video_utils, hunyuan, ltx +from modules.video_models import models_def, video_utils def engine_change(engine): @@ -12,6 +12,10 @@ def engine_change(engine): def model_change(engine, model): found = [model.name for model in models_def.models.get(engine, [])] selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + yield ['Video model loading', + gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False), + video_utils.get_url(selected.url if selected else None), + ] if selected: if 'None' in selected.name: sd_models.unload_model_weights() @@ -21,7 +25,10 @@ def model_change(engine, model): else: sd_models.unload_model_weights() msg = 'Video model unloaded' - return [msg, gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False)] + return [msg, + gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False), + video_utils.get_url(selected.url if selected else None), + ] def run_video(*args): @@ -29,10 +36,20 @@ def run_video(*args): found = [model.name for model in models_def.models.get(engine, [])] selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if selected and 'Hunyuan' in selected.name: - return hunyuan.generate(*args) + from modules.video_models import run_hunyuan + return run_hunyuan.generate(*args) elif selected and 'LTX' in selected.name: - pass - # return ltx.generate(*args) + from modules.video_models import run_ltx + return run_ltx.generate(*args) + elif selected and 'Mochi' in selected.name: + from modules.video_models import run_mochi + return run_mochi.generate(*args) + elif selected and 'Cog' in selected.name: + from modules.video_models import run_cog + return run_cog.generate(*args) + elif selected and 'Allegro' in selected.name: + from modules.video_models import run_allegro + return run_allegro.generate(*args) shared.log.error(f'Video model not found: args={args}') return [], None, '', '', f'Video model not found: engine={engine} model={model}' @@ -55,6 +72,8 @@ def create_ui(): with gr.Row(): engine = gr.Dropdown(label='Engine', choices=list(models_def.models), value='None', elem_id="video_engine") model = gr.Dropdown(label='Model', choices=[''], value=None, elem_id="video_model") + with gr.Row(): + url = gr.HTML(label='Model URL', elem_id='video_model_url', value='') with gr.Row(): width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) with gr.Row(): @@ -65,7 +84,7 @@ def create_ui(): steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") with gr.Row(): sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") - dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift") + dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) with gr.Row(): guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") @@ -95,7 +114,7 @@ def create_ui(): random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) # handle engine and model change engine.change(fn=engine_change, inputs=[engine], outputs=[model]) - model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group]) + model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group, url]) # setup extra networks ui_extra_networks.setup_ui(extra_networks_ui, gallery) diff --git a/modules/video_models/ltx.py b/modules/video_models/ltx.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index aba1b76b9..82e6f108a 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -6,6 +6,7 @@ import transformers @dataclass class Model(): name: str + url: str = '' repo: str = None repo_cls: classmethod = None dit: str = None @@ -23,16 +24,19 @@ models = { 'Hunyuan Video': [ Model(name='None'), Model(name='Hunyuan Video T2V', + url='https://huggingface.co/tencent/HunyuanVideo', repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='Hunyuan Video I2V', # https://github.com/huggingface/diffusers/pull/10983 + url='https://huggingface.co/tencent/HunyuanVideo-I2V', repo='hunyuanvideo-community/HunyuanVideo-I2V', repo_cls=diffusers.HunyuanVideoImageToVideoPipeline, te_cls=transformers.LlavaForConditionalGeneration, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='SkyReels Hunyuan T2V', # https://github.com/huggingface/diffusers/pull/10837 + url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-T2V', repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, @@ -40,28 +44,108 @@ models = { dit_folder=None, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='SkyReels Hunyuan I2V', # https://github.com/huggingface/diffusers/pull/10837 + url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-I2V', repo='hunyuanvideo-community/HunyuanVideo', te_cls=transformers.LlamaModel, dit='Skywork/SkyReels-V1-Hunyuan-I2V', dit_folder=None, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='Fast Hunyuan T2V', # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 + url='https://huggingface.co/FastVideo/FastHunyuan', repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, dit='FastVideo/FastHunyuan-diffusers', dit_cls=diffusers.HunyuanVideoTransformer3DModel), ], + 'LTX Video': [ + Model(name='None'), + Model(name='LTXVideo 0.9.5 T2V', # https://github.com/huggingface/diffusers/pull/10968 + url='https://huggingface.co/Lightricks/LTX-Video-0.9.5', + repo='YiYiXu/ltx-95', + repo_cls=diffusers.LTXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + Model(name='LTXVideo 0.9.5 I2V', + url='https://huggingface.co/Lightricks/LTX-Video-0.9.5', + repo='YiYiXu/ltx-95', + repo_cls=diffusers.LTXConditionPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + Model(name='LTXVideo 0.9.1 T2V', + url='https://huggingface.co/a-r-r-o-w/LTX-Video-0.9.1-diffusers', + repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', + repo_cls=diffusers.LTXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + Model(name='LTXVideo 0.9.1 I2V', + url='https://huggingface.co/a-r-r-o-w/LTX-Video-0.9.1-diffusers', + repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', + repo_cls=diffusers.LTXImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + Model(name='LTXVideo 0.9.0 T2V', + url='https://huggingface.co/a-r-r-o-w/LTX-Video-diffusers', + repo='a-r-r-o-w/LTX-Video-diffusers', + repo_cls=diffusers.LTXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + Model(name='LTXVideo 0.9.0 I2V', + url='https://huggingface.co/a-r-r-o-w/LTX-Video-diffusers', + repo='a-r-r-o-w/LTX-Video-diffusers', + repo_cls=diffusers.LTXImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LTXVideoTransformer3DModel), + ], + 'Mochi Video': [ + Model(name='None'), + Model(name='Mochi 1 T2V', + url='https://huggingface.co/genmo/mochi-1-preview', + repo='genmo/mochi-1-preview', + repo_cls=diffusers.MochiPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.MochiTransformer3DModel), + ], + 'Allegro Video': [ + Model(name='None'), + Model(name='Allegro T2V', + url='https://huggingface.co/rhymes-ai/Allegro', + repo='rhymes-ai/Allegro', + repo_cls=diffusers.AllegroPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.AllegroTransformer3DModel), + ], + 'Cog Video': [ + Model(name='None'), + Model(name='CogVideoX 1.0 2B T2V', + url='https://huggingface.co/THUDM/CogVideoX-2b', + repo='THUDM/CogVideoX-2b', + repo_cls=diffusers.CogVideoXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.CogVideoXTransformer3DModel), + Model(name='CogVideoX 1.0 5B T2V', + url='https://huggingface.co/THUDM/CogVideoX-5b', + repo='THUDM/THUDM/CogVideoX-5b', + repo_cls=diffusers.CogVideoXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.CogVideoXTransformer3DModel), + Model(name='CogVideoX 1.0 5B I2V', + url='https://huggingface.co/THUDM/CogVideoX-5b-I2V', + repo='THUDM/CogVideoX-5b-I2V', + repo_cls=diffusers.CogVideoXImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.CogVideoXTransformer3DModel), + Model(name='CogVideoX 1.5 5B T2V', + url='https://huggingface.co/THUDM/THUDM/CogVideoX1.5-5B', + repo='THUDM/CogVideoX1.5-5B', + repo_cls=diffusers.CogVideoXPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.CogVideoXTransformer3DModel), + Model(name='CogVideoX 1.5 5B I2V', + url='https://huggingface.co/THUDM/CogVideoX1.5-5B-I2V', + repo='THUDM/CogVideoX1.5-5B-I2V', + repo_cls=diffusers.CogVideoXImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.CogVideoXTransformer3DModel), + ], } - -""" -'LTX Video': [ - Model(name='None'), - Model(name='LTXVideo 0.9.0 T2V', repo='a-r-r-o-w/LTX-Video-diffusers', subfolder='transformer'), - Model(name='LTXVideo 0.9.1 T2V', repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', subfolder='transformer'), - Model(name='LTXVideo 0.9.5 T2V', repo='Lightricks/LTX-Video-0.9.5'), # https://github.com/huggingface/diffusers/pull/10968 - Model(name='LTXVideo 0.9.0 I2V', repo='a-r-r-o-w/LTX-Video-diffusers', subfolder='transformer'), - Model(name='LTXVideo 0.9.1 I2V', repo='a-r-r-o-w/LTX-Video-0.9.1-diffusers', subfolder='transformer'), - Model(name='LTXVideo 0.9.5 I2V', repo='Lightricks/LTX-Video-0.9.5', subfolder='transformer'), -]," -""" diff --git a/modules/video_models/run_allegro.py b/modules/video_models/run_allegro.py new file mode 100644 index 000000000..be1b0da8d --- /dev/null +++ b/modules/video_models/run_allegro.py @@ -0,0 +1,91 @@ +import os +import time +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def generate(*args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: + video_utils.load_model(selected) + if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: + return video_utils.queue_err('model not loaded') + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=8 * int(width // 8), + height=8 * int(height // 8), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + p.scripts = None + p.script_args = None + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + + # cleanup memory + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # set args + processing.fix_seed(p) + video_utils.set_vae_params(p.frames, vae_tile_frames) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'pil' + p.ops.append('video') + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_shift'] = sampler_shift + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None + t0 = time.time() + try: + processed = processing.process_images(p) + except Exception as e: + err = str(e) + errors.display(e, 'video') + t1 = time.time() + shared.state.disable_preview = False + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + p.close() + + # done + if err: + return video_utils.queue_err(err) + if processed is None or len(processed.images) == 0: + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + generation_info_js = processed.js() if processed is not None else '' + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_cog.py b/modules/video_models/run_cog.py new file mode 100644 index 000000000..5fbe1d320 --- /dev/null +++ b/modules/video_models/run_cog.py @@ -0,0 +1,91 @@ +import os +import time +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def generate(*args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: + video_utils.load_model(selected) + if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: + return video_utils.queue_err('model not loaded') + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=8 * int(width // 8), + height=8 * int(height // 8), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + p.scripts = None + p.script_args = None + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + + # cleanup memory + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # set args + processing.fix_seed(p) + video_utils.set_vae_params(p.frames, vae_tile_frames) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'pil' + p.ops.append('video') + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_shift'] = sampler_shift + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None + t0 = time.time() + try: + processed = processing.process_images(p) + except Exception as e: + err = str(e) + errors.display(e, 'video') + t1 = time.time() + shared.state.disable_preview = False + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + p.close() + + # done + if err: + return video_utils.queue_err(err) + if processed is None or len(processed.images) == 0: + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + generation_info_js = processed.js() if processed is not None else '' + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/hunyuan.py b/modules/video_models/run_hunyuan.py similarity index 77% rename from modules/video_models/hunyuan.py rename to modules/video_models/run_hunyuan.py index e42139189..aa076f935 100644 --- a/modules/video_models/hunyuan.py +++ b/modules/video_models/run_hunyuan.py @@ -11,9 +11,9 @@ def generate(*args, **kwargs): task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None video_utils.load_model(selected) if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') @@ -21,6 +21,8 @@ def generate(*args, **kwargs): p = processing.StableDiffusionProcessingVideo( sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, styles=styles, seed=int(seed), sampler_name = processing.get_sampler_name(sampler_index), @@ -50,29 +52,16 @@ def generate(*args, **kwargs): shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) devices.torch_gc(force=True) - # handle sampler and seed + # set args + processing.fix_seed(p) + video_utils.set_vae_params(p.frames, vae_tile_frames) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'pil' + p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift shared.opts.data['schedulers_shift'] = sampler_shift - - # handle vae - if vae_tile_frames > p.frames: - shared.sd_model.vae.tile_sample_min_num_frames = vae_tile_frames - shared.sd_model.vae.use_framewise_decoding = True - shared.sd_model.vae.enable_tiling() - else: - shared.sd_model.vae.use_framewise_decoding = False - shared.sd_model.vae.disable_tiling() - - # set args - processing.fix_seed(p) - p.prompt = shared.prompt_styles.apply_styles_to_prompt(prompt, p.styles) - p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(negative, p.styles) - p.task_args['prompt'] = p.prompt - p.task_args['negative_prompt'] = p.negative_prompt - p.task_args['output_type'] = 'pil' - p.ops.append('video') debug(f'Video: task_args={p.task_args}') # run processing @@ -90,15 +79,13 @@ def generate(*args, **kwargs): shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift shared.opts.data['schedulers_shift'] = orig_sampler_shift p.close() + + # done if err: return video_utils.queue_err(err) if processed is None or len(processed.images) == 0: return video_utils.queue_err('processing failed') shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - else: - video_file = None - + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) generation_info_js = processed.js() if processed is not None else '' return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_ltx.py b/modules/video_models/run_ltx.py new file mode 100644 index 000000000..ba01e45a4 --- /dev/null +++ b/modules/video_models/run_ltx.py @@ -0,0 +1,91 @@ +import os +import time +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def generate(*args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: + video_utils.load_model(selected) + if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: + return video_utils.queue_err('model not loaded') + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=32 * int(width // 32), + height=32 * int(height // 32), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + p.scripts = None + p.script_args = None + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + + # cleanup memory + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # set args + processing.fix_seed(p) + video_utils.set_vae_params(p.frames, vae_tile_frames) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'pil' + p.ops.append('video') + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_shift'] = sampler_shift + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None + t0 = time.time() + try: + processed = processing.process_images(p) + except Exception as e: + err = str(e) + errors.display(e, 'video') + t1 = time.time() + shared.state.disable_preview = False + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + p.close() + + # done + if err: + return video_utils.queue_err(err) + if processed is None or len(processed.images) == 0: + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + generation_info_js = processed.js() if processed is not None else '' + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_mochi.py b/modules/video_models/run_mochi.py new file mode 100644 index 000000000..fff3a6d04 --- /dev/null +++ b/modules/video_models/run_mochi.py @@ -0,0 +1,91 @@ +import os +import time +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def generate(*args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: + video_utils.load_model(selected) + if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: + return video_utils.queue_err('model not loaded') + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=8 * int(width // 8), + height=8 * int(height // 8), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + p.scripts = None + p.script_args = None + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + + # cleanup memory + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # set args + processing.fix_seed(p) + video_utils.set_vae_params(p.frames, vae_tile_frames) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'pil' + p.ops.append('video') + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_shift'] = sampler_shift + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None + t0 = time.time() + try: + processed = processing.process_images(p) + except Exception as e: + err = str(e) + errors.display(e, 'video') + t1 = time.time() + shared.state.disable_preview = False + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + p.close() + + # done + if err: + return video_utils.queue_err(err) + if processed is None or len(processed.images) == 0: + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + generation_info_js = processed.js() if processed is not None else '' + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 3313cab69..c2b3878cf 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -1,6 +1,6 @@ import os import time -from modules import shared, timer, sd_models, sd_checkpoint, model_quant, devices +from modules import shared, errors, timer, sd_models, sd_checkpoint, model_quant, devices from modules.video_models import models_def @@ -18,6 +18,32 @@ def get_quant(args): return None +def get_url(url): + return f'  {url}
' if url else '' + + +def set_prompt(p): + p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + p.task_args['prompt'] = p.prompt + p.task_args['negative_prompt'] = p.negative_prompt + + +def set_vae_params(frames, tile_frames): + if tile_frames > frames: + if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'): + shared.sd_model.vae.tile_sample_min_num_frames = tile_frames + if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): + shared.sd_model.vae.use_framewise_decoding = True + if hasattr(shared.sd_model.vae, 'enable_tiling'): + shared.sd_model.vae.enable_tiling() + else: + if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): + shared.sd_model.vae.use_framewise_decoding = False + if hasattr(shared.sd_model.vae, 'disable_tiling'): + shared.sd_model.vae.disable_tiling() + + def hijack_vae_decode(*args, **kwargs): t0 = time.time() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) @@ -62,6 +88,7 @@ def load_model(selected: models_def.Model): ) except Exception as e: shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}') + errors.display(e, 'video') text_encoder = None # transformer @@ -77,6 +104,7 @@ def load_model(selected: models_def.Model): ) except Exception as e: shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}') + errors.display(e, 'video') transformer = None # model @@ -91,6 +119,7 @@ def load_model(selected: models_def.Model): ) except Exception as e: shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}') + errors.display(e, 'video') t1 = time.time() sd_models.set_diffuser_options(shared.sd_model) From 9a85e45cb4a45927df5b2102a9e9a430f1921792 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 19 Mar 2025 18:26:44 -0400 Subject: [PATCH 034/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4a6ba094e..0e3cac31c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,12 +1,15 @@ # Change Log for SD.Next -## Update for 2025-03-17 +## Update for 2025-03-19 ### TODO - Gemma3 requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - Remote VAE encode for SD15 and Flux.1: - HunyuanVideo-I2V: - - LTXVideo condition input + - HunyuanVideo: Remote VAE + - HunyuanVideo: Tiny VAE + - LTXVideo-095: Condition input + - LTXVideo-095: Broken offloading ### Highlights for 2025-03-17 @@ -77,6 +80,7 @@ Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP mode - fix hires with latent upscale - fix legacy diffusion latent upscalers - fix upscaler selection in postprocessing + - fix sd35 with batch processing ## Update for 2025-02-28 From 5c46904fd869a7aae9384c031d342b2b78db4758 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 20 Mar 2025 14:31:30 -0400 Subject: [PATCH 035/122] video tab alpha releasee Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 50 +++-- modules/lora/network_overrides.py | 8 +- modules/modeldata.py | 15 +- modules/processing_vae.py | 2 + modules/sd_offload.py | 6 +- modules/sd_vae_remote.py | 18 +- modules/sd_vae_taesd.py | 55 ++++-- modules/taesd/taehv.py | 284 ++++++++++++++++++++++++++++ modules/taesd/taem1.py | 272 ++++++++++++++++++++++++++ modules/ui_video.py | 6 +- modules/video_models/models_def.py | 33 ++++ modules/video_models/run_allegro.py | 8 +- modules/video_models/run_cog.py | 8 +- modules/video_models/run_hunyuan.py | 13 +- modules/video_models/run_ltx.py | 8 +- modules/video_models/run_mochi.py | 8 +- modules/video_models/video_load.py | 80 ++++++++ modules/video_models/video_utils.py | 102 +--------- modules/video_models/video_vae.py | 62 ++++++ scripts/cogvideo.py | 8 +- 20 files changed, 874 insertions(+), 172 deletions(-) create mode 100644 modules/taesd/taehv.py create mode 100644 modules/taesd/taem1.py create mode 100644 modules/video_models/video_load.py create mode 100644 modules/video_models/video_vae.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 0e3cac31c..f784e5f56 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,21 +2,49 @@ ## Update for 2025-03-19 -### TODO - - Gemma3 requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - - Remote VAE encode for SD15 and Flux.1: - - HunyuanVideo-I2V: - - HunyuanVideo: Remote VAE - - HunyuanVideo: Tiny VAE - - LTXVideo-095: Condition input - - LTXVideo-095: Broken offloading +### ToDo/Limitations -### Highlights for 2025-03-17 + - VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` + - VAE Remote encode: SD15 and Flux.1 issues: + - Video: ModernUI support is TBD + - Video: API support is TBD + - Video: Wiki page is TBD + - Video: HunyuanVideo-I2V incompatible with latest transformers + - Video: LTXVideo-095 support for conditioned input + - Video: LTXVideo-095 support for offloading -Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP models, improvements to remote VAE, additional docs/guides. +### Highlights for 2025-03-20 -### Details for 2025-03-17 +Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1** and more! +Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides. +### Details for 2025-03-20 + +- **Video tab** + - initial release so consider this as alpha version + - new top-level tab, replaces previous *video* script in text/image tabs + old scripts are still present, but will be removed in the future + - support for all latest models: + - [Hunyuan](https://huggingface.co/Tencent/HunyuanVideo): *HunyuanVideo, FastHunyuan, SkyReels* | *T2V, I2V* + - [WAN21](https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers): *1.3B, 13B* | *T2V, I2V* + - [LTXVideo](https://huggingface.co/Lightricks/LTX-Video): *0.9.0, 0.9.1, 0.9.5* | *T2V, I2V* + - [CogVideoX](https://huggingface.co/THUDM/CogVideoX-5b): *2B, 5B* | *T2V, I2V* + - [Allegro](https://huggingface.co/rhymes-ai/Allegro): *T2V* + - [Mochi1](https://huggingface.co/genmo/mochi-1-preview): *T2V* + - decoding: + - **Default**: use vae from model + - **Tiny VAE**: support for *Hunyuan, WAN, Mochi* + - **Remote VAE**: support for *Hunyuan* + - **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog* + - additional key points: + - all models are auto-downloaded upon first use + - optional video interpolation while creating video files + - optional video preview in ui + - support for balanced offloading and model offloading + - on-the-fly quantization: *BnB, Quanto, TorchAO* + - different video models support different video resolutions, frame counts, etc. + and may require specific settings - see model links for details + - see *ToDo/Limitations* section for additional notes - **Models** - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer diff --git a/modules/lora/network_overrides.py b/modules/lora/network_overrides.py index 65448ef2d..22d251c47 100644 --- a/modules/lora/network_overrides.py +++ b/modules/lora/network_overrides.py @@ -30,8 +30,14 @@ force_models = [ # forced always 'sc', 'kandinsky', 'hunyuandit', - 'hunyuanvideo', 'auraflow', + # video models + 'hunyuanvideo', + 'cogvideo', + 'wanvideo', + 'ltxvideo', + 'mochivideo', + 'allegrovideo', ] force_classes = [ # forced always diff --git a/modules/modeldata.py b/modules/modeldata.py index 078a5b372..f066eaa68 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -29,8 +29,6 @@ def get_model_type(pipe): model_type = 'auraflow' elif "Flux" in name: model_type = 'f1' - elif "Mochi" in name: - model_type = 'mochi' elif "Lumina2" in name: model_type = 'lumina2' elif "Lumina" in name: @@ -41,12 +39,21 @@ def get_model_type(pipe): model_type = 'cogview3' elif "CogView4" in name: model_type = 'cogview4' - elif "CogVideo" in name: - model_type = 'cogvideox' elif "Sana" in name: model_type = 'sana' + # video models + elif "CogVideo" in name: + model_type = 'cogvideo' elif 'HunyuanVideoPipeline' in name or 'HunyuanSkyreels' in name: model_type = 'hunyuanvideo' + elif 'Wan' in name: + model_type = 'wanvideo' + elif 'LTX' in name: + model_type = 'ltxvideo' + elif "Mochi" in name: + model_type = 'mochivideo' + elif "Allegro" in name: + model_type = 'allegrovideo' else: model_type = name return model_type diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 348bbc0c5..00eecc092 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -239,6 +239,8 @@ def vae_postprocess(tensor, model, output_type='np'): if len(tensor.shape) == 3 and tensor.shape[0] == 3: tensor = tensor.unsqueeze(0) if hasattr(model, 'video_processor'): + if len(tensor.shape) == 6 and tensor.shape[1] == 1: + tensor = tensor.squeeze(0) images = model.video_processor.postprocess_video(tensor, output_type='pil') elif hasattr(model, 'image_processor'): images = model.image_processor.postprocess(tensor, output_type=output_type) diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 2c4126209..cadfd3c02 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -10,7 +10,7 @@ from modules.timer import process as process_timer debug_move = shared.log.trace if os.environ.get('SD_MOVE_DEBUG', None) is not None else lambda *args, **kwargs: None -should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'hunyuanvideo', 'cogvideox', 'mochi', 'cogview4'] +should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'cogview4'] offload_hook_instance = None balanced_offload_exclude = ['OmniGenPipeline', 'CogView4Pipeline'] @@ -66,7 +66,7 @@ def set_diffuser_offload(sd_model, op:str='model', quiet:bool=False): if not (hasattr(sd_model, "has_accelerate") and sd_model.has_accelerate): sd_model.has_accelerate = False if shared.opts.diffusers_offload_mode == "none": - if shared.sd_model_type in should_offload: + if shared.sd_model_type in should_offload or 'video' in shared.sd_model_type: shared.log.warning(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} type={shared.sd_model.__class__.__name__} large model') else: shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') @@ -183,7 +183,7 @@ class OffloadHook(accelerate.hooks.ModelHook): return module -def apply_balanced_offload(sd_model, exclude=[]): +def apply_balanced_offload(sd_model=None, exclude=[]): global offload_hook_instance # pylint: disable=global-statement if shared.opts.diffusers_offload_mode != "balanced": return sd_model diff --git a/modules/sd_vae_remote.py b/modules/sd_vae_remote.py index c3591af7b..3c1846ea8 100644 --- a/modules/sd_vae_remote.py +++ b/modules/sd_vae_remote.py @@ -38,11 +38,14 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ return tensors t0 = time.time() modelloader.hf_login() - latents = latents.unsqueeze(0) if len(latents.shape) == 3 else latents + latent_copy = latents.detach().clone().to(device=devices.cpu, dtype=devices.dtype) + latent_copy = latents.unsqueeze(0) if len(latents.shape) == 3 else latents + if model_type == 'hunyuanvideo': + latent_copy = latent_copy.unsqueeze(0) - for i in range(latents.shape[0]): + for i in range(latent_copy.shape[0]): try: - latent = latents[i].detach().clone().to(device=devices.cpu, dtype=devices.dtype) + latent = latent_copy[i] if model_type != 'f1': latent = latent.unsqueeze(0) params = { @@ -51,7 +54,12 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ "dtype": str(latent.dtype).split(".", maxsplit=1)[-1], } headers = { "Content-Type": "tensor/binary" } - if shared.opts.remote_vae_type == 'png': + if 'video' in model_type: + params["partial_postprocess"] = False + params["output_type"] = "pt" + params["output_tensor_type"] = "binary" + headers["Accept"] = "tensor/binary" + elif shared.opts.remote_vae_type == 'png': params["image_format"] = "png" params["output_type"] = "pil" headers["Accept"] = "image/png" @@ -81,7 +89,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ shared.log.error(f'Decode: type="remote" model={model_type} code={response.status_code} shape={latent.shape} url="{url}" args={params} headers={response.headers} response={response.json()}') else: content += len(response.content) - if shared.opts.remote_vae_type == 'raw': + if shared.opts.remote_vae_type == 'raw' or 'video' in model_type: shape = json.loads(response.headers["shape"]) dtype = response.headers["dtype"] tensor = torch.frombuffer(bytearray(response.content), dtype=dtypes[dtype]).reshape(shape) diff --git a/modules/sd_vae_taesd.py b/modules/sd_vae_taesd.py index c8a1b882f..a2a447a3e 100644 --- a/modules/sd_vae_taesd.py +++ b/modules/sd_vae_taesd.py @@ -16,6 +16,9 @@ TAESD_MODELS = { 'TAESD 1.2 Chocolate-Dipped Shortbread': { 'fn': 'taesd_12_', 'uri': 'https://github.com/madebyollin/taesd/raw/8909b44e3befaa0efa79c5791e4fe1c4d4f7884e', 'model': None }, 'TAESD 1.1 Fruit Loops': { 'fn': 'taesd_11_', 'uri': 'https://github.com/madebyollin/taesd/raw/3e8a8a2ab4ad4079db60c1c7dc1379b4cc0c6b31', 'model': None }, 'TAESD 1.0': { 'fn': 'taesd_10_', 'uri': 'https://github.com/madebyollin/taesd/raw/88012e67cf0454e6d90f98911fe9d4aef62add86', 'model': None }, + 'TAE HunyuanVideo': { 'fn': 'taehv.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taehv.pth', 'model': None }, + 'TAE WanVideo': { 'fn': 'taew1.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taew2_1.pth', 'model': None }, + 'TAE MochiVideo': { 'fn': 'taem1.pth', 'uri': 'https://github.com/madebyollin/taem1/raw/refs/heads/main/taem1.pth', 'model': None }, } CQYAN_MODELS = { 'Hybrid-Tiny SD': { @@ -35,49 +38,63 @@ prev_model = '' lock = threading.Lock() -def warn_once(msg): +def warn_once(msg, variant=None): from modules import shared + variant = variant or shared.opts.taesd_variant global prev_warnings # pylint: disable=global-statement if not prev_warnings: prev_warnings = True - shared.log.error(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}": {msg}') + shared.log.error(f'Decode: type="taesd" variant="{variant}": {msg}') return Image.new('RGB', (8, 8), color = (0, 0, 0)) -def get_model(model_type = 'decoder'): +def get_model(model_type = 'decoder', variant = None): global prev_cls, prev_type, prev_model # pylint: disable=global-statement from modules import shared cls = shared.sd_model_type if cls == 'ldm': cls = 'sd' + variant = variant or shared.opts.taesd_variant folder = os.path.join(paths.models_path, "TAESD") os.makedirs(folder, exist_ok=True) - if 'sd' not in cls and 'f1' not in cls: + if 'video' in cls: + return None + if ('sd' not in cls) and ('f1' not in cls): warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported') return None - if shared.opts.taesd_variant.startswith('TAESD'): - cfg = TAESD_MODELS[shared.opts.taesd_variant] - if (cls == prev_cls) and (model_type == prev_type) and (shared.opts.taesd_variant == prev_model) and (cfg['model'] is not None): + if variant.startswith('TAESD'): + cfg = TAESD_MODELS[variant] + if (cls == prev_cls) and (model_type == prev_type) and (variant == prev_model) and (cfg['model'] is not None): return cfg['model'] fn = os.path.join(folder, cfg['fn'] + cls + '_' + model_type + '.pth') if not os.path.exists(fn): uri = cfg['uri'] + '/tae' + cls + '_' + model_type + '.pth' try: - shared.log.info(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}": uri="{uri}" fn="{fn}" download') + shared.log.info(f'Decode: type="taesd" variant="{variant}": uri="{uri}" fn="{fn}" download') torch.hub.download_url_to_file(uri, fn) except Exception as e: warn_once(f'download uri={uri} {e}') if os.path.exists(fn): prev_cls = cls prev_type = model_type - prev_model = shared.opts.taesd_variant - shared.log.debug(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}" fn="{fn}" load') - from modules.taesd.taesd import TAESD - TAESD_MODELS[shared.opts.taesd_variant]['model'] = TAESD(decoder_path=fn if model_type=='decoder' else None, encoder_path=fn if model_type=='encoder' else None) - return TAESD_MODELS[shared.opts.taesd_variant]['model'] - elif shared.opts.taesd_variant.startswith('Hybrid'): - cfg = CQYAN_MODELS[shared.opts.taesd_variant].get(cls, None) - if (cls == prev_cls) and (model_type == prev_type) and (shared.opts.taesd_variant == prev_model) and (cfg['model'] is not None): + prev_model = variant + shared.log.debug(f'Decode: type="taesd" variant="{variant}" fn="{fn}" load') + if 'TAEHV' in variant: + from modules.taesd.taehv import TAEHV + TAESD_MODELS[variant]['model'] = TAEHV(checkpoint_path=fn) + if 'TAEW2' in variant: + from modules.taesd.taehv import TAEHV + TAESD_MODELS[variant]['model'] = TAEHV(checkpoint_path=fn) + elif 'TAEM1' in variant: + from modules.taesd.taem1 import TAEM1 + TAESD_MODELS[variant]['model'] = TAEM1(checkpoint_path=fn) + else: + from modules.taesd.taesd import TAESD + TAESD_MODELS[variant]['model'] = TAESD(decoder_path=fn if model_type=='decoder' else None, encoder_path=fn if model_type=='encoder' else None) + return TAESD_MODELS[variant]['model'] + elif variant.startswith('Hybrid'): + cfg = CQYAN_MODELS[variant].get(cls, None) + if (cls == prev_cls) and (model_type == prev_type) and (variant == prev_model) and (cfg['model'] is not None): return cfg['model'] if cfg is None: warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported') @@ -85,8 +102,8 @@ def get_model(model_type = 'decoder'): repo = cfg['repo'] prev_cls = cls prev_type = model_type - prev_model = shared.opts.taesd_variant - shared.log.debug(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}" id="{repo}" load') + prev_model = variant + shared.log.debug(f'Decode: type="taesd" variant="{variant}" id="{repo}" load') dtype = devices.dtype_vae if devices.dtype_vae != torch.bfloat16 else torch.float16 # taesd does not support bf16 if 'tiny' in repo: from diffusers.models import AutoencoderTiny @@ -95,7 +112,7 @@ def get_model(model_type = 'decoder'): from modules.taesd.hybrid_small import AutoencoderSmall vae = AutoencoderSmall.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir, torch_dtype=dtype) vae = vae.to(devices.device, dtype=dtype) - CQYAN_MODELS[shared.opts.taesd_variant][cls]['model'] = vae + CQYAN_MODELS[variant][cls]['model'] = vae return vae else: warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported') diff --git a/modules/taesd/taehv.py b/modules/taesd/taehv.py new file mode 100644 index 000000000..4a424f137 --- /dev/null +++ b/modules/taesd/taehv.py @@ -0,0 +1,284 @@ +#!/usr/bin/env python3 +""" +Tiny AutoEncoder for Hunyuan Video +(DNN for encoding / decoding videos to Hunyuan Video's latent space) +""" +from collections import namedtuple +import torch +import torch.nn as nn +import torch.nn.functional as F +from tqdm.auto import tqdm + +DecoderResult = namedtuple("DecoderResult", ("frame", "memory")) +TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index")) + +def conv(n_in, n_out, **kwargs): + return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs) + +class Clamp(nn.Module): + def forward(self, x): + return torch.tanh(x / 3) * 3 + +class MemBlock(nn.Module): + def __init__(self, n_in, n_out): + super().__init__() + self.conv = nn.Sequential(conv(n_in * 2, n_out), nn.ReLU(inplace=True), conv(n_out, n_out), nn.ReLU(inplace=True), conv(n_out, n_out)) + self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() + self.act = nn.ReLU(inplace=True) + def forward(self, x, past): + return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x)) + +class TPool(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f*stride,n_f, 1, bias=False) + def forward(self, x): + _NT, C, H, W = x.shape + return self.conv(x.reshape(-1, self.stride * C, H, W)) + +class TGrow(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False) + def forward(self, x): + _NT, C, H, W = x.shape + x = self.conv(x) + return x.reshape(-1, C, H, W) + +def apply_model_with_memblocks(model, x, parallel, show_progress_bar): + """ + Apply a sequential model with memblocks to the given input. + Args: + - model: nn.Sequential of blocks to apply + - x: input data, of dimensions NTCHW + - parallel: if True, parallelize over timesteps (fast but uses O(T) memory) + if False, each timestep will be processed sequentially (slow but uses O(1) memory) + - show_progress_bar: if True, enables tqdm progressbar display + + Returns NTCHW tensor of output data. + """ + assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor" + N, T, C, H, W = x.shape + if parallel: + x = x.reshape(N*T, C, H, W) + # parallel over input timesteps, iterate over blocks + for b in tqdm(model, disable=not show_progress_bar): + if isinstance(b, MemBlock): + NT, C, H, W = x.shape + T = NT // N + _x = x.reshape(N, T, C, H, W) + mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape) + x = b(x, mem) + else: + x = b(x) + NT, C, H, W = x.shape + T = NT // N + x = x.view(N, T, C, H, W) + else: + # TODO(oboerbohan): at least on macos this still gradually uses more memory during decode... + # need to fix :( + out = [] + # iterate over input timesteps and also iterate over blocks. + # because of the cursed TPool/TGrow blocks, this is not a nested loop, + # it's actually a ***graph traversal*** problem! so let's make a queue + work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))] + # in addition to manually managing our queue, we also need to manually manage our progressbar. + # we'll update it for every source node that we consume. + progress_bar = tqdm(range(T), disable=not show_progress_bar) + # we'll also need a separate addressable memory per node as well + mem = [None] * len(model) + while work_queue: + xt, i = work_queue.pop(0) + if i == 0: + # new source node consumed + progress_bar.update(1) + if i == len(model): + # reached end of the graph, append result to output list + out.append(xt) + else: + # fetch the block to process + b = model[i] + if isinstance(b, MemBlock): + # mem blocks are simple since we're visiting the graph in causal order + if mem[i] is None: + xt_new = b(xt, xt * 0) + mem[i] = xt + else: + xt_new = b(xt, mem[i]) + mem[i].copy_(xt) # inplace might reduce mysterious pytorch memory allocations? doesn't help though + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_new, i+1)) + elif isinstance(b, TPool): + # pool blocks are miserable + if mem[i] is None: + mem[i] = [] # pool memory is itself a queue of inputs to pool + mem[i].append(xt) + if len(mem[i]) > b.stride: + # pool mem is in invalid state, we should have pooled before this + raise ValueError("???") + elif len(mem[i]) < b.stride: + # pool mem is not yet full, go back to processing the work queue + pass + else: + # pool mem is ready, run the pool block + N, C, H, W = xt.shape + xt = b(torch.cat(mem[i], 1).view(N*b.stride, C, H, W)) + # reset the pool mem + mem[i] = [] + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i+1)) + elif isinstance(b, TGrow): + xt = b(xt) + NT, C, H, W = xt.shape + # each tgrow has multiple successor nodes + for xt_next in reversed(xt.view(N, b.stride*C, H, W).chunk(b.stride, 1)): + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_next, i+1)) + else: + # normal block with no funny business + xt = b(xt) + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i+1)) + progress_bar.close() + x = torch.stack(out, 1) + return x + +class TAEHV(nn.Module): + latent_channels = 16 + image_channels = 3 + def __init__(self, checkpoint_path="taehv.pth", decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True)): + """Initialize pretrained TAEHV from the given checkpoint. + + Arg: + checkpoint_path: path to weight file to load. taehv.pth for Hunyuan, taew2_1.pth for Wan 2.1. + decoder_time_upscale: whether temporal upsampling is enabled for each block. upsampling can be disabled for a cheaper preview. + decoder_space_upscale: whether spatial upsampling is enabled for each block. upsampling can be disabled for a cheaper preview. + """ + super().__init__() + self.encoder = nn.Sequential( + conv(TAEHV.image_channels, 64), nn.ReLU(inplace=True), + TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + TPool(64, 1), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + conv(64, TAEHV.latent_channels), + ) + n_f = [256, 128, 64, 64] + self.frames_to_trim = 2**sum(decoder_time_upscale) - 1 + self.decoder = nn.Sequential( + Clamp(), conv(TAEHV.latent_channels, n_f[0]), nn.ReLU(inplace=True), + MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False), + MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1), TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1), conv(n_f[1], n_f[2], bias=False), + MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1), TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1), conv(n_f[2], n_f[3], bias=False), + nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels), + ) + if checkpoint_path is not None: + self.load_state_dict(self.patch_tgrow_layers(torch.load(checkpoint_path, map_location="cpu", weights_only=True))) + + def patch_tgrow_layers(self, sd): + """Patch TGrow layers to use a smaller kernel if needed. + + Args: + sd: state dict to patch + """ + new_sd = self.state_dict() + for i, layer in enumerate(self.decoder): + if isinstance(layer, TGrow): + key = f"decoder.{i}.conv.weight" + if sd[key].shape[0] > new_sd[key].shape[0]: + # take the last-timestep output channels + sd[key] = sd[key][-new_sd[key].shape[0]:] + return sd + + def encode_video(self, x, parallel=True, show_progress_bar=True): + """Encode a sequence of frames. + + Args: + x: input NTCHW RGB (C=3) tensor with values in [0, 1]. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW latent tensor with ~Gaussian values. + """ + return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar) + + def decode_video(self, x, parallel=True, show_progress_bar=True): + """Decode a sequence of frames. + + Args: + x: input NTCHW latent (C=12) tensor with ~Gaussian values. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW RGB tensor with ~[0, 1] values. + """ + x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar) + return x[:, self.frames_to_trim:] + + def forward(self, x): + return self.c(x) + +@torch.no_grad() +def main(): + """Run TAEHV roundtrip reconstruction on the given video paths.""" + import sys + import cv2 # no highly esteemed deed is commemorated here + + class VideoTensorReader: + def __init__(self, video_file_path): + self.cap = cv2.VideoCapture(video_file_path) + assert self.cap.isOpened(), f"Could not load {video_file_path}" + self.fps = self.cap.get(cv2.CAP_PROP_FPS) + def __iter__(self): + return self + def __next__(self): + ret, frame = self.cap.read() + if not ret: + self.cap.release() + raise StopIteration # End of video or error + return torch.from_numpy(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)).permute(2, 0, 1) # BGR HWC -> RGB CHW + + class VideoTensorWriter: + def __init__(self, video_file_path, width_height, fps=30): + self.writer = cv2.VideoWriter(video_file_path, cv2.VideoWriter_fourcc(*'mp4v'), fps, width_height) + assert self.writer.isOpened(), f"Could not create writer for {video_file_path}" + def write(self, frame_tensor): + assert frame_tensor.ndim == 3 and frame_tensor.shape[0] == 3, f"{frame_tensor.shape}??" + self.writer.write(cv2.cvtColor(frame_tensor.permute(1, 2, 0).numpy(), cv2.COLOR_RGB2BGR)) # RGB CHW -> BGR HWC + def __del__(self): + if hasattr(self, 'writer'): + self.writer.release() + + dev = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu") + dtype = torch.float16 + print("Using device", dev, "and dtype", dtype) + taehv = TAEHV().to(dev, dtype) + for video_path in sys.argv[1:]: + print(f"Processing {video_path}...") + video_in = VideoTensorReader(video_path) + video = torch.stack(list(video_in), 0)[None] + vid_dev = video.to(dev, dtype).div_(255.0) + # convert to device tensor + if video.numel() < 100_000_000: + print(f" {video_path} seems small enough, will process all frames in parallel") + # convert to device tensor + vid_enc = taehv.encode_video(vid_dev) + print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...") + vid_dec = taehv.decode_video(vid_enc) + print(f" Decoded {video_path} -> {vid_dec.shape}") + else: + print(f" {video_path} seems large, will process each frame sequentially") + # convert to device tensor + vid_enc = taehv.encode_video(vid_dev, parallel=False) + print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...") + vid_dec = taehv.decode_video(vid_enc, parallel=False) + print(f" Decoded {video_path} -> {vid_dec.shape}") + video_out_path = video_path + ".reconstructed_by_taehv.mp4" + video_out = VideoTensorWriter(video_out_path, (vid_dec.shape[-1], vid_dec.shape[-2]), fps=int(round(video_in.fps))) + for frame in vid_dec.clamp_(0, 1).mul_(255).round_().byte().cpu()[0]: + video_out.write(frame) + print(f" Saved to {video_out_path}") + +if __name__ == "__main__": + main() diff --git a/modules/taesd/taem1.py b/modules/taesd/taem1.py new file mode 100644 index 000000000..7d59ca2b6 --- /dev/null +++ b/modules/taesd/taem1.py @@ -0,0 +1,272 @@ +#!/usr/bin/env python3 +""" +Tiny AutoEncoder for Mochi 1 +(DNN for encoding / decoding videos to Mochi 1's latent space) +""" +from collections import namedtuple +import torch +import torch.nn as nn +import torch.nn.functional as F +from tqdm.auto import tqdm + +DecoderResult = namedtuple("DecoderResult", ("frame", "memory")) +TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index")) + +def conv(n_in, n_out, **kwargs): + return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs) + +class Clamp(nn.Module): + def forward(self, x): + return torch.tanh(x / 3) * 3 + +class MemBlock(nn.Module): + def __init__(self, n_in, n_out): + super().__init__() + self.conv = nn.Sequential(conv(n_in * 2, n_out), nn.ReLU(inplace=True), conv(n_out, n_out), nn.ReLU(inplace=True), conv(n_out, n_out)) + self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() + self.act = nn.ReLU(inplace=True) + def forward(self, x, past): + return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x)) + +class TPool(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f*stride,n_f, 1, bias=False) + def forward(self, x): + _NT, C, H, W = x.shape + return self.conv(x.reshape(-1, self.stride * C, H, W)) + +class TGrow(nn.Module): + def __init__(self, n_f, stride): + super().__init__() + self.stride = stride + self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False) + def forward(self, x): + _NT, C, H, W = x.shape + x = self.conv(x) + return x.reshape(-1, C, H, W) + +def apply_model_with_memblocks(model, x, parallel, show_progress_bar): + """ + Apply a sequential model with memblocks to the given input. + Args: + - model: nn.Sequential of blocks to apply + - x: input data, of dimensions NTCHW + - parallel: if True, parallelize over timesteps (fast but uses O(T) memory) + if False, each timestep will be processed sequentially (slow but uses O(1) memory) + - show_progress_bar: if True, enables tqdm progressbar display + + Returns NTCHW tensor of output data. + """ + assert x.ndim == 5, f"TAEM1 operates on NTCHW tensors, but got {x.ndim}-dim tensor" + N, T, C, H, W = x.shape + if parallel: + x = x.reshape(N*T, C, H, W) + # parallel over input timesteps, iterate over blocks + for b in tqdm(model, disable=not show_progress_bar): + if isinstance(b, MemBlock): + NT, C, H, W = x.shape + T = NT // N + _x = x.reshape(N, T, C, H, W) + mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape) + x = b(x, mem) + else: + x = b(x) + NT, C, H, W = x.shape + T = NT // N + x = x.view(N, T, C, H, W) + else: + # TODO(oboerbohan): at least on macos this still gradually uses more memory during decode... + # need to fix :( + out = [] + # iterate over input timesteps and also iterate over blocks. + # because of the cursed TPool/TGrow blocks, this is not a nested loop, + # it's actually a ***graph traversal*** problem! so let's make a queue + work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))] + # in addition to manually managing our queue, we also need to manually manage our progressbar. + # we'll update it for every source node that we consume. + progress_bar = tqdm(range(T), disable=not show_progress_bar) + # we'll also need a separate addressable memory per node as well + mem = [None] * len(model) + while work_queue: + xt, i = work_queue.pop(0) + if i == 0: + # new source node consumed + progress_bar.update(1) + if i == len(model): + # reached end of the graph, append result to output list + out.append(xt) + else: + # fetch the block to process + b = model[i] + if isinstance(b, MemBlock): + # mem blocks are simple since we're visiting the graph in causal order + if mem[i] is None: + xt_new = b(xt, xt * 0) + mem[i] = xt + else: + xt_new = b(xt, mem[i]) + mem[i].copy_(xt) # inplace might reduce mysterious pytorch memory allocations? doesn't help though + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_new, i+1)) + elif isinstance(b, TPool): + # pool blocks are miserable + if mem[i] is None: + mem[i] = [] # pool memory is itself a queue of inputs to pool + mem[i].append(xt) + if len(mem[i]) > b.stride: + # pool mem is in invalid state, we should have pooled before this + raise ValueError("???") + elif len(mem[i]) < b.stride: + # pool mem is not yet full, go back to processing the work queue + pass + else: + # pool mem is ready, run the pool block + N, C, H, W = xt.shape + xt = b(torch.cat(mem[i], 1).view(N*b.stride, C, H, W)) + # reset the pool mem + mem[i] = [] + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i+1)) + elif isinstance(b, TGrow): + xt = b(xt) + NT, C, H, W = xt.shape + # each tgrow has multiple successor nodes + for xt_next in reversed(xt.view(N, b.stride*C, H, W).chunk(b.stride, 1)): + # add successor to work queue + work_queue.insert(0, TWorkItem(xt_next, i+1)) + else: + # normal block with no funny business + xt = b(xt) + # add successor to work queue + work_queue.insert(0, TWorkItem(xt, i+1)) + progress_bar.close() + x = torch.stack(out, 1) + return x + +class TAEM1(nn.Module): + latent_channels = 12 + image_channels = 3 + def __init__(self, checkpoint_path="taem1.pth"): + """Initialize pretrained TAEM1 from the given checkpoints.""" + super().__init__() + self.encoder = nn.Sequential( + conv(TAEM1.image_channels, 64), nn.ReLU(inplace=True), + TPool(64, 3), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + TPool(64, 1), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64), + conv(64, TAEM1.latent_channels), + ) + n_f = [256, 128, 64, 64] + self.decoder = nn.Sequential( + Clamp(), conv(TAEM1.latent_channels, n_f[0]), nn.ReLU(inplace=True), + MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False), + MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2), TGrow(n_f[1], 2), conv(n_f[1], n_f[2], bias=False), + MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), nn.Upsample(scale_factor=2), TGrow(n_f[2], 3), conv(n_f[2], n_f[3], bias=False), + nn.ReLU(inplace=True), conv(n_f[3], TAEM1.image_channels), + ) + if checkpoint_path is not None: + self.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True)) + + def encode_video(self, x, parallel=True, show_progress_bar=True): + """Encode a sequence of frames. + + Args: + x: input NTCHW RGB (C=3) tensor with values in [0, 1]. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW latent tensor with ~Gaussian values. + """ + return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar) + + def decode_video(self, x, parallel=True, show_progress_bar=True): + """Decode a sequence of frames. + + Args: + x: input NTCHW latent (C=12) tensor with ~Gaussian values. + parallel: if True, all frames will be processed at once. + (this is faster but may require more memory). + if False, frames will be processed sequentially. + Returns NTCHW RGB tensor with ~[0, 1] values. + """ + x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar) + # NOTE: + # the Mochi VAE does not preserve shape along the time axis; + # videos are encoded to floor((n_in - 1)/6)+1 latent frames + # (which makes sense, it's stride 6, so 12 -> 2 and 13->3) + # but then they're decoded to only the *minimal* number + # of input frames (3 latents get decoded to 13 frames, not 18) + # in order to achieve the intended causal structure... + # anyway, that's why we have to remove some frames here. + # mochi-VAE does the slicing at each TGrow (save compute/mem?) + # but I think it's basically the same + return x[:, 5:] + + def forward(self, x): + return self.c(x) + +@torch.no_grad() +def main(): + """Run TAEM1 roundtrip reconstruction on the given video paths.""" + import sys + import cv2 # no highly esteemed deed is commemorated here + + class VideoTensorReader: + def __init__(self, video_file_path): + self.cap = cv2.VideoCapture(video_file_path) + assert self.cap.isOpened(), f"Could not load {video_file_path}" + self.fps = self.cap.get(cv2.CAP_PROP_FPS) + def __iter__(self): + return self + def __next__(self): + ret, frame = self.cap.read() + if not ret: + self.cap.release() + raise StopIteration # End of video or error + return torch.from_numpy(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)).permute(2, 0, 1) # BGR HWC -> RGB CHW + + class VideoTensorWriter: + def __init__(self, video_file_path, width_height, fps=30): + self.writer = cv2.VideoWriter(video_file_path, cv2.VideoWriter_fourcc(*'mp4v'), fps, width_height) + assert self.writer.isOpened(), f"Could not create writer for {video_file_path}" + def write(self, frame_tensor): + assert frame_tensor.ndim == 3 and frame_tensor.shape[0] == 3, f"{frame_tensor.shape}??" + self.writer.write(cv2.cvtColor(frame_tensor.permute(1, 2, 0).numpy(), cv2.COLOR_RGB2BGR)) # RGB CHW -> BGR HWC + def __del__(self): + if hasattr(self, 'writer'): + self.writer.release() + + dev = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu") + dtype = torch.float16 + print("Using device", dev, "and dtype", dtype) + taem1 = TAEM1().to(dev, dtype) + for video_path in sys.argv[1:]: + print(f"Processing {video_path}...") + video_in = VideoTensorReader(video_path) + video = torch.stack(list(video_in), 0)[None] + vid_dev = video.to(dev, dtype).div_(255.0) + # convert to device tensor + if video.numel() < 100_000_000: + print(f" {video_path} seems small enough, will process all frames in parallel") + # convert to device tensor + vid_enc = taem1.encode_video(vid_dev) + print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...") + vid_dec = taem1.decode_video(vid_enc) + print(f" Decoded {video_path} -> {vid_dec.shape}") + else: + print(f" {video_path} seems large, will process each frame sequentially") + # convert to device tensor + vid_enc = taem1.encode_video(vid_dev, parallel=False) + print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...") + vid_dec = taem1.decode_video(vid_enc, parallel=False) + print(f" Decoded {video_path} -> {vid_dec.shape}") + video_out_path = video_path + ".reconstructed_by_taem1.mp4" + video_out = VideoTensorWriter(video_out_path, (vid_dec.shape[-1], vid_dec.shape[-2]), fps=int(round(video_in.fps))) + for frame in vid_dec.clamp_(0, 1).mul_(255).round_().byte().cpu()[0]: + video_out.write(frame) + print(f" Saved to {video_out_path}") + +if __name__ == "__main__": + main() diff --git a/modules/ui_video.py b/modules/ui_video.py index 05a3bfc09..905cdfb95 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -1,7 +1,7 @@ import gradio as gr from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste from modules.ui_components import ToolButton -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load def engine_change(engine): @@ -21,7 +21,7 @@ def model_change(engine, model): sd_models.unload_model_weights() msg = 'Video model unloaded' else: - msg = video_utils.load_model(selected) + msg = video_load.load_model(selected) else: sd_models.unload_model_weights() msg = 'Video model unloaded' @@ -84,7 +84,7 @@ def create_ui(): steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") with gr.Row(): sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") - dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) + dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) # TODO video: dynamic shift with gr.Row(): guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 82e6f108a..5ff58fd82 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -17,6 +17,7 @@ class Model(): te_folder: str = 'text_encoder' te_hijack: bool = True vae_hijack: bool = True + vae_remote: bool = False models = { @@ -25,18 +26,21 @@ models = { Model(name='None'), Model(name='Hunyuan Video T2V', url='https://huggingface.co/tencent/HunyuanVideo', + vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='Hunyuan Video I2V', # https://github.com/huggingface/diffusers/pull/10983 url='https://huggingface.co/tencent/HunyuanVideo-I2V', + vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo-I2V', repo_cls=diffusers.HunyuanVideoImageToVideoPipeline, te_cls=transformers.LlavaForConditionalGeneration, dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='SkyReels Hunyuan T2V', # https://github.com/huggingface/diffusers/pull/10837 url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-T2V', + vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, @@ -45,6 +49,7 @@ models = { dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='SkyReels Hunyuan I2V', # https://github.com/huggingface/diffusers/pull/10837 url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-I2V', + vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo', te_cls=transformers.LlamaModel, dit='Skywork/SkyReels-V1-Hunyuan-I2V', @@ -52,6 +57,7 @@ models = { dit_cls=diffusers.HunyuanVideoTransformer3DModel), Model(name='Fast Hunyuan T2V', # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213 url='https://huggingface.co/FastVideo/FastHunyuan', + vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo', repo_cls=diffusers.HunyuanVideoPipeline, te_cls=transformers.LlamaModel, @@ -97,6 +103,33 @@ models = { te_cls=transformers.T5EncoderModel, dit_cls=diffusers.LTXVideoTransformer3DModel), ], + 'WAN Video': [ + Model(name='None'), + Model(name='WAN 2.1 1.3B T2V', + url='https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers', + repo='Wan-AI/Wan2.1-T2V-1.3B-Diffusers', + repo_cls=diffusers.WanPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.WanTransformer3DModel), + Model(name='WAN 2.1 14B T2V', + url='https://huggingface.co/Wan-AI/Wan2.1-T2V-14B-Diffusers', + repo='Wan-AI/Wan2.1-T2V-14B-Diffusers', + repo_cls=diffusers.WanPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.WanTransformer3DModel), + Model(name='WAN 2.1 14B I2V 480p', + url='https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-480P-Diffusers', + repo='Wan-AI/Wan2.1-I2V-14B-480P-Diffusers', + repo_cls=diffusers.WanImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.WanTransformer3DModel), + Model(name='WAN 2.1 14B I2V 720p', + url='https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-720P-Diffusers', + repo='Wan-AI/Wan2.1-I2V-14B-720P-Diffusers', + repo_cls=diffusers.WanImageToVideoPipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.WanTransformer3DModel), + ], 'Mochi Video': [ Model(name='None'), Model(name='Mochi 1 T2V', diff --git a/modules/video_models/run_allegro.py b/modules/video_models/run_allegro.py index be1b0da8d..091be7333 100644 --- a/modules/video_models/run_allegro.py +++ b/modules/video_models/run_allegro.py @@ -1,7 +1,7 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load, video_vae debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,7 +14,7 @@ def generate(*args, **kwargs): found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: - video_utils.load_model(selected) + video_load.load_model(selected) if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -54,13 +54,13 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) - video_utils.set_vae_params(p.frames, vae_tile_frames) + video_vae.set_vae_params(p) video_utils.set_prompt(p) p.task_args['output_type'] = 'pil' p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift shared.opts.data['schedulers_shift'] = sampler_shift debug(f'Video: task_args={p.task_args}') diff --git a/modules/video_models/run_cog.py b/modules/video_models/run_cog.py index 5fbe1d320..0aaee4b97 100644 --- a/modules/video_models/run_cog.py +++ b/modules/video_models/run_cog.py @@ -1,7 +1,7 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load, video_vae debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,7 +14,7 @@ def generate(*args, **kwargs): found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: - video_utils.load_model(selected) + video_load.load_model(selected) if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -54,13 +54,13 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) - video_utils.set_vae_params(p.frames, vae_tile_frames) + video_vae.set_vae_params(p) video_utils.set_prompt(p) p.task_args['output_type'] = 'pil' p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift shared.opts.data['schedulers_shift'] = sampler_shift debug(f'Video: task_args={p.task_args}') diff --git a/modules/video_models/run_hunyuan.py b/modules/video_models/run_hunyuan.py index aa076f935..0f722b890 100644 --- a/modules/video_models/run_hunyuan.py +++ b/modules/video_models/run_hunyuan.py @@ -1,7 +1,7 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load, video_vae debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,7 +14,7 @@ def generate(*args, **kwargs): found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - video_utils.load_model(selected) + video_load.load_model(selected) if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -38,6 +38,9 @@ def generate(*args, **kwargs): vae_tile_frames=int(vae_tile_frames), override_settings=override_settings, ) + if p.vae_type == 'Remote' and not selected.vae_remote: + shared.log.warning(f'Video: model={selected.name} remote vae not supported') + p.vae_type = 'Default' p.scripts = None p.script_args = None p.state = ui_state @@ -54,13 +57,13 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) - video_utils.set_vae_params(p.frames, vae_tile_frames) + video_vae.set_vae_params(p) video_utils.set_prompt(p) - p.task_args['output_type'] = 'pil' + p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil' p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift shared.opts.data['schedulers_shift'] = sampler_shift debug(f'Video: task_args={p.task_args}') diff --git a/modules/video_models/run_ltx.py b/modules/video_models/run_ltx.py index ba01e45a4..ef3fa08b8 100644 --- a/modules/video_models/run_ltx.py +++ b/modules/video_models/run_ltx.py @@ -1,7 +1,7 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load, video_vae debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,7 +14,7 @@ def generate(*args, **kwargs): found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: - video_utils.load_model(selected) + video_load.load_model(selected) if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -54,13 +54,13 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) - video_utils.set_vae_params(p.frames, vae_tile_frames) + video_vae.set_vae_params(p) video_utils.set_prompt(p) p.task_args['output_type'] = 'pil' p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift shared.opts.data['schedulers_shift'] = sampler_shift debug(f'Video: task_args={p.task_args}') diff --git a/modules/video_models/run_mochi.py b/modules/video_models/run_mochi.py index fff3a6d04..ef705880e 100644 --- a/modules/video_models/run_mochi.py +++ b/modules/video_models/run_mochi.py @@ -1,7 +1,7 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils +from modules.video_models import models_def, video_utils, video_load, video_vae debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,7 +14,7 @@ def generate(*args, **kwargs): found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: - video_utils.load_model(selected) + video_load.load_model(selected) if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -54,13 +54,13 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) - video_utils.set_vae_params(p.frames, vae_tile_frames) + video_vae.set_vae_params(p) video_utils.set_prompt(p) p.task_args['output_type'] = 'pil' p.ops.append('video') orig_dynamic_shift = shared.opts.schedulers_dynamic_shift orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift shared.opts.data['schedulers_shift'] = sampler_shift debug(f'Video: task_args={p.task_args}') diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py new file mode 100644 index 000000000..2c8634f1e --- /dev/null +++ b/modules/video_models/video_load.py @@ -0,0 +1,80 @@ +import os +import time +from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices +from modules.video_models import models_def, video_utils, video_vae + + +loaded_model = None +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def load_model(selected: models_def.Model): + if selected is None: + return + global loaded_model # pylint: disable=global-statement + if loaded_model == selected.name: + return + sd_models.unload_model_weights() + t0 = time.time() + + # text encoder + try: + quant_args = model_quant.create_config(module='Text Encoder') + debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={video_utils.get_quant(quant_args)}') + text_encoder = selected.te_cls.from_pretrained( + pretrained_model_name_or_path=selected.te or selected.repo, + subfolder=selected.te_folder, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}') + errors.display(e, 'video') + text_encoder = None + + # transformer + try: + quant_args = model_quant.create_config(module='Model') + debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={video_utils.get_quant(quant_args)}') + transformer = selected.dit_cls.from_pretrained( + pretrained_model_name_or_path=selected.dit or selected.repo, + subfolder=selected.dit_folder, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args + ) + except Exception as e: + shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}') + errors.display(e, 'video') + transformer = None + + # model + try: + debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}') + shared.sd_model = selected.repo_cls.from_pretrained( + pretrained_model_name_or_path=selected.repo, + transformer=transformer, + text_encoder=text_encoder, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + ) + except Exception as e: + shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}') + errors.display(e, 'video') + + t1 = time.time() + shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) + shared.sd_model.sd_model_hash = None + sd_models.set_diffuser_options(shared.sd_model) + if selected.vae_hijack: + shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode + shared.sd_model.vae.decode = video_vae.hijack_vae_decode + if selected.te_hijack: + shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt + shared.sd_model.encode_prompt = video_utils.hijack_encode_prompt + shared.sd_model.vae.enable_slicing() + loaded_model = selected.name + msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' + shared.log.info(msg) + return msg diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index c2b3878cf..0ed4791eb 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -1,7 +1,6 @@ import os import time -from modules import shared, errors, timer, sd_models, sd_checkpoint, model_quant, devices -from modules.video_models import models_def +from modules import shared, sd_models, timer debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -29,31 +28,6 @@ def set_prompt(p): p.task_args['negative_prompt'] = p.negative_prompt -def set_vae_params(frames, tile_frames): - if tile_frames > frames: - if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'): - shared.sd_model.vae.tile_sample_min_num_frames = tile_frames - if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): - shared.sd_model.vae.use_framewise_decoding = True - if hasattr(shared.sd_model.vae, 'enable_tiling'): - shared.sd_model.vae.enable_tiling() - else: - if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): - shared.sd_model.vae.use_framewise_decoding = False - if hasattr(shared.sd_model.vae, 'disable_tiling'): - shared.sd_model.vae.disable_tiling() - - -def hijack_vae_decode(*args, **kwargs): - t0 = time.time() - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - debug(f'Video decode: vae={shared.sd_model.vae.__class__.__name__} time={t1-t0:.2f}') - return res - - def hijack_encode_prompt(*args, **kwargs): t0 = time.time() res = shared.sd_model.orig_encode_prompt(*args, **kwargs) @@ -62,77 +36,3 @@ def hijack_encode_prompt(*args, **kwargs): debug(f'Video encode: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) return res - -loaded_model = None - - -def load_model(selected: models_def.Model): - if selected is None: - return - global loaded_model # pylint: disable=global-statement - if loaded_model == selected.name: - return - sd_models.unload_model_weights() - t0 = time.time() - - # text encoder - try: - quant_args = model_quant.create_config(module='Text Encoder') - debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={get_quant(quant_args)}') - text_encoder = selected.te_cls.from_pretrained( - pretrained_model_name_or_path=selected.te or selected.repo, - subfolder=selected.te_folder, - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - except Exception as e: - shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}') - errors.display(e, 'video') - text_encoder = None - - # transformer - try: - quant_args = model_quant.create_config(module='Model') - debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={get_quant(quant_args)}') - transformer = selected.dit_cls.from_pretrained( - pretrained_model_name_or_path=selected.dit or selected.repo, - subfolder=selected.dit_folder, - torch_dtype=devices.dtype, - cache_dir=shared.opts.hfcache_dir, - **quant_args - ) - except Exception as e: - shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}') - errors.display(e, 'video') - transformer = None - - # model - try: - debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}') - shared.sd_model = selected.repo_cls.from_pretrained( - pretrained_model_name_or_path=selected.repo, - transformer=transformer, - text_encoder=text_encoder, - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - ) - except Exception as e: - shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}') - errors.display(e, 'video') - - t1 = time.time() - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) - shared.sd_model.sd_model_hash = None - if selected.vae_hijack: - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.vae.decode = hijack_vae_decode - if selected.te_hijack: - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.encode_prompt = hijack_encode_prompt - shared.sd_model.vae.enable_slicing() - loaded_model = selected.name - msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' - shared.log.info(msg) - return msg diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py new file mode 100644 index 000000000..bd552f66c --- /dev/null +++ b/modules/video_models/video_vae.py @@ -0,0 +1,62 @@ +import os +import time +from modules import shared, sd_models, devices, timer + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None +vae_type = None + + +def set_vae_params(p): + global vae_type # pylint: disable=global-statement + vae_type = p.vae_type + if p.vae_tile_frames > p.frames: + if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'): + shared.sd_model.vae.tile_sample_min_num_frames = p.vae_tile_frames + if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): + shared.sd_model.vae.use_framewise_decoding = True + if hasattr(shared.sd_model.vae, 'enable_tiling'): + shared.sd_model.vae.enable_tiling() + else: + if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): + shared.sd_model.vae.use_framewise_decoding = False + if hasattr(shared.sd_model.vae, 'disable_tiling'): + shared.sd_model.vae.disable_tiling() + + +def vae_decode_tiny(latents): + if 'Hunyuan' in shared.sd_model.__class__.__name__: + variant = 'TAE HunyuanVideo' + elif 'Mochi' in shared.sd_model.__class__.__name__: + variant = 'TAE MochiVideo' + elif 'WAN' in shared.sd_model.__class__.__name__: + variant = 'TAE WanVideo' + else: + shared.log.warning(f'Video VAE: type=Tiny cls={shared.sd_model.__class__.__name__} not supported') + return None + from modules import sd_vae_taesd + vae = sd_vae_taesd.get_model(variant) + if vae is None: + return None + debug(f'Video VAE: type=Tiny cls={vae.__class__.__name__} variant="{variant}" latents={latents.shape}') + vae = vae.to(device=devices.device, dtype=devices.dtype) + latents = latents.transpose(1, 2).to(device=devices.device, dtype=devices.dtype) + images = vae.decode_video(latents, parallel=False).transpose(1, 2).mul_(2).sub_(1) + images = images.transpose(1, 2).mul_(2).sub_(1) + return (images, None) + + +def hijack_vae_decode(*args, **kwargs): + t0 = time.time() + res = None + if vae_type == 'Tiny': + res = vae_decode_tiny(args[0]) + if vae_type == 'Remote': + pass + if res is None: + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) + res = shared.sd_model.vae.orig_decode(*args, **kwargs) + t1 = time.time() + timer.process.add('vae', t1-t0) + debug(f'Video decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}') + return res diff --git a/scripts/cogvideo.py b/scripts/cogvideo.py index 7184dd946..c18b4eb2f 100644 --- a/scripts/cogvideo.py +++ b/scripts/cogvideo.py @@ -50,7 +50,7 @@ class Script(scripts.Script): return [model, sampler, frames, guidance, offload, override, video_type, duration, loop, pad, interpolate, image, video] def load(self, model): - if (shared.sd_model_type != 'cogvideox' or shared.sd_model.sd_model_checkpoint != model) and model != 'None': + if (shared.sd_model_type != 'cogvideo' or shared.sd_model.sd_model_checkpoint != model) and model != 'None': sd_models.unload_model_weights('model') shared.log.info(f'CogVideoX load: model="{model}"') try: @@ -64,7 +64,7 @@ class Script(scripts.Script): shared.log.error(f'Load CogVideoX: {e}') if debug: errors.display(e, 'CogVideoX') - if shared.sd_model_type == 'cogvideox' and model != 'None': + if shared.sd_model_type == 'cogvideo' and model != 'None': shared.sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m', ncols=80, colour='#327fba') shared.log.debug(f'CogVideoX load: class="{shared.sd_model.__class__.__name__}"') if shared.sd_model is not None and model == 'None': @@ -74,7 +74,7 @@ class Script(scripts.Script): devices.torch_gc() def offload(self, offload): - if shared.sd_model_type != 'cogvideox': + if shared.sd_model_type != 'cogvideo': return if offload == 'none': sd_models.move_model(shared.sd_model, devices.device) @@ -131,7 +131,7 @@ class Script(scripts.Script): return img def generate(self, p: processing.StableDiffusionProcessing, model: str): - if shared.sd_model_type != 'cogvideox': + if shared.sd_model_type != 'cogvideo': return [] shared.log.info(f'CogVideoX: sampler={p.sampler} steps={p.steps} frames={p.frames} width={p.width} height={p.height} seed={p.seed} guidance={p.guidance}') if p.sampler == 'DDIM': From 1b51f32251fa2b9924b9caa7c949d9aaff55097d Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 20 Mar 2025 14:39:35 -0400 Subject: [PATCH 036/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 5 +++++ modules/processing_diffusers.py | 2 +- 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f784e5f56..742bf342b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -38,10 +38,15 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio - **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog* - additional key points: - all models are auto-downloaded upon first use + uses *system paths -> huggingface* folder + - support for many video types - optional video interpolation while creating video files - optional video preview in ui + present if video output is selected - support for balanced offloading and model offloading + uses system settings - on-the-fly quantization: *BnB, Quanto, TorchAO* + uses system settings, granular for *transformer* and *text-encoder* separately - different video models support different video resolutions, frame counts, etc. and may require specific settings - see model links for details - see *ToDo/Limitations* section for additional notes diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 366a62995..684484217 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -74,7 +74,7 @@ def process_base(p: processing.StableDiffusionProcessing): true_cfg_scale=p.diffusers_guidance_rescale, denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None, denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None, - num_frames=getattr(p, 'frames', None), + num_frames=getattr(p, 'frames', 1), output_type='latent', clip_skip=p.clip_skip, desc='Base', From c3580069870dfd59c79004103586649d0de90755 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 20 Mar 2025 14:48:41 -0400 Subject: [PATCH 037/122] add wan21 Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 2 +- modules/ui_video.py | 3 + modules/video_models/run_wan.py | 94 ++++++++++++++++++++++++++++++ modules/video_models/video_load.py | 4 +- 4 files changed, 100 insertions(+), 3 deletions(-) create mode 100644 modules/video_models/run_wan.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 742bf342b..1d099df92 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,7 +26,7 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio old scripts are still present, but will be removed in the future - support for all latest models: - [Hunyuan](https://huggingface.co/Tencent/HunyuanVideo): *HunyuanVideo, FastHunyuan, SkyReels* | *T2V, I2V* - - [WAN21](https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers): *1.3B, 13B* | *T2V, I2V* + - [WAN21](https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers): *1.3B, 14B* | *T2V, I2V* - [LTXVideo](https://huggingface.co/Lightricks/LTX-Video): *0.9.0, 0.9.1, 0.9.5* | *T2V, I2V* - [CogVideoX](https://huggingface.co/THUDM/CogVideoX-5b): *2B, 5B* | *T2V, I2V* - [Allegro](https://huggingface.co/rhymes-ai/Allegro): *T2V* diff --git a/modules/ui_video.py b/modules/ui_video.py index 905cdfb95..81f68c50a 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -50,6 +50,9 @@ def run_video(*args): elif selected and 'Allegro' in selected.name: from modules.video_models import run_allegro return run_allegro.generate(*args) + elif selected and 'WAN' in selected.name: + from modules.video_models import run_wan + return run_wan.generate(*args) shared.log.error(f'Video model not found: args={args}') return [], None, '', '', f'Video model not found: engine={engine} model={model}' diff --git a/modules/video_models/run_wan.py b/modules/video_models/run_wan.py new file mode 100644 index 000000000..7868aa020 --- /dev/null +++ b/modules/video_models/run_wan.py @@ -0,0 +1,94 @@ +import os +import time +from modules import shared, errors, sd_models, processing, devices, images, ui_common +from modules.video_models import models_def, video_utils, video_load, video_vae + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def generate(*args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args + if engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + found = [model.name for model in models_def.models.get(engine, [])] + selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not shared.sd_loaded or 'WAN' not in shared.sd_model.__class__.__name__: + video_load.load_model(selected) + if not shared.sd_loaded or 'WAN' not in shared.sd_model.__class__.__name__: + return video_utils.queue_err('model not loaded') + debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') + + p = processing.StableDiffusionProcessingVideo( + sd_model=shared.sd_model, + prompt=prompt, + negative_prompt=negative, + styles=styles, + seed=int(seed), + sampler_name = processing.get_sampler_name(sampler_index), + sampler_shift=float(sampler_shift), + steps=int(steps), + width=16 * int(width // 16), + height=16 * int(height // 16), + frames=int(frames), + init_image=init_image, + cfg_scale=float(guidance_scale), + diffusers_guidance_rescale=float(guidance_true), + vae_type=vae_type, + vae_tile_frames=int(vae_tile_frames), + override_settings=override_settings, + ) + if p.vae_type == 'Remote' and not selected.vae_remote: + shared.log.warning(f'Video: model={selected.name} remote vae not supported') + p.vae_type = 'Default' + p.scripts = None + p.script_args = None + p.state = ui_state + p.do_not_save_grid = True + p.do_not_save_samples = not save_frames + if 'I2V' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + + # cleanup memory + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + # set args + processing.fix_seed(p) + video_vae.set_vae_params(p) + video_utils.set_prompt(p) + p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil' + p.ops.append('video') + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift + shared.opts.data['schedulers_shift'] = sampler_shift + debug(f'Video: task_args={p.task_args}') + + # run processing + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') + err = None + t0 = time.time() + try: + processed = processing.process_images(p) + except Exception as e: + err = str(e) + errors.display(e, 'video') + t1 = time.time() + shared.state.disable_preview = False + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + p.close() + + # done + if err: + return video_utils.queue_err(err) + if processed is None or len(processed.images) == 0: + return video_utils.queue_err('processing failed') + shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') + video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + generation_info_js = processed.js() if processed is not None else '' + return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 2c8634f1e..55016cfc9 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -10,10 +10,10 @@ debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None e def load_model(selected: models_def.Model): if selected is None: - return + return '' global loaded_model # pylint: disable=global-statement if loaded_model == selected.name: - return + return '' sd_models.unload_model_weights() t0 = time.time() From 2a077414bcf93969592b6768e9de7a183e1eaa3b Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Fri, 21 Mar 2025 14:56:58 +0900 Subject: [PATCH 038/122] zluda experimental torch.compile --- modules/zluda.py | 4 +-- modules/zluda_hijacks.py | 14 ++++++++-- modules/zluda_installer.py | 54 ++++++++++++++++++++++++++++++++------ 3 files changed, 60 insertions(+), 12 deletions(-) diff --git a/modules/zluda.py b/modules/zluda.py index 0203a6398..431ab2c8c 100644 --- a/modules/zluda.py +++ b/modules/zluda.py @@ -3,7 +3,8 @@ from typing import Union import torch from torch._prims_common import DeviceLikeType import onnxruntime as ort -from modules import shared, devices +from modules import shared, devices, zluda_installer +from modules.zluda_installer import core, default_agent # pylint: disable=unused-import from modules.onnx_impl.execution_providers import available_execution_providers, ExecutionProvider @@ -32,7 +33,6 @@ def initialize_zluda(): from modules.zluda_hijacks import do_hijack do_hijack() - from modules import zluda_installer torch.backends.cudnn.enabled = zluda_installer.MIOpen_available if not zluda_installer.MIOpen_available: torch.backends.cuda.enable_cudnn_sdp(False) diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index 9f622be94..1b3f750a9 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -1,5 +1,6 @@ import torch -from modules import rocm +import torch._dynamo.device_interface +from modules import rocm, zluda _topk = torch.topk @@ -10,7 +11,7 @@ def topk(input: torch.Tensor, *args, **kwargs): # pylint: disable=redefined-buil class DeviceProperties: - PROPERTIES_OVERRIDE = {"regs_per_multiprocessor": 65535} + PROPERTIES_OVERRIDE = {"regs_per_multiprocessor": 65535, "gcnArchName": "UNKNOWN ARCHITECTURE"} internal: torch._C._CudaDeviceProperties def __init__(self, props: torch._C._CudaDeviceProperties): @@ -27,11 +28,20 @@ def torch_cuda__get_device_properties(device): return DeviceProperties(__get_device_properties(device)) +_cuda_getCurrentRawStream = torch._C._cuda_getCurrentRawStream # pylint: disable=protected-access +def torch__C__cuda_getCurrentRawStream(device): + return zluda.core.to_hip_stream(_cuda_getCurrentRawStream(device)) + + def do_hijack(): torch.version.hip = rocm.version torch.topk = topk + if zluda.default_agent is not None: + DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = zluda.default_agent.name torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access + torch._C._cuda_getCurrentRawStream = torch__C__cuda_getCurrentRawStream # pylint: disable=protected-access + torch._dynamo.device_interface.CudaInterface.get_raw_stream = staticmethod(torch__C__cuda_getCurrentRawStream) # pylint: disable=protected-access try: import triton _get_device_properties = triton.runtime.driver.active.utils.get_device_properties diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 4a0b9ffa4..35a84cfd2 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -17,7 +17,6 @@ DLL_MAPPING = { 'nvrtc.dll': 'nvrtc64_112_0.dll', } HIPSDK_TARGETS = ['rocblas.dll', 'rocsolver.dll', 'hipfft.dll',] -ZLUDA_TARGETS = ('nvcuda.dll', 'nvml.dll',) hipBLASLt_available = False MIOpen_available = False @@ -27,7 +26,49 @@ default_agent: Union[rocm.Agent, None] = None hipBLASLt_enabled = False nightly = os.environ.get("ZLUDA_NIGHTLY", "0") == "1" -skip_arch_test = os.environ.get("ZLUDA_SKIP_ARCH_TEST", "0") == "1" + + +class ZLUDAResult(ctypes.Structure): + _fields_ = [ + ('return_code', ctypes.c_int), + ('value', ctypes.c_ulonglong), + ] + + +class ZLUDALibrary: + internal: ctypes.WinDLL + + def __init__(self, internal: ctypes.WinDLL): + self.internal = internal + + +class Core(ZLUDALibrary): + internal: ctypes.WinDLL + + def __init__(self, internal: ctypes.WinDLL): + internal.zluda_get_hip_object.restype = ZLUDAResult + internal.zluda_get_hip_object.argtypes = [ctypes.c_void_p, ctypes.c_int] + + internal.zluda_get_nightly_flag.restype = ctypes.c_int + internal.zluda_get_nightly_flag.argtypes = [] + + super().__init__(internal) + + def to_hip_stream(self, zluda_object: ctypes.c_void_p): + return self.internal.zluda_get_hip_object(zluda_object, 1).value + + def get_nightly_flag(self) -> int: + return self.internal.zluda_get_nightly_flag().value + + +core = None +ml = None + + +def load_core_modules(): + global core, ml # pylint: disable=global-statement + core = Core(ctypes.windll.LoadLibrary(os.path.join(path, 'nvcuda.dll'))) + ml = ZLUDALibrary(ctypes.windll.LoadLibrary(os.path.join(path, 'nvml.dll'))) def set_default_agent(agent: rocm.Agent): @@ -36,10 +77,8 @@ def set_default_agent(agent: rocm.Agent): is_nightly = False try: - nvcuda = ctypes.windll.LoadLibrary(os.path.join(path, 'nvcuda.dll')) - nvcuda.zluda_get_nightly_flag.restype = ctypes.c_int - nvcuda.zluda_get_nightly_flag.argtypes = [] - is_nightly = nvcuda.zluda_get_nightly_flag() == 1 + load_core_modules() + is_nightly = core.get_nightly_flag() except Exception: pass @@ -113,10 +152,9 @@ def load() -> None: os.environ["ZLUDA_COMGR_LOG_LEVEL"] = "1" os.environ["ZLUDA_NVRTC_LIB"] = os.path.join([v for v in site.getsitepackages() if v.endswith("site-packages")][0], "torch", "lib", "nvrtc64_112_0.dll") + load_core_modules() for v in HIPSDK_TARGETS: ctypes.windll.LoadLibrary(os.path.join(rocm.path, 'bin', v)) - for v in ZLUDA_TARGETS: - ctypes.windll.LoadLibrary(os.path.join(path, v)) for v in DLL_MAPPING.values(): ctypes.windll.LoadLibrary(os.path.join(path, v)) From e5f62501533736008cf954fcb79f6e25b9bd60d9 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Fri, 21 Mar 2025 17:42:35 +0900 Subject: [PATCH 039/122] fix zluda nightly --- modules/zluda_installer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 35a84cfd2..a280a215a 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -58,7 +58,7 @@ class Core(ZLUDALibrary): return self.internal.zluda_get_hip_object(zluda_object, 1).value def get_nightly_flag(self) -> int: - return self.internal.zluda_get_nightly_flag().value + return self.internal.zluda_get_nightly_flag() core = None From b2432db88e71d05cdfa6644a9bb7e233b11bb9ae Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Mar 2025 09:39:00 -0400 Subject: [PATCH 040/122] fix wan and add latte Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 3 ++ installer.py | 2 +- modules/ui_video.py | 44 ++++++++++++++++-------------- modules/video_models/models_def.py | 9 ++++++ modules/video_models/run_wan.py | 4 +-- modules/video_models/video_load.py | 3 +- wiki | 2 +- 7 files changed, 42 insertions(+), 25 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1d099df92..6bc01e323 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,8 @@ - Video: HunyuanVideo-I2V incompatible with latest transformers - Video: LTXVideo-095 support for conditioned input - Video: LTXVideo-095 support for offloading + - Video: FasterCache: https://github.com/huggingface/diffusers/pull/10163 + - Video: PyramidAttention: https://github.com/huggingface/diffusers/pull/9562 ### Highlights for 2025-03-20 @@ -31,6 +33,7 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio - [CogVideoX](https://huggingface.co/THUDM/CogVideoX-5b): *2B, 5B* | *T2V, I2V* - [Allegro](https://huggingface.co/rhymes-ai/Allegro): *T2V* - [Mochi1](https://huggingface.co/genmo/mochi-1-preview): *T2V* + - [Latte1](https://huggingface.co/maxin-cn/Latte-1): *T2V - decoding: - **Default**: use vae from model - **Tiny VAE**: support for *Hunyuan, WAN, Mochi* diff --git a/installer.py b/installer.py index 0411ef496..cacb0a62d 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git: return - sha = '56f740051dae2d410677292a5c9e5b66e60f87dc' # diffusers commit hash + sha = '844221ae4e20a8939ee052f75874e284f75d4c5c' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/ui_video.py b/modules/ui_video.py index 81f68c50a..20edfc89e 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -26,7 +26,6 @@ def model_change(engine, model): sd_models.unload_model_weights() msg = 'Video model unloaded' return [msg, - gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False), video_utils.get_url(selected.url if selected else None), ] @@ -84,24 +83,29 @@ def create_ui(): seed = gr.Number(label='Initial seed', value=-1, elem_id="video_seed", container=True) random_seed = ToolButton(ui_symbols.random, elem_id="video_random_seed", label='Random seed') reuse_seed = ToolButton(ui_symbols.reuse, elem_id="video_reuse_seed", label='Reuse seed') - steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") - with gr.Row(): - sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") - dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) # TODO video: dynamic shift - with gr.Row(): - guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") - guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") - with gr.Row(): - vae_type = gr.Dropdown(label='VAE decode', choices=['Default', 'Tiny', 'Remote'], value='Default', elem_id="video_vae_type") - vae_tile_frames = gr.Slider(label='Tile frames', minimum=1, maximum=64, step=1, value=16, elem_id="video_vae_tile_frames") - with gr.Row(): - with gr.Group(visible=False, elem_id='video_init_image') as image_group: - gr.HTML("
  Init image") - init_image = gr.Image(elem_id="video_image", show_label=False, type="pil", image_mode="RGB", height=512) - with gr.Row(): - save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") - with gr.Row(): - video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video') + with gr.Accordion(open=True, label="Parameters", elem_id='video_parameters_accordion'): + steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") + with gr.Row(): + sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") + dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) # TODO video: dynamic shift + with gr.Row(): + guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") + guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") + with gr.Accordion(open=True, label="Decode", elem_id='video_decode_accordion'): + with gr.Row(): + vae_type = gr.Dropdown(label='VAE decode', choices=['Default', 'Tiny', 'Remote'], value='Default', elem_id="video_vae_type") + vae_tile_frames = gr.Slider(label='Tile frames', minimum=1, maximum=64, step=1, value=16, elem_id="video_vae_tile_frames") + with gr.Accordion(open=False, label="Init image", elem_id='video_init_accordion'): + gr.HTML("
  Init image") + init_image = gr.Image(elem_id="video_image", show_label=False, type="pil", image_mode="RGB", height=512) + with gr.Accordion(open=False, label="Accelerate", elem_id='video_accelerate_accordion'): + faster_cache = gr.Checkbox(label='FasterCache', value=False, elem_id="video_faster_cache") + pyramid_attention = gr.Checkbox(label='PyramidAttention', value=False, elem_id="video_pyramid_attention") + with gr.Accordion(open=True, label="Output", elem_id='video_output_accordion'): + with gr.Row(): + save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") + with gr.Row(): + video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video') override_settings = ui_common.create_override_inputs('video') # output panel with gallery and video tabs @@ -117,7 +121,7 @@ def create_ui(): random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) # handle engine and model change engine.change(fn=engine_change, inputs=[engine], outputs=[model]) - model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, image_group, url]) + model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, url]) # setup extra networks ui_extra_networks.setup_ui(extra_networks_ui, gallery) diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 5ff58fd82..4395eb817 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -139,6 +139,15 @@ models = { te_cls=transformers.T5EncoderModel, dit_cls=diffusers.MochiTransformer3DModel), ], + 'Latte Video': [ + Model(name='None'), + Model(name='Latte 1 T2V', + url='https://huggingface.co/maxin-cn/Latte-1', + repo='maxin-cn/Latte-1', + repo_cls=diffusers.LattePipeline, + te_cls=transformers.T5EncoderModel, + dit_cls=diffusers.LatteTransformer3DModel), + ], 'Allegro Video': [ Model(name='None'), Model(name='Allegro T2V', diff --git a/modules/video_models/run_wan.py b/modules/video_models/run_wan.py index 7868aa020..fd50c1fc5 100644 --- a/modules/video_models/run_wan.py +++ b/modules/video_models/run_wan.py @@ -13,9 +13,9 @@ def generate(*args, **kwargs): return video_utils.queue_err('model not selected') found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'WAN' not in shared.sd_model.__class__.__name__: + if not shared.sd_loaded or 'Wan' not in shared.sd_model.__class__.__name__: video_load.load_model(selected) - if not shared.sd_loaded or 'WAN' not in shared.sd_model.__class__.__name__: + if not shared.sd_loaded or 'Wan' not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 55016cfc9..f380a8d88 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -73,7 +73,8 @@ def load_model(selected: models_def.Model): if selected.te_hijack: shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.encode_prompt = video_utils.hijack_encode_prompt - shared.sd_model.vae.enable_slicing() + if hasattr(shared.sd_model.vae, 'enable_slicing'): + shared.sd_model.vae.enable_slicing() loaded_model = selected.name msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}' shared.log.info(msg) diff --git a/wiki b/wiki index d50882dcb..3f46b4f74 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit d50882dcb83a1441591b0d491efd143e12c1930a +Subproject commit 3f46b4f742e439dee1d012c9e5e2ddf2a6b79aa6 From 46bc0834b11cd9bfaa2b782fe61745de7564dc00 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Mar 2025 14:53:52 -0400 Subject: [PATCH 041/122] video tab major update Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 21 +++-- modules/processing_args.py | 15 +-- modules/processing_class.py | 6 +- modules/processing_info.py | 3 + modules/shared.py | 2 +- modules/ui_symbols.py | 1 + modules/ui_video.py | 84 ++++++++++------- modules/video_models/models_def.py | 31 +++++- modules/video_models/run_allegro.py | 91 ------------------ modules/video_models/run_cog.py | 91 ------------------ modules/video_models/run_hunyuan.py | 94 ------------------- modules/video_models/run_ltx.py | 91 ------------------ modules/video_models/run_mochi.py | 91 ------------------ modules/video_models/video_cache.py | 45 +++++++++ .../video_models/{run_wan.py => video_run.py} | 11 ++- modules/video_models/video_utils.py | 2 +- 16 files changed, 163 insertions(+), 516 deletions(-) delete mode 100644 modules/video_models/run_allegro.py delete mode 100644 modules/video_models/run_cog.py delete mode 100644 modules/video_models/run_hunyuan.py delete mode 100644 modules/video_models/run_ltx.py delete mode 100644 modules/video_models/run_mochi.py create mode 100644 modules/video_models/video_cache.py rename modules/video_models/{run_wan.py => video_run.py} (90%) diff --git a/CHANGELOG.md b/CHANGELOG.md index 6bc01e323..5bcc50a84 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2025-03-19 +## Update for 2025-03-21 ### ToDo/Limitations @@ -12,15 +12,19 @@ - Video: HunyuanVideo-I2V incompatible with latest transformers - Video: LTXVideo-095 support for conditioned input - Video: LTXVideo-095 support for offloading - - Video: FasterCache: https://github.com/huggingface/diffusers/pull/10163 - - Video: PyramidAttention: https://github.com/huggingface/diffusers/pull/9562 + - Video: FasterCache and PyramidAttentionBroadcast granular config + - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN + - Video: HunyuanVideo-I2V-16ch + - Video: CogVideo-15 support -### Highlights for 2025-03-20 +### Highlights for 2025-03-21 -Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1** and more! -Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides. +Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** +And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! -### Details for 2025-03-20 +Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides + +### Details for 2025-03-21 - **Video tab** - initial release so consider this as alpha version @@ -39,6 +43,9 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio - **Tiny VAE**: support for *Hunyuan, WAN, Mochi* - **Remote VAE**: support for *Hunyuan* - **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog* + - acceleration: + - [FasterCache](https://huggingface.co/papers/2410.19355): support for *Hunyuan, Mochi, Latte, Allegro, Cog* + - [PyramidAttentionBroadcast](https://huggingface.co/papers/2408.12588): support for *Hunyuan, Mochi, Latte, Allegro, Cog* - additional key points: - all models are auto-downloaded upon first use uses *system paths -> huggingface* folder diff --git a/modules/processing_args.py b/modules/processing_args.py index 261a996b0..a5f04eb7d 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -26,7 +26,7 @@ def task_specific_kwargs(p, model): p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images] if isinstance(p.init_images[0], Image.Image): p.init_images = [i.convert('RGB') if i.mode != 'RGB' else i for i in p.init_images if i is not None] - if (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE or len(getattr(p, 'init_images', [])) == 0) and not is_img2img_model: + if (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE or len(getattr(p, 'init_images', [])) == 0) and not is_img2img_model and 'video' not in p.ops: p.ops.append('txt2img') if hasattr(p, 'width') and hasattr(p, 'height'): task_args = { @@ -238,13 +238,13 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t if hasattr(model, 'scheduler') and hasattr(model.scheduler, 'noise_sampler_seed') and hasattr(model.scheduler, 'noise_sampler'): model.scheduler.noise_sampler = None # noise needs to be reset instead of using cached values model.scheduler.noise_sampler_seed = p.seeds # some schedulers have internal noise generator and do not use pipeline generator - if 'seed' in possible: + if 'seed' in possible and p.seed is not None: args['seed'] = p.seed - if 'noise_sampler_seed' in possible: + if 'noise_sampler_seed' in possible and p.seeds is not None: args['noise_sampler_seed'] = p.seeds - if 'guidance_scale' in possible: + if 'guidance_scale' in possible and p.cfg_scale is not None and p.cfg_scale > 0: args['guidance_scale'] = p.cfg_scale - if 'img_guidance_scale' in possible and hasattr(p, 'image_cfg_scale'): + if 'img_guidance_scale' in possible and hasattr(p, 'image_cfg_scale') and p.image_cfg_scale is not None and p.image_cfg_scale > 0: args['img_guidance_scale'] = p.image_cfg_scale if 'generator' in possible: generator = get_generator(p) @@ -304,9 +304,10 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t # handle remaining args for arg in kwargs: if arg in possible: # add kwargs + if type(kwargs[arg]) == float or type(kwargs[arg]) == int: + if kwargs[arg] <= -1: # skip -1 as default value + continue args[arg] = kwargs[arg] - else: - pass task_kwargs = task_specific_kwargs(p, model) for arg in task_kwargs: diff --git a/modules/processing_class.py b/modules/processing_class.py index a0d2104b5..eb4e333ec 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -236,9 +236,13 @@ class StableDiffusionProcessing: self.height = firstphase_height self.sampler_name = sampler_name or processing_helpers.get_sampler_name(sampler_index, img=True) self.hr_sampler_name: str = hr_sampler_name if hr_sampler_name != 'Same as primary' else self.sampler_name - self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts} self.inpaint_full_res = inpaint_full_res if isinstance(inpaint_full_res, bool) else self.inpaint_full_res self.inpaint_full_res = inpaint_full_res != 0 if isinstance(inpaint_full_res, int) else self.inpaint_full_res + try: + self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts} + except Exception as e: + shared.log.error(f'Override: {override_settings} {e}') + self.override_settings = {} # null items initialized later self.prompts = None diff --git a/modules/processing_info.py b/modules/processing_info.py index 3b0e1bd24..4db8ccaa5 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -186,6 +186,9 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No for k, v in args.copy().items(): if v is None: del args[k] + if type(v) is float or type(v) is int: + if v <= -1: + del args[k] if isinstance(v, str): if len(v) == 0 or v == '0x0': del args[k] diff --git a/modules/shared.py b/modules/shared.py index 8b2b2b9b0..a4a448966 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -779,7 +779,7 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"), 'schedulers_beta_end': OptionInfo(0, "Beta end", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.00001, "visible": native}), 'schedulers_timesteps_range': OptionInfo(1000, "Timesteps range", gr.Slider, {"minimum": 250, "maximum": 4000, "step": 1, "visible": native}), 'schedulers_shift': OptionInfo(3, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": False}), - 'schedulers_dynamic_shift': OptionInfo(True, "Sampler dynamic shift", gr.Checkbox, {"visible": False}), + 'schedulers_dynamic_shift': OptionInfo(False, "Sampler dynamic shift", gr.Checkbox, {"visible": False}), # managed from ui.py for backend original k-diffusion "always_batch_cond_uncond": OptionInfo(False, "Disable conditional batching", gr.Checkbox, {"visible": not native}), diff --git a/modules/ui_symbols.py b/modules/ui_symbols.py index c92ddb8f6..ef426e7e7 100644 --- a/modules/ui_symbols.py +++ b/modules/ui_symbols.py @@ -20,6 +20,7 @@ reuse = '♻️' info = 'ℹ' # noqa reset = '🔄' upload = '⬆️' +loading = '↺' reuse = '⬅️' search = '🔍' preview = '🖼️' diff --git a/modules/ui_video.py b/modules/ui_video.py index 20edfc89e..5044d0c36 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -1,59 +1,69 @@ +import os import gradio as gr from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste from modules.ui_components import ToolButton -from modules.video_models import models_def, video_utils, video_load +from modules.video_models import models_def, video_utils + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None def engine_change(engine): + debug(f'Video change: engine="{engine}"') found = [model.name for model in models_def.models.get(engine, [])] return gr.update(choices=found, value=found[0] if len(found) > 0 else None) def model_change(engine, model): + debug(f'Video change: engine="{engine}" model="{model}"') found = [model.name for model in models_def.models.get(engine, [])] selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - yield ['Video model loading', - gr.update(visible='I2V' in selected.name) if selected else gr.update(visible=False), - video_utils.get_url(selected.url if selected else None), - ] + return video_utils.get_url(selected.url if selected else None) + + +def model_load(engine, model): + debug(f'Video load: engine="{engine}" model="{model}"') + found = [model.name for model in models_def.models.get(engine, [])] + selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + yield f'Video model loading: {selected.name}' if selected: if 'None' in selected.name: sd_models.unload_model_weights() msg = 'Video model unloaded' else: + from modules.video_models import video_load msg = video_load.load_model(selected) else: sd_models.unload_model_weights() msg = 'Video model unloaded' - return [msg, - video_utils.get_url(selected.url if selected else None), - ] + yield msg + return msg def run_video(*args): engine, model = args[2], args[3] + debug(f'Video run: engine="{engine}" model="{model}"') found = [model.name for model in models_def.models.get(engine, [])] selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + if not selected or engine is None or model is None or engine == 'None' or model == 'None': + return video_utils.queue_err('model not selected') + debug(f'Video run: {str(selected)}') + from modules.video_models import video_run if selected and 'Hunyuan' in selected.name: - from modules.video_models import run_hunyuan - return run_hunyuan.generate(*args) + return video_run.generate('Hunyuan', *args) elif selected and 'LTX' in selected.name: - from modules.video_models import run_ltx - return run_ltx.generate(*args) + return video_run.generate('LTX', *args) elif selected and 'Mochi' in selected.name: - from modules.video_models import run_mochi - return run_mochi.generate(*args) + return video_run.generate('Mochi', *args) elif selected and 'Cog' in selected.name: - from modules.video_models import run_cog - return run_cog.generate(*args) + return video_run.generate('Cog', *args) elif selected and 'Allegro' in selected.name: - from modules.video_models import run_allegro - return run_allegro.generate(*args) + return video_run.generate('Allegro', *args) elif selected and 'WAN' in selected.name: - from modules.video_models import run_wan - return run_wan.generate(*args) - shared.log.error(f'Video model not found: args={args}') - return [], None, '', '', f'Video model not found: engine={engine} model={model}' + return video_run.generate('Wan', *args) + elif selected and 'Latte' in selected.name: + return video_run.generate('Latte', *args) + return video_utils.queue_err(f'model not found: engine="{engine}" model="{model}"') def create_ui(): @@ -74,23 +84,25 @@ def create_ui(): with gr.Row(): engine = gr.Dropdown(label='Engine', choices=list(models_def.models), value='None', elem_id="video_engine") model = gr.Dropdown(label='Model', choices=[''], value=None, elem_id="video_model") + btn_load = ToolButton(ui_symbols.loading, elem_id="video_model_load", label='Load model') with gr.Row(): - url = gr.HTML(label='Model URL', elem_id='video_model_url', value='') - with gr.Row(): - width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) - with gr.Row(): - frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15, elem_id="video_frames") - seed = gr.Number(label='Initial seed', value=-1, elem_id="video_seed", container=True) - random_seed = ToolButton(ui_symbols.random, elem_id="video_random_seed", label='Random seed') - reuse_seed = ToolButton(ui_symbols.reuse, elem_id="video_reuse_seed", label='Reuse seed') + url = gr.HTML(label='Model URL', elem_id='video_model_url', value='

') + with gr.Accordion(open=True, label="Size", elem_id='video_size_accordion'): + with gr.Row(): + width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) + with gr.Row(): + frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15, elem_id="video_frames") + seed = gr.Number(label='Initial seed', value=-1, elem_id="video_seed", container=True) + random_seed = ToolButton(ui_symbols.random, elem_id="video_random_seed", label='Random seed') + reuse_seed = ToolButton(ui_symbols.reuse, elem_id="video_reuse_seed", label='Reuse seed') with gr.Accordion(open=True, label="Parameters", elem_id='video_parameters_accordion'): steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video") with gr.Row(): - sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift") - dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) # TODO video: dynamic shift + sampler_shift = gr.Slider(label='Sampler shift', minimum=-1.0, maximum=20.0, step=0.1, value=-1.0, elem_id="video_scheduler_shift") + dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift") with gr.Row(): - guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale") - guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true") + guidance_scale = gr.Slider(label='Guidance scale', minimum=-1.0, maximum=14.0, step=0.1, value=-1.0, elem_id="video_guidance_scale") + guidance_true = gr.Slider(label='True guidance', minimum=-1.0, maximum=14.0, step=0.1, value=-1.0, elem_id="video_guidance_true") with gr.Accordion(open=True, label="Decode", elem_id='video_decode_accordion'): with gr.Row(): vae_type = gr.Dropdown(label='VAE decode', choices=['Default', 'Tiny', 'Remote'], value='Default', elem_id="video_vae_type") @@ -121,7 +133,8 @@ def create_ui(): random_seed.click(fn=lambda: -1, show_progress=False, inputs=[], outputs=[seed]) # handle engine and model change engine.change(fn=engine_change, inputs=[engine], outputs=[model]) - model.change(fn=model_change, inputs=[engine, model], outputs=[html_log, url]) + model.change(fn=model_change, inputs=[engine, model], outputs=[url]) + btn_load.click(fn=model_load, inputs=[engine, model], outputs=[html_log]) # setup extra networks ui_extra_networks.setup_ui(extra_networks_ui, gallery) @@ -154,6 +167,7 @@ def create_ui(): vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, + faster_cache, pyramid_attention, override_settings, ] # generate function diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 4395eb817..271318b69 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -3,6 +3,32 @@ import diffusers import transformers +""" +Hunyuan Video T2V: pass/pass/pass +Hunyuan Video I2V: pass/pass/pass, transformers incompatibility +SkyReels Hunyuan T2V: +SkyReels Hunyuan I2V: +Fast Hunyuan T2V: +LTXVideo 0.9.5 T2V: +LTXVideo 0.9.5 I2V: +LTXVideo 0.9.1 T2V: +LTXVideo 0.9.1 I2V: +LTXVideo 0.9.0 T2V: +LTXVideo 0.9.0 I2V: +WAN 2.1 1.3B T2V: pass/pass/pass +WAN 2.1 14B T2V: pass/fail/fail, error loading shard +WAN 2.1 14B I2V 480p: +WAN 2.1 14B I2V 720p: +Mochi 1 T2V: pass/pass/pass +Latte 1 T2V: pass/fail/fail, float vs bfloat during generate +Allegro T2V: pass/pass/fail, output is pure gray +CogVideoX 1.0 2B T2V: pass/pass/pass +CogVideoX 1.0 5B T2V: +CogVideoX 1.0 5B I2V: +CogVideoX 1.5 5B T2V: pass/pass/fail, output is pure black +CogVideoX 1.5 5B I2V: pass/pass/pass +""" + @dataclass class Model(): name: str @@ -19,6 +45,9 @@ class Model(): vae_hijack: bool = True vae_remote: bool = False + def __str__(self): + return f'name="{self.name}" url="{self.url}" repo="{self.repo}" repo_cls="{self.repo_cls}" dit="{self.dit}" dit_cls="{self.dit_cls}" dit_folder="{self.dit_folder}" te="{self.te}" te_cls="{self.te_cls}" te_folder="{self.te_folder}" te_hijack={self.te_hijack} vae_hijack={self.vae_hijack} vae_remote={self.vae_remote}' + models = { 'None': [], @@ -167,7 +196,7 @@ models = { dit_cls=diffusers.CogVideoXTransformer3DModel), Model(name='CogVideoX 1.0 5B T2V', url='https://huggingface.co/THUDM/CogVideoX-5b', - repo='THUDM/THUDM/CogVideoX-5b', + repo='THUDM/CogVideoX-5b', repo_cls=diffusers.CogVideoXPipeline, te_cls=transformers.T5EncoderModel, dit_cls=diffusers.CogVideoXTransformer3DModel), diff --git a/modules/video_models/run_allegro.py b/modules/video_models/run_allegro.py deleted file mode 100644 index 091be7333..000000000 --- a/modules/video_models/run_allegro.py +++ /dev/null @@ -1,91 +0,0 @@ -import os -import time -from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args - if engine is None or model is None or engine == 'None' or model == 'None': - return video_utils.queue_err('model not selected') - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: - video_load.load_model(selected) - if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__: - return video_utils.queue_err('model not loaded') - debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') - - p = processing.StableDiffusionProcessingVideo( - sd_model=shared.sd_model, - prompt=prompt, - negative_prompt=negative, - styles=styles, - seed=int(seed), - sampler_name = processing.get_sampler_name(sampler_index), - sampler_shift=float(sampler_shift), - steps=int(steps), - width=8 * int(width // 8), - height=8 * int(height // 8), - frames=int(frames), - init_image=init_image, - cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), - vae_type=vae_type, - vae_tile_frames=int(vae_tile_frames), - override_settings=override_settings, - ) - p.scripts = None - p.script_args = None - p.state = ui_state - p.do_not_save_grid = True - p.do_not_save_samples = not save_frames - if 'I2V' in model: - if init_image is None: - return video_utils.queue_err('init image not set') - p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - - # cleanup memory - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - # set args - processing.fix_seed(p) - video_vae.set_vae_params(p) - video_utils.set_prompt(p) - p.task_args['output_type'] = 'pil' - p.ops.append('video') - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - debug(f'Video: task_args={p.task_args}') - - # run processing - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') - err = None - t0 = time.time() - try: - processed = processing.process_images(p) - except Exception as e: - err = str(e) - errors.display(e, 'video') - t1 = time.time() - shared.state.disable_preview = False - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - p.close() - - # done - if err: - return video_utils.queue_err(err) - if processed is None or len(processed.images) == 0: - return video_utils.queue_err('processing failed') - shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - generation_info_js = processed.js() if processed is not None else '' - return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_cog.py b/modules/video_models/run_cog.py deleted file mode 100644 index 0aaee4b97..000000000 --- a/modules/video_models/run_cog.py +++ /dev/null @@ -1,91 +0,0 @@ -import os -import time -from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args - if engine is None or model is None or engine == 'None' or model == 'None': - return video_utils.queue_err('model not selected') - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: - video_load.load_model(selected) - if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__: - return video_utils.queue_err('model not loaded') - debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') - - p = processing.StableDiffusionProcessingVideo( - sd_model=shared.sd_model, - prompt=prompt, - negative_prompt=negative, - styles=styles, - seed=int(seed), - sampler_name = processing.get_sampler_name(sampler_index), - sampler_shift=float(sampler_shift), - steps=int(steps), - width=8 * int(width // 8), - height=8 * int(height // 8), - frames=int(frames), - init_image=init_image, - cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), - vae_type=vae_type, - vae_tile_frames=int(vae_tile_frames), - override_settings=override_settings, - ) - p.scripts = None - p.script_args = None - p.state = ui_state - p.do_not_save_grid = True - p.do_not_save_samples = not save_frames - if 'I2V' in model: - if init_image is None: - return video_utils.queue_err('init image not set') - p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - - # cleanup memory - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - # set args - processing.fix_seed(p) - video_vae.set_vae_params(p) - video_utils.set_prompt(p) - p.task_args['output_type'] = 'pil' - p.ops.append('video') - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - debug(f'Video: task_args={p.task_args}') - - # run processing - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') - err = None - t0 = time.time() - try: - processed = processing.process_images(p) - except Exception as e: - err = str(e) - errors.display(e, 'video') - t1 = time.time() - shared.state.disable_preview = False - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - p.close() - - # done - if err: - return video_utils.queue_err(err) - if processed is None or len(processed.images) == 0: - return video_utils.queue_err('processing failed') - shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - generation_info_js = processed.js() if processed is not None else '' - return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_hunyuan.py b/modules/video_models/run_hunyuan.py deleted file mode 100644 index 0f722b890..000000000 --- a/modules/video_models/run_hunyuan.py +++ /dev/null @@ -1,94 +0,0 @@ -import os -import time -from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args - if engine is None or model is None or engine == 'None' or model == 'None': - return video_utils.queue_err('model not selected') - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - video_load.load_model(selected) - if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__: - return video_utils.queue_err('model not loaded') - debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') - - p = processing.StableDiffusionProcessingVideo( - sd_model=shared.sd_model, - prompt=prompt, - negative_prompt=negative, - styles=styles, - seed=int(seed), - sampler_name = processing.get_sampler_name(sampler_index), - sampler_shift=float(sampler_shift), - steps=int(steps), - width=16 * int(width // 16), - height=16 * int(height // 16), - frames=int(frames), - init_image=init_image, - cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), - vae_type=vae_type, - vae_tile_frames=int(vae_tile_frames), - override_settings=override_settings, - ) - if p.vae_type == 'Remote' and not selected.vae_remote: - shared.log.warning(f'Video: model={selected.name} remote vae not supported') - p.vae_type = 'Default' - p.scripts = None - p.script_args = None - p.state = ui_state - p.do_not_save_grid = True - p.do_not_save_samples = not save_frames - if 'I2V' in model: - if init_image is None: - return video_utils.queue_err('init image not set') - p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - - # cleanup memory - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - # set args - processing.fix_seed(p) - video_vae.set_vae_params(p) - video_utils.set_prompt(p) - p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil' - p.ops.append('video') - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - debug(f'Video: task_args={p.task_args}') - - # run processing - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') - err = None - t0 = time.time() - try: - processed = processing.process_images(p) - except Exception as e: - err = str(e) - errors.display(e, 'video') - t1 = time.time() - shared.state.disable_preview = False - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - p.close() - - # done - if err: - return video_utils.queue_err(err) - if processed is None or len(processed.images) == 0: - return video_utils.queue_err('processing failed') - shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - generation_info_js = processed.js() if processed is not None else '' - return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_ltx.py b/modules/video_models/run_ltx.py deleted file mode 100644 index ef3fa08b8..000000000 --- a/modules/video_models/run_ltx.py +++ /dev/null @@ -1,91 +0,0 @@ -import os -import time -from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args - if engine is None or model is None or engine == 'None' or model == 'None': - return video_utils.queue_err('model not selected') - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: - video_load.load_model(selected) - if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__: - return video_utils.queue_err('model not loaded') - debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') - - p = processing.StableDiffusionProcessingVideo( - sd_model=shared.sd_model, - prompt=prompt, - negative_prompt=negative, - styles=styles, - seed=int(seed), - sampler_name = processing.get_sampler_name(sampler_index), - sampler_shift=float(sampler_shift), - steps=int(steps), - width=32 * int(width // 32), - height=32 * int(height // 32), - frames=int(frames), - init_image=init_image, - cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), - vae_type=vae_type, - vae_tile_frames=int(vae_tile_frames), - override_settings=override_settings, - ) - p.scripts = None - p.script_args = None - p.state = ui_state - p.do_not_save_grid = True - p.do_not_save_samples = not save_frames - if 'I2V' in model: - if init_image is None: - return video_utils.queue_err('init image not set') - p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - - # cleanup memory - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - # set args - processing.fix_seed(p) - video_vae.set_vae_params(p) - video_utils.set_prompt(p) - p.task_args['output_type'] = 'pil' - p.ops.append('video') - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - debug(f'Video: task_args={p.task_args}') - - # run processing - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') - err = None - t0 = time.time() - try: - processed = processing.process_images(p) - except Exception as e: - err = str(e) - errors.display(e, 'video') - t1 = time.time() - shared.state.disable_preview = False - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - p.close() - - # done - if err: - return video_utils.queue_err(err) - if processed is None or len(processed.images) == 0: - return video_utils.queue_err('processing failed') - shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - generation_info_js = processed.js() if processed is not None else '' - return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/run_mochi.py b/modules/video_models/run_mochi.py deleted file mode 100644 index ef705880e..000000000 --- a/modules/video_models/run_mochi.py +++ /dev/null @@ -1,91 +0,0 @@ -import os -import time -from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args - if engine is None or model is None or engine == 'None' or model == 'None': - return video_utils.queue_err('model not selected') - found = [model.name for model in models_def.models.get(engine, [])] - selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: - video_load.load_model(selected) - if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__: - return video_utils.queue_err('model not loaded') - debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') - - p = processing.StableDiffusionProcessingVideo( - sd_model=shared.sd_model, - prompt=prompt, - negative_prompt=negative, - styles=styles, - seed=int(seed), - sampler_name = processing.get_sampler_name(sampler_index), - sampler_shift=float(sampler_shift), - steps=int(steps), - width=8 * int(width // 8), - height=8 * int(height // 8), - frames=int(frames), - init_image=init_image, - cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), - vae_type=vae_type, - vae_tile_frames=int(vae_tile_frames), - override_settings=override_settings, - ) - p.scripts = None - p.script_args = None - p.state = ui_state - p.do_not_save_grid = True - p.do_not_save_samples = not save_frames - if 'I2V' in model: - if init_image is None: - return video_utils.queue_err('init image not set') - p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - - # cleanup memory - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - # set args - processing.fix_seed(p) - video_vae.set_vae_params(p) - video_utils.set_prompt(p) - p.task_args['output_type'] = 'pil' - p.ops.append('video') - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - debug(f'Video: task_args={p.task_args}') - - # run processing - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') - err = None - t0 = time.time() - try: - processed = processing.process_images(p) - except Exception as e: - err = str(e) - errors.display(e, 'video') - t1 = time.time() - shared.state.disable_preview = False - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - p.close() - - # done - if err: - return video_utils.queue_err(err) - if processed is None or len(processed.images) == 0: - return video_utils.queue_err('processing failed') - shared.log.info(f'Video: name="{selected.name}" cls={shared.sd_model.__class__.__name__} frames={len(processed.images)} time={t1-t0:.2f}') - video_file = images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) - generation_info_js = processed.js() if processed is not None else '' - return processed.images, video_file, generation_info_js, processed.info, ui_common.plaintext_to_html(processed.comments) diff --git a/modules/video_models/video_cache.py b/modules/video_models/video_cache.py new file mode 100644 index 000000000..be43068d1 --- /dev/null +++ b/modules/video_models/video_cache.py @@ -0,0 +1,45 @@ +import os +import diffusers +from modules import shared, errors + + +debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def set_cache(faster_cache=False, pyramid_attention_broadcast=False): + if not shared.sd_loaded or not hasattr(shared.sd_model, 'transformer'): + return + if not hasattr(shared.sd_model.transformer, 'enable_cache'): + debug(f'Video cache: cls={shared.sd_model.transformer.__class__.__name__} not supported') + return + try: + if faster_cache: # https://github.com/huggingface/diffusers/pull/10163 + config = diffusers.FasterCacheConfig( + spatial_attention_block_skip_range=2, + spatial_attention_timestep_skip_range=(-1, 681), + current_timestep_callback=lambda: shared.sd_model.current_timestep, + attention_weight_callback=lambda _: 0.3, + unconditional_batch_skip_range=5, + unconditional_batch_timestep_skip_range=(-1, 781), + tensor_format="BFCHW", + ) + shared.sd_model.transformer.disable_cache() + shared.sd_model.transformer.enable_cache(config) + shared.log.debug(f'Video cache: type={config.__class__.__name__}') + debug(f'Video cache: {vars(config)}') + elif pyramid_attention_broadcast: # https://github.com/huggingface/diffusers/pull/9562 + config = diffusers.PyramidAttentionBroadcastConfig( + spatial_attention_block_skip_range=2, + spatial_attention_timestep_skip_range=(100, 800), + current_timestep_callback=lambda: shared.sd_model.current_timestep, + ) + shared.sd_model.transformer.disable_cache() + shared.sd_model.transformer.enable_cache(config) + shared.log.debug(f'Video cache: type={config.__class__.__name__}') + debug(f'Video cache: {vars(config)}') + else: + debug('Video cache: not enabled') + shared.sd_model.transformer.disable_cache() + except Exception as e: + shared.log.error(f'Video cache: error={e}') + errors.display(e, 'video cache') diff --git a/modules/video_models/run_wan.py b/modules/video_models/video_run.py similarity index 90% rename from modules/video_models/run_wan.py rename to modules/video_models/video_run.py index fd50c1fc5..09dee3826 100644 --- a/modules/video_models/run_wan.py +++ b/modules/video_models/video_run.py @@ -1,21 +1,21 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae +from modules.video_models import models_def, video_utils, video_load, video_vae, video_cache debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None -def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args +def generate(keyword, *args, **kwargs): + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, faster_cache, pyramid_attention, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or 'Wan' not in shared.sd_model.__class__.__name__: + if not shared.sd_loaded or keyword not in shared.sd_model.__class__.__name__: video_load.load_model(selected) - if not shared.sd_loaded or 'Wan' not in shared.sd_model.__class__.__name__: + if not shared.sd_loaded or keyword not in shared.sd_model.__class__.__name__: return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') @@ -58,6 +58,7 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) video_vae.set_vae_params(p) + video_cache.set_cache(faster_cache=faster_cache, pyramid_attention_broadcast=pyramid_attention) video_utils.set_prompt(p) p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil' p.ops.append('video') diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 0ed4791eb..8a69e19ec 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -18,7 +18,7 @@ def get_quant(args): def get_url(url): - return f'  {url}
' if url else '' + return f'  {url}

' if url else '

' def set_prompt(p): From 8930d1dcb073e509f1847ca462cfc9b95d22b1a6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Mar 2025 18:22:00 -0400 Subject: [PATCH 042/122] fix video-i2v Signed-off-by: Vladimir Mandic --- modules/processing_args.py | 3 ++ modules/video_models/models_def.py | 53 +++++++++++++++++------------- 2 files changed, 33 insertions(+), 23 deletions(-) diff --git a/modules/processing_args.py b/modules/processing_args.py index a5f04eb7d..db8794214 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -336,6 +336,9 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t if isinstance(args['image'], torch.Tensor) or isinstance(args['image'], np.ndarray): args['width'] = 8 * args['image'].shape[-1] args['height'] = 8 * args['image'].shape[-2] + elif isinstance(args['image'], Image.Image): + args['width'] = args['image'].width + args['height'] = args['image'].height elif isinstance(args['image'][0], torch.Tensor) or isinstance(args['image'][0], np.ndarray): args['width'] = 8 * args['image'][0].shape[-1] args['height'] = 8 * args['image'][0].shape[-2] diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 271318b69..0e55b31af 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -4,29 +4,35 @@ import transformers """ -Hunyuan Video T2V: pass/pass/pass -Hunyuan Video I2V: pass/pass/pass, transformers incompatibility -SkyReels Hunyuan T2V: -SkyReels Hunyuan I2V: -Fast Hunyuan T2V: -LTXVideo 0.9.5 T2V: -LTXVideo 0.9.5 I2V: -LTXVideo 0.9.1 T2V: -LTXVideo 0.9.1 I2V: -LTXVideo 0.9.0 T2V: -LTXVideo 0.9.0 I2V: -WAN 2.1 1.3B T2V: pass/pass/pass -WAN 2.1 14B T2V: pass/fail/fail, error loading shard -WAN 2.1 14B I2V 480p: -WAN 2.1 14B I2V 720p: -Mochi 1 T2V: pass/pass/pass -Latte 1 T2V: pass/fail/fail, float vs bfloat during generate -Allegro T2V: pass/pass/fail, output is pure gray -CogVideoX 1.0 2B T2V: pass/pass/pass -CogVideoX 1.0 5B T2V: -CogVideoX 1.0 5B I2V: -CogVideoX 1.5 5B T2V: pass/pass/fail, output is pure black -CogVideoX 1.5 5B I2V: pass/pass/pass +# Model tests: download/load/generate + +- Hunyuan Video T2V: pass/pass/pass +- Hunyuan Video I2V: pass/pass/fail, transformers incompatibility +- SkyReels Hunyuan T2V: pass/pass/pass +- SkyReels Hunyuan I2V: pass/pass/pass +- Fast Hunyuan T2V: pass/pass/pass + +- LTXVideo 0.9.5 T2V: pass/pass/fail, v095 pipeline is tbd +- LTXVideo 0.9.5 I2V: pass/pass/fail, v095 pipeline is tbd +- LTXVideo 0.9.1 T2V: pass/pass/pass +- LTXVideo 0.9.1 I2V: pass/pass/tbd +- LTXVideo 0.9.0 T2V: pass/pass/pass +- LTXVideo 0.9.0 I2V: pass/pass/tbd + +- WAN 2.1 1.3B T2V: pass/pass/pass +- WAN 2.1 14B T2V: pass/pass/fail, error loading shard +- WAN 2.1 14B I2V 480p: pass/pass/tbd +- WAN 2.1 14B I2V 720p: pass/pass/tbd + +- CogVideoX 1.0 2B T2V: pass/pass/pass +- CogVideoX 1.0 5B T2V: pass/pass/pass +- CogVideoX 1.0 5B I2V: pass/pass/pass +- CogVideoX 1.5 5B T2V: download/load/fail, v15 pipeline is tbd +- CogVideoX 1.5 5B I2V: download/load/fail, v15 pipeline is tbd + +- Mochi 1 T2V: pass/pass/pass +- Latte 1 T2V: pass/pass/fail, float vs bfloat during generate +- Allegro T2V: pass/pass/fail, output is pure gray """ @dataclass @@ -80,6 +86,7 @@ models = { url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-I2V', vae_remote=True, repo='hunyuanvideo-community/HunyuanVideo', + repo_cls=diffusers.HunyuanSkyreelsImageToVideoPipeline, te_cls=transformers.LlamaModel, dit='Skywork/SkyReels-V1-Hunyuan-I2V', dit_folder=None, From 990822c69c22d9484b7f62b9aa0e194f48a68e56 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Mar 2025 18:55:33 -0400 Subject: [PATCH 043/122] fix video model change logic Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 4 ++-- extensions-builtin/sdnext-modernui | 2 +- modules/ui_video.py | 14 +++++++------- modules/video_models/models_def.py | 2 +- modules/video_models/video_run.py | 11 ++++++++--- 5 files changed, 19 insertions(+), 14 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5bcc50a84..f98afce51 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,12 +6,12 @@ - VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - VAE Remote encode: SD15 and Flux.1 issues: - - Video: ModernUI support is TBD - Video: API support is TBD - Video: Wiki page is TBD - - Video: HunyuanVideo-I2V incompatible with latest transformers + - Video: LTXVideo-095 params - Video: LTXVideo-095 support for conditioned input - Video: LTXVideo-095 support for offloading + - Video: HunyuanVideo-I2V incompatible with latest transformers - Video: FasterCache and PyramidAttentionBroadcast granular config - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN - Video: HunyuanVideo-I2V-16ch diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 7fc52bb97..a9e015993 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 7fc52bb976783322bdf381e046ceb689261dbe2c +Subproject commit a9e015993194fe7b7c6222c9f32efbb557287a67 diff --git a/modules/ui_video.py b/modules/ui_video.py index 5044d0c36..cff8a54b0 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -50,19 +50,19 @@ def run_video(*args): debug(f'Video run: {str(selected)}') from modules.video_models import video_run if selected and 'Hunyuan' in selected.name: - return video_run.generate('Hunyuan', *args) + return video_run.generate(*args) elif selected and 'LTX' in selected.name: - return video_run.generate('LTX', *args) + return video_run.generate(*args) elif selected and 'Mochi' in selected.name: - return video_run.generate('Mochi', *args) + return video_run.generate(*args) elif selected and 'Cog' in selected.name: - return video_run.generate('Cog', *args) + return video_run.generate(*args) elif selected and 'Allegro' in selected.name: - return video_run.generate('Allegro', *args) + return video_run.generate(*args) elif selected and 'WAN' in selected.name: - return video_run.generate('Wan', *args) + return video_run.generate(*args) elif selected and 'Latte' in selected.name: - return video_run.generate('Latte', *args) + return video_run.generate(*args) return video_utils.queue_err(f'model not found: engine="{engine}" model="{model}"') diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 0e55b31af..da0ff8cc2 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -20,7 +20,7 @@ import transformers - LTXVideo 0.9.0 I2V: pass/pass/tbd - WAN 2.1 1.3B T2V: pass/pass/pass -- WAN 2.1 14B T2V: pass/pass/fail, error loading shard +- WAN 2.1 14B T2V: pass/pass/pass - WAN 2.1 14B I2V 480p: pass/pass/tbd - WAN 2.1 14B I2V 720p: pass/pass/tbd diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 09dee3826..2172fe053 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -7,15 +7,20 @@ from modules.video_models import models_def, video_utils, video_load, video_vae, debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None -def generate(keyword, *args, **kwargs): +def generate(*args, **kwargs): task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, faster_cache, pyramid_attention, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') found = [model.name for model in models_def.models.get(engine, [])] selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None - if not shared.sd_loaded or keyword not in shared.sd_model.__class__.__name__: + if not shared.sd_loaded: + debug('Video: model not yet loaded') video_load.load_model(selected) - if not shared.sd_loaded or keyword not in shared.sd_model.__class__.__name__: + if selected.name != video_load.loaded_model: + debug('Video: force reload') + video_load.load_model(selected) + if not shared.sd_loaded: + debug('Video: model still not loaded') return video_utils.queue_err('model not loaded') debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}') From c9391d67eb56de72d25a76e21fe2ef134532c902 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Mar 2025 21:17:41 -0400 Subject: [PATCH 044/122] video tab modernui layout Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 2 +- TODO.md | 8 -------- extensions-builtin/sdnext-modernui | 2 +- javascript/imageParams.js | 1 + javascript/progressBar.js | 2 ++ javascript/sdnext.css | 8 ++++---- modules/video_models/models_def.py | 4 ++-- scripts/prompt_enhance.py | 2 +- 8 files changed, 12 insertions(+), 17 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f98afce51..71a623c51 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,7 +14,7 @@ - Video: HunyuanVideo-I2V incompatible with latest transformers - Video: FasterCache and PyramidAttentionBroadcast granular config - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN - - Video: HunyuanVideo-I2V-16ch + - Video: HunyuanVideo-I2V-16/33ch - Video: CogVideo-15 support ### Highlights for 2025-03-21 diff --git a/TODO.md b/TODO.md index f262d8b78..0575e97c7 100644 --- a/TODO.md +++ b/TODO.md @@ -4,14 +4,6 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Current -- Video tab: - - remote vae - - tiny vae - - lora - - accelerators: teacache, pab, fastercache, paraattention, perflow - - modernui tab -- Detailer daemon: https://github.com/muerrilla/sd-webui-detail-daemon/blob/main/scripts/detail_daemon.py - ## Future Candidates - Redesign postprocessing diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index a9e015993..4d0bde42e 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit a9e015993194fe7b7c6222c9f32efbb557287a67 +Subproject commit 4d0bde42e95c801dba3a48b9161e12b8ddc5d5bd diff --git a/javascript/imageParams.js b/javascript/imageParams.js index 9ac42ace2..69057ba5e 100644 --- a/javascript/imageParams.js +++ b/javascript/imageParams.js @@ -9,6 +9,7 @@ async function initDragDrop() { if (tab === 0) promptTarget = 'txt2img_prompt_image'; else if (tab === 1) promptTarget = 'img2img_prompt_image'; else if (tab === 2) promptTarget = 'control_prompt_image'; + else if (tab === 3) promptTarget = 'video_prompt_image'; else return; const imgParent = gradioApp().getElementById(promptTarget); const fileInput = imgParent.querySelector('input[type="file"]'); diff --git a/javascript/progressBar.js b/javascript/progressBar.js index cfc9a9039..1edcbb501 100644 --- a/javascript/progressBar.js +++ b/javascript/progressBar.js @@ -14,9 +14,11 @@ function checkPaused(state) { lastState.paused = state ? !state : !lastState.paused; const t_el = document.getElementById('txt2img_pause'); const i_el = document.getElementById('img2img_pause'); + const c_el = document.getElementById('control_pause'); const v_el = document.getElementById('video_pause'); if (t_el) t_el.innerText = lastState.paused ? 'Resume' : 'Pause'; if (i_el) i_el.innerText = lastState.paused ? 'Resume' : 'Pause'; + if (c_el) c_el.innerText = lastState.paused ? 'Resume' : 'Pause'; if (v_el) v_el.innerText = lastState.paused ? 'Resume' : 'Pause'; } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index ec5e4fa3d..10f1f66d4 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -102,13 +102,13 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_prompt_container, #img2img_prompt_container, #control_prompt_container, #video_prompt_container { margin-right: var(--layout-gap) } #txt2img_footer, #img2img_footer, #control_footer { height: fit-content; display: none; } #txt2img_generate_box, #img2img_generate_box { gap: 0.5em; flex-wrap: unset; min-width: unset; width: 66.6%; } -#control_generate_box { gap: 0.5em; flex-wrap: unset; min-width: unset; width: 100%; } -#control_generate_box button:nth-child(1) { flex-grow: 2; } -#control_generate_box button:nth-child(2) { flex-grow: 1; } +#control_generate_box, #video_generate_box { gap: 0.5em; flex-wrap: unset; min-width: unset; width: 100%; } +#control_generate_box button:nth-child(1), #video_generate_box button:nth-child(1) { flex-grow: 2; } +#control_generate_box button:nth-child(2), #video_generate_box button:nth-child(2) { flex-grow: 1; } #txt2img_actions_column, #img2img_actions_column, #control_actions_column, #video_actions_column { gap: 0.3em; height: fit-content; } #txt2img_generate_box>button, #img2img_generate_box>button, #control_generate_box>button, #video_generate_box>button, #txt2img_enqueue, #img2img_enqueue, #txt2img_enqueue>button, #img2img_enqueue>button { min-height: 44px !important; max-height: 44px !important; line-height: 1em; white-space: break-spaces; min-width: unset; } #txt2img_enqueue_wrapper, #img2img_enqueue_wrapper, #control_enqueue_wrapper { min-width: unset !important; width: 31%; } -#txt2img_generate_line2, #img2img_generate_line2, #txt2img_tools, #img2img_tools, #control_generate_line2, #control_tools { display: flex; } +#txt2img_generate_line2, #img2img_generate_line2, #txt2img_tools, #img2img_tools, #control_generate_line2, #control_tools, #video_generate_line2, #video_tools { display: flex; } #txt2img_generate_line2>button, #img2img_generate_line2>button, #extras_generate_box>button, #control_generate_line2>button, #txt2img_tools>button, #img2img_tools>button, #control_tools>button { height: 2em; line-height: 0; font-size: var(--text-md); min-width: unset; display: block !important; } #txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt, #video_prompt, #video_neg_prompt { display: contents; } diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index da0ff8cc2..8c1c3f81a 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -15,9 +15,9 @@ import transformers - LTXVideo 0.9.5 T2V: pass/pass/fail, v095 pipeline is tbd - LTXVideo 0.9.5 I2V: pass/pass/fail, v095 pipeline is tbd - LTXVideo 0.9.1 T2V: pass/pass/pass -- LTXVideo 0.9.1 I2V: pass/pass/tbd +- LTXVideo 0.9.1 I2V: pass/pass/pass - LTXVideo 0.9.0 T2V: pass/pass/pass -- LTXVideo 0.9.0 I2V: pass/pass/tbd +- LTXVideo 0.9.0 I2V: pass/pass/pass - WAN 2.1 1.3B T2V: pass/pass/pass - WAN 2.1 14B T2V: pass/pass/pass diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 0ad17bba4..17613964a 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -95,5 +95,5 @@ class Script(scripts.Script): shared.log.debug(f'Prompt enhance: prompt="{p.prompt}"') def after_component(self, component, **kwargs): # searching for actual ui prompt components - if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt']: + if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt', 'video_prompt']: self.prompt = component From 8fdcc5711c41455e251139962e4897e8b5fe1ee0 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Sat, 22 Mar 2025 11:37:43 +0900 Subject: [PATCH 045/122] zluda fix bug --- modules/zluda_installer.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index a280a215a..ae801b4b8 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -43,8 +43,6 @@ class ZLUDALibrary: class Core(ZLUDALibrary): - internal: ctypes.WinDLL - def __init__(self, internal: ctypes.WinDLL): internal.zluda_get_hip_object.restype = ZLUDAResult internal.zluda_get_hip_object.argtypes = [ctypes.c_void_p, ctypes.c_int] @@ -78,7 +76,7 @@ def set_default_agent(agent: rocm.Agent): is_nightly = False try: load_core_modules() - is_nightly = core.get_nightly_flag() + is_nightly = core.get_nightly_flag() == 1 except Exception: pass From 85bda171f491ae263bb96e7692c806ca86adb01a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 22 Mar 2025 09:26:03 -0400 Subject: [PATCH 046/122] update video notes and todo Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 22 +++------------------- TODO.md | 18 ++++++++++++++++++ installer.py | 2 +- modules/video_models/models_def.py | 14 +++++++------- wiki | 2 +- 5 files changed, 30 insertions(+), 28 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 71a623c51..c839cf403 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,33 +1,17 @@ # Change Log for SD.Next -## Update for 2025-03-21 +## Update for 2025-03-22 -### ToDo/Limitations - - - VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - - VAE Remote encode: SD15 and Flux.1 issues: - - Video: API support is TBD - - Video: Wiki page is TBD - - Video: LTXVideo-095 params - - Video: LTXVideo-095 support for conditioned input - - Video: LTXVideo-095 support for offloading - - Video: HunyuanVideo-I2V incompatible with latest transformers - - Video: FasterCache and PyramidAttentionBroadcast granular config - - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN - - Video: HunyuanVideo-I2V-16/33ch - - Video: CogVideo-15 support - -### Highlights for 2025-03-21 +### Highlights for 2025-03-22 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides -### Details for 2025-03-21 +### Details for 2025-03-22 - **Video tab** - - initial release so consider this as alpha version - new top-level tab, replaces previous *video* script in text/image tabs old scripts are still present, but will be removed in the future - support for all latest models: diff --git a/TODO.md b/TODO.md index 0575e97c7..95426f72b 100644 --- a/TODO.md +++ b/TODO.md @@ -4,6 +4,24 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Current +### Issues/Limitations + +- VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` +- VAE Remote encode: SD15 and Flux.1 issues: +- Video: API support is TBD +- Video: Hunyuan Video I2V: transformers incompatibility +- Video: Hunyuan Video I2V: 16ch vs 33ch processing +- Video: WAN 2.1 14B I2V 480p: broken offload +- Video: WAN 2.1 14B I2V 720p: broken offload +- Video: CogVideoX 1.5 5B T2V/I2V: requires pipeline update +- Video: CogVideoX 1.5 5B I2V: requires pipeline update +- Video: LTXVideo 0.9.5 T2V/I2V: broken offload, new pipeline +- Video: LTXVideo 0.9.5 T2V/I2V: set preset params +- Video: LTXVideo 0.9.5 T2V/I2V: support for conditioned input +- Video: LTXVideo 0.9.1 I2V: generator list mismatch +- Video: FasterCache and PyramidAttentionBroadcast granular config +- Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN + ## Future Candidates - Redesign postprocessing diff --git a/installer.py b/installer.py index cacb0a62d..1ae99f737 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git: return - sha = '844221ae4e20a8939ee052f75874e284f75d4c5c' # diffusers commit hash + sha = 'a7d53a59394d5d8367826663601b69828e9f74fc' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 8c1c3f81a..f5a7804d1 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -12,23 +12,23 @@ import transformers - SkyReels Hunyuan I2V: pass/pass/pass - Fast Hunyuan T2V: pass/pass/pass -- LTXVideo 0.9.5 T2V: pass/pass/fail, v095 pipeline is tbd -- LTXVideo 0.9.5 I2V: pass/pass/fail, v095 pipeline is tbd +- LTXVideo 0.9.5 T2V: pass/pass/fail, completely broken offload, new pipeline +- LTXVideo 0.9.5 I2V: pass/pass/fail, completely broken offload, new pipeline - LTXVideo 0.9.1 T2V: pass/pass/pass -- LTXVideo 0.9.1 I2V: pass/pass/pass +- LTXVideo 0.9.1 I2V: pass/pass/fail, generator list mismatch - LTXVideo 0.9.0 T2V: pass/pass/pass - LTXVideo 0.9.0 I2V: pass/pass/pass - WAN 2.1 1.3B T2V: pass/pass/pass - WAN 2.1 14B T2V: pass/pass/pass -- WAN 2.1 14B I2V 480p: pass/pass/tbd -- WAN 2.1 14B I2V 720p: pass/pass/tbd +- WAN 2.1 14B I2V 480p: pass/pass/fail, offloading cpu vs cuda +- WAN 2.1 14B I2V 720p: pass/pass/fail, offloading cpu vs cuda - CogVideoX 1.0 2B T2V: pass/pass/pass - CogVideoX 1.0 5B T2V: pass/pass/pass - CogVideoX 1.0 5B I2V: pass/pass/pass -- CogVideoX 1.5 5B T2V: download/load/fail, v15 pipeline is tbd -- CogVideoX 1.5 5B I2V: download/load/fail, v15 pipeline is tbd +- CogVideoX 1.5 5B T2V: download/load/fail, pipeline is tbd +- CogVideoX 1.5 5B I2V: download/load/fail, pipeline is tbd - Mochi 1 T2V: pass/pass/pass - Latte 1 T2V: pass/pass/fail, float vs bfloat during generate diff --git a/wiki b/wiki index 3f46b4f74..00145b30c 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 3f46b4f742e439dee1d012c9e5e2ddf2a6b79aa6 +Subproject commit 00145b30c5ed318423487f8aa6336d834b578db7 From ab9e87d848d83d20ccbf0ed34ce9b7bceef55f00 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Sat, 22 Mar 2025 23:10:59 +0900 Subject: [PATCH 047/122] zluda flash attention 2 via triton --- html/licenses.html | 34 + modules/flash_attn_triton_amd/__init__.py | 0 modules/flash_attn_triton_amd/bwd_prefill.py | 606 +++++++++++++++ modules/flash_attn_triton_amd/bwd_ref.py | 271 +++++++ modules/flash_attn_triton_amd/fwd_decode.py | 700 ++++++++++++++++++ modules/flash_attn_triton_amd/fwd_prefill.py | 634 ++++++++++++++++ modules/flash_attn_triton_amd/fwd_ref.py | 258 +++++++ modules/flash_attn_triton_amd/interface_fa.py | 394 ++++++++++ modules/flash_attn_triton_amd/utils.py | 280 +++++++ modules/zluda_hijacks.py | 64 +- 10 files changed, 3233 insertions(+), 8 deletions(-) create mode 100644 modules/flash_attn_triton_amd/__init__.py create mode 100644 modules/flash_attn_triton_amd/bwd_prefill.py create mode 100644 modules/flash_attn_triton_amd/bwd_ref.py create mode 100644 modules/flash_attn_triton_amd/fwd_decode.py create mode 100644 modules/flash_attn_triton_amd/fwd_prefill.py create mode 100644 modules/flash_attn_triton_amd/fwd_ref.py create mode 100644 modules/flash_attn_triton_amd/interface_fa.py create mode 100644 modules/flash_attn_triton_amd/utils.py diff --git a/html/licenses.html b/html/licenses.html index 6597fa3ae..dc0e1fdbe 100644 --- a/html/licenses.html +++ b/html/licenses.html @@ -637,6 +637,40 @@ SOFTWARE. limitations under the License. +

Flash Attention

+Fast and memory-efficient exact attention +
+BSD 3-Clause License
+
+Copyright (c) 2022, the respective contributors, as shown by the AUTHORS file.
+All rights reserved.
+
+Redistribution and use in source and binary forms, with or without
+modification, are permitted provided that the following conditions are met:
+
+* Redistributions of source code must retain the above copyright notice, this
+   list of conditions and the following disclaimer.
+
+* Redistributions in binary form must reproduce the above copyright notice,
+   this list of conditions and the following disclaimer in the documentation
+   and/or other materials provided with the distribution.
+
+* Neither the name of the copyright holder nor the names of its
+   contributors may be used to endorse or promote products derived from
+   this software without specific prior written permission.
+
+THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
+FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
+DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
+SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
+CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
+OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
+OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+
+

Curated transformers

The MPS workaround for nn.Linear on macOS 13.2.X is based on the MPS workaround for nn.Linear created by danieldk for Curated transformers
diff --git a/modules/flash_attn_triton_amd/__init__.py b/modules/flash_attn_triton_amd/__init__.py
new file mode 100644
index 000000000..e69de29bb
diff --git a/modules/flash_attn_triton_amd/bwd_prefill.py b/modules/flash_attn_triton_amd/bwd_prefill.py
new file mode 100644
index 000000000..7f5be379b
--- /dev/null
+++ b/modules/flash_attn_triton_amd/bwd_prefill.py
@@ -0,0 +1,606 @@
+import torch
+import triton
+import triton.language as tl
+from modules.flash_attn_triton_amd.utils import get_shape_from_layout, get_strides_from_layout
+
+
+@triton.jit
+def _bwd_preprocess_use_o(
+    Out,
+    DO,
+    Delta,
+    stride_oz, stride_oh, stride_om, stride_ok,
+    stride_doz, stride_doh, stride_dom, stride_dok, # pylint: disable=unused-argument
+    stride_deltaz, stride_deltah, stride_deltam,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q,
+    max_seqlen_k,
+    BLOCK_M: tl.constexpr,
+    BLOCK_DMODEL: tl.constexpr,
+    ACTUAL_BLOCK_DMODEL: tl.constexpr,
+    N_CTX_Q: tl.constexpr,
+    Z: tl.constexpr, # pylint: disable=unused-argument
+    H: tl.constexpr,
+    IS_VARLEN: tl.constexpr
+):
+    pid_m = tl.program_id(0)
+    pid_bh = tl.program_id(1)
+
+    # Compute batch and head indices
+    off_z = pid_bh // H
+    off_h = pid_bh % H
+
+    if IS_VARLEN:
+        # Compute sequence lengths for the current batch
+        q_start = tl.load(cu_seqlens_q + off_z)
+        q_end = tl.load(cu_seqlens_q + off_z + 1)
+        k_start = tl.load(cu_seqlens_k + off_z)
+        k_end = tl.load(cu_seqlens_k + off_z + 1)
+
+        # Compute actual sequence lengths
+        N_CTX_Q = q_end - q_start
+        N_CTX_K = k_end - k_start # pylint: disable=unused-variable
+    else:
+        q_start = 0
+        k_start = 0
+        N_CTX_Q = max_seqlen_q
+        N_CTX_K = max_seqlen_k # pylint: disable=unused-variable
+
+    off_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+    off_d = tl.arange(0, BLOCK_DMODEL)
+
+    # create masks
+    mask_m = off_m < N_CTX_Q
+    mask_d = off_d < ACTUAL_BLOCK_DMODEL
+
+    # compute offsets
+    o_offset = Out + off_z * stride_oz + off_h * stride_oh + q_start * stride_om
+    do_offset = DO + off_z * stride_oz + off_h * stride_oh + q_start * stride_om
+
+    # compute pointers
+    out_ptrs = o_offset + off_m[:, None] * stride_om + off_d[None, :] * stride_ok
+    do_ptrs = do_offset + off_m[:, None] * stride_dom + off_d[None, :] * stride_dok
+
+    # load
+    o = tl.load(out_ptrs, mask=mask_m[:, None] & mask_d[None, :], other=0.0).to(tl.float32)
+    do = tl.load(do_ptrs, mask=mask_m[:, None] & mask_d[None, :], other=0.0).to(tl.float32)
+
+    # compute delta
+    delta = tl.sum(o * do, axis=1)
+
+    # write-back delta
+    delta_offset = Delta + off_z * stride_deltaz + off_h * stride_deltah + q_start * stride_deltam
+    delta_ptrs = delta_offset + off_m * stride_deltam
+    tl.store(delta_ptrs, delta, mask=mask_m)
+
+
+@triton.jit
+def _bwd_kernel_one_col_block(
+    Q,
+    K,
+    V,
+    sm_scale,
+    Out, DO, DQ, DK, DV, L, D, # pylint: disable=unused-argument
+    q_offset,
+    k_offset,
+    v_offset,
+    do_offset,
+    dq_offset,
+    dk_offset,
+    dv_offset,
+    d_offset,
+    l_offset,
+    stride_dq_all, stride_qz, stride_qh, # pylint: disable=unused-argument
+    stride_qm,
+    stride_qk,
+    stride_kz, stride_kh, # pylint: disable=unused-argument
+    stride_kn,
+    stride_kk,
+    stride_vz, stride_vh, # pylint: disable=unused-argument
+    stride_vn,
+    stride_vk,
+    stride_deltaz,  stride_deltah, # pylint: disable=unused-argument
+    stride_deltam,
+    Z, H, # pylint: disable=unused-argument
+    N_CTX_Q,
+    N_CTX_K,
+    off_h, off_z, off_hz, # pylint: disable=unused-argument
+    start_n,
+    num_block_m,
+    num_block_n, # pylint: disable=unused-argument
+    BLOCK_M: tl.constexpr,
+    BLOCK_DMODEL: tl.constexpr,
+    ACTUAL_BLOCK_DMODEL: tl.constexpr,
+    BLOCK_N: tl.constexpr,
+    SEQUENCE_PARALLEL: tl.constexpr,
+    CAUSAL: tl.constexpr,
+    USE_EXP2: tl.constexpr,
+):
+    if CAUSAL:
+        # TODO: Causal can skip more blocks with something like lo = start_m * BLOCK_M
+        lo = 0
+    else:
+        lo = 0
+
+    # initialize col and head offsets
+    offs_n = start_n * BLOCK_N + tl.arange(0, BLOCK_N)
+    offs_d = tl.arange(0, BLOCK_DMODEL)
+
+    # masks
+    mask_n = offs_n < N_CTX_K
+    mask_d = offs_d < ACTUAL_BLOCK_DMODEL
+    kv_mask = mask_n[:, None] & mask_d[None, :]
+
+    # initialize grad accumulators
+    dv = tl.zeros([BLOCK_N, BLOCK_DMODEL], dtype=tl.float32)
+    dk = tl.zeros([BLOCK_N, BLOCK_DMODEL], dtype=tl.float32)
+
+    # load k and v once per column block
+    k_ptrs = k_offset + offs_n[:, None] * stride_kn + offs_d[None, :] * stride_kk
+    v_ptrs = v_offset + offs_n[:, None] * stride_vn + offs_d[None, :] * stride_vk
+    k = tl.load(k_ptrs, mask=kv_mask, other=0.0)
+    v = tl.load(v_ptrs, mask=kv_mask, other=0.0)
+
+    # loop over rows
+    for start_m in range(lo, num_block_m * BLOCK_M, BLOCK_M):
+        offs_m = start_m + tl.arange(0, BLOCK_M)
+        q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qk
+        dq_ptrs = dq_offset + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qk
+        do_ptrs = do_offset + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qk
+
+        # update mask as row block changes
+        mask_m = offs_m < N_CTX_Q
+        q_mask = mask_m[:, None] & mask_d[None, :]
+
+        # load q, k, v, do on-chip
+        q = tl.load(q_ptrs, mask=q_mask, other=0.0)
+        do = tl.load(do_ptrs, mask=q_mask, other=0.0)
+
+        # recompute p = softmax(qk, dim=-1).T
+        qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
+        qk += tl.dot(q, tl.trans(k))
+
+        if CAUSAL:
+            col_offset = N_CTX_Q - N_CTX_K
+            causal_mask = offs_m[:, None] >= (col_offset + offs_n[None, :])
+            qk = tl.where(causal_mask, qk, float("-inf"))
+
+        l_ptrs = l_offset + offs_m * stride_deltam
+        l_i = tl.load(l_ptrs, mask=mask_m)
+
+        # compute p
+        if USE_EXP2:
+            RCP_LN2: tl.constexpr = 1.4426950408889634
+            qk *= sm_scale * RCP_LN2
+            l_i *= RCP_LN2
+            p = tl.math.exp2(qk - l_i[:, None])
+        else:
+            qk *= sm_scale
+            p = tl.math.exp(qk - l_i[:, None])
+
+        # mask block in the cases where the data is smaller the block size
+        p_mask = mask_m[:, None] & mask_n[None, :]
+        p = tl.where(p_mask, p, 0.0)
+
+        # compute dv
+        dv += tl.dot(tl.trans(p.to(Q.dtype.element_ty)), do)
+
+        # compute dp
+        dp = tl.dot(do, tl.trans(v))
+
+        # compute ds , ds = p * (dp - delta[:, None])
+        d_ptrs = d_offset + offs_m * stride_deltam
+        Di = tl.load(d_ptrs, mask=mask_m)
+        ds = (p * (dp - Di[:, None])) * sm_scale
+        ds = tl.where(p_mask, ds, 0.0).to(Q.dtype.element_ty)
+
+        # compute dk = dot(ds.T, q)
+        dk += tl.dot(tl.trans(ds), q)
+
+        # compute dq
+        if SEQUENCE_PARALLEL:
+            dq = tl.dot(ds, k)
+        else:
+            dq = tl.load(dq_ptrs, mask=q_mask, other=0.0)
+            dq += tl.dot(ds, k)
+        tl.store(dq_ptrs, dq.to(Q.dtype.element_ty), mask=q_mask)
+
+    # write-back dv and dk
+    dk_ptrs = dk_offset + offs_n[:, None] * stride_kn + offs_d[None, :] * stride_kk
+    dv_ptrs = dv_offset + offs_n[:, None] * stride_vn + offs_d[None, :] * stride_vk
+
+    # write-back
+    tl.store(dk_ptrs, dk.to(K.dtype.element_ty), mask=kv_mask)
+    tl.store(dv_ptrs, dv.to(V.dtype.element_ty), mask=kv_mask)
+
+@triton.jit
+def _bwd_kernel(
+    Q,
+    K,
+    V,
+    sm_scale,
+    Out,
+    DO,
+    DQ,
+    DK,
+    DV,
+    L,
+    D,
+    stride_dq_all,
+    stride_qz,
+    stride_qh,
+    stride_qm,
+    stride_qk,
+    stride_kz,
+    stride_kh,
+    stride_kn,
+    stride_kk,
+    stride_vz,
+    stride_vh,
+    stride_vn,
+    stride_vk,
+    stride_deltaz,
+    stride_deltah,
+    stride_deltam,
+    Z,
+    H,
+    num_block_m,
+    num_block_n,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q,
+    max_seqlen_k,
+    BLOCK_M: tl.constexpr,
+    BLOCK_DMODEL: tl.constexpr,
+    ACTUAL_BLOCK_DMODEL: tl.constexpr,
+    BLOCK_N: tl.constexpr,
+    SEQUENCE_PARALLEL: tl.constexpr,
+    CAUSAL: tl.constexpr,
+    USE_EXP2: tl.constexpr,
+    IS_VARLEN: tl.constexpr,
+):
+    # program ids
+    off_hz = tl.program_id(0)
+    if SEQUENCE_PARALLEL:
+        start_n = tl.program_id(1)
+    off_z = off_hz // H
+    off_h = off_hz % H
+
+    if IS_VARLEN:
+        # Compute sequence lengths for the current batch
+        q_start = tl.load(cu_seqlens_q + off_z)
+        q_end = tl.load(cu_seqlens_q + off_z + 1)
+        k_start = tl.load(cu_seqlens_k + off_z)
+        k_end = tl.load(cu_seqlens_k + off_z + 1)
+
+        # Compute actual sequence lengths
+        N_CTX_Q = q_end - q_start
+        N_CTX_K = k_end - k_start
+    else:
+        q_start = 0
+        k_start = 0
+        N_CTX_Q = max_seqlen_q
+        N_CTX_K = max_seqlen_k
+
+    # input tensor offsets
+    q_offset = Q + off_z * stride_qz + off_h * stride_qh + q_start * stride_qm
+    k_offset = K + off_z * stride_kz + off_h * stride_kh + k_start * stride_kn
+    v_offset = V + off_z * stride_vz + off_h * stride_vh + k_start * stride_vn
+    do_offset = DO + off_z * stride_qz + off_h * stride_qh + q_start * stride_qm
+    l_offset = L + off_z * stride_deltaz + off_h * stride_deltah + q_start * stride_deltam
+    d_offset = D + off_z * stride_deltaz + off_h * stride_deltah + q_start * stride_deltam
+
+    # output tensor offsets
+    dk_offset = DK + off_z * stride_kz + off_h * stride_kh + k_start * stride_kn
+    dv_offset = DV + off_z * stride_vz + off_h * stride_vh + k_start * stride_vn
+    if SEQUENCE_PARALLEL:
+        dq_offset = DQ + start_n * stride_dq_all + off_z * stride_qz + off_h * stride_qh + q_start * stride_qm
+    else:
+        dq_offset = DQ + off_z * stride_qz + off_h * stride_qh + q_start * stride_qm
+
+    # inner loop
+    if SEQUENCE_PARALLEL:
+        _bwd_kernel_one_col_block(
+            Q,
+            K,
+            V,
+            sm_scale,
+            Out,
+            DO,
+            DQ,
+            DK,
+            DV,
+            L,
+            D,
+            q_offset,
+            k_offset,
+            v_offset,
+            do_offset,
+            dq_offset,
+            dk_offset,
+            dv_offset,
+            d_offset,
+            l_offset,
+            stride_dq_all,
+            stride_qz,
+            stride_qh,
+            stride_qm,
+            stride_qk,
+            stride_kz,
+            stride_kh,
+            stride_kn,
+            stride_kk,
+            stride_vz,
+            stride_vh,
+            stride_vn,
+            stride_vk,
+            stride_deltaz,
+            stride_deltah,
+            stride_deltam,
+            Z,
+            H,
+            N_CTX_Q,
+            N_CTX_K,
+            off_h,
+            off_z,
+            off_hz,
+            start_n,
+            num_block_m,
+            num_block_n,
+            BLOCK_M=BLOCK_M,
+            BLOCK_DMODEL=BLOCK_DMODEL,
+            ACTUAL_BLOCK_DMODEL=ACTUAL_BLOCK_DMODEL,
+            BLOCK_N=BLOCK_N,
+            SEQUENCE_PARALLEL=SEQUENCE_PARALLEL,
+            CAUSAL=CAUSAL,
+            USE_EXP2=USE_EXP2,
+        )
+    else:
+        for start_n in range(0, num_block_n):
+            _bwd_kernel_one_col_block(
+                Q,
+                K,
+                V,
+                sm_scale,
+                Out,
+                DO,
+                DQ,
+                DK,
+                DV,
+                L,
+                D,
+                q_offset,
+                k_offset,
+                v_offset,
+                do_offset,
+                dq_offset,
+                dk_offset,
+                dv_offset,
+                d_offset,
+                l_offset,
+                stride_dq_all,
+                stride_qz,
+                stride_qh,
+                stride_qm,
+                stride_qk,
+                stride_kz,
+                stride_kh,
+                stride_kn,
+                stride_kk,
+                stride_vz,
+                stride_vh,
+                stride_vn,
+                stride_vk,
+                stride_deltaz,
+                stride_deltah,
+                stride_deltam,
+                Z,
+                H,
+                N_CTX_Q,
+                N_CTX_K,
+                off_h,
+                off_z,
+                off_hz,
+                start_n,
+                num_block_m,
+                num_block_n,
+                BLOCK_M=BLOCK_M,
+                BLOCK_DMODEL=BLOCK_DMODEL,
+                ACTUAL_BLOCK_DMODEL=ACTUAL_BLOCK_DMODEL,
+                BLOCK_N=BLOCK_N,
+                SEQUENCE_PARALLEL=SEQUENCE_PARALLEL,
+                CAUSAL=CAUSAL,
+                USE_EXP2=USE_EXP2,
+            )
+
+
+# NOTE: smaller blocks have lower accuracy. more accumlation error probably 128 * 128 seems good but leads to oom. 64 * 64 has accumlation errors but no oom.
+def attention_prefill_backward_triton_impl(
+    do,
+    q,
+    k,
+    v,
+    o,
+    softmax_lse,
+    dq,
+    dk,
+    dv,
+    sm_scale: float,
+    alibi_slopes, # pylint: disable=unused-argument
+    causal,
+    layout: str,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q: int,
+    max_seqlen_k: int,
+    use_exp2: bool,
+    sequence_parallel = True,
+):
+    # make contigious
+    q = q.contiguous()
+    k = k.contiguous()
+    v = v.contiguous()
+    softmax_lse = softmax_lse.contiguous()
+
+    # get strides and shape
+    batch, nheads_q, nheads_k, head_size, max_seqlen_q, max_seqlen_k = get_shape_from_layout(q, k, layout, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k) # pylint: disable=unused-variable
+    q_strides, k_strides, v_strides, o_strides = get_strides_from_layout(q, k, v, o, layout)
+    stride_qz, stride_qh, stride_qm, stride_qk =  q_strides
+    stride_kz, stride_kh, stride_kn, stride_kk = k_strides
+    stride_vz, stride_vh, stride_vn, stride_vk = v_strides
+    stride_oz, stride_oh, stride_om, stride_ok = o_strides
+    batch_headsize = batch * nheads_q
+    is_varlen = layout == "thd"
+
+    # FIXME: some configs lead to oom for some reason when using 64 x 64 blocks
+    if max_seqlen_q <= 32 or max_seqlen_k <= 32:
+        BLOCK_M = 32
+        BLOCK_N = 32
+    else:
+        BLOCK_M = 64
+        BLOCK_N = 64
+    num_warps = 4 # NOTE: originial is 8. changing it to 1 caused issues be careful
+    num_stages = 1
+    waves_per_eu = 1
+
+    # divide up the problem
+    num_blocks_m = triton.cdiv(max_seqlen_q, BLOCK_M)
+    num_blocks_n = triton.cdiv(max_seqlen_k, BLOCK_N)
+
+    # get closest power of 2 over or equal to 32.
+    padded_d_model = 1 << (head_size - 1).bit_length()
+    padded_d_model = max(padded_d_model, 16)
+    BLOCK_DMODEL = padded_d_model
+    ACTUAL_BLOCK_DMODEL = head_size
+
+    do = do.contiguous()
+    # NOTE: we might need to copy the output tensor if they are not continuous or have other issues
+    copy_back = {"dq": False, "dk": False, "dv": False}
+
+    dq_og = None
+    # deal with dq
+    if dq is None:
+        if sequence_parallel:
+            dq = torch.zeros((num_blocks_n,) + q.shape, device=q.device, dtype=q.dtype)
+        else:
+            dq = torch.zeros(q.shape, device=q.device, dtype=q.dtype)
+    else:
+        dq_og = dq
+        if not dq.is_contiguous():
+            dq = dq.contiguous()
+            copy_back["dq"] = True
+
+        if sequence_parallel:
+            dq = torch.zeros((num_blocks_n,) + q.shape, device=q.device, dtype=q.dtype)
+            copy_back["dq"] = True
+        else:
+            # NOTE: the kernel does inplace accumlation so dq has to be zeros. This avoids the case where we are passed empty dq and it is not all zeros
+            dq.zero_()
+    stride_dq_all = dq.stride()[0]
+
+    dk_og = None
+    dv_og = None
+    # deal with dk, dv
+    if (dk is None) or (dv is None):
+        dk = torch.empty_like(k)
+        dv = torch.empty_like(v)
+    else:
+        if not dk.is_contiguous():
+            dk_og = dk
+            dk = dk.contiguous()
+            copy_back["dk"] = True
+
+        if not dv.is_contiguous():
+            dv_og = dv
+            dv = dv.contiguous()
+            copy_back["dv"] = True
+
+    # assert contigious
+    assert do.is_contiguous()
+    assert q.is_contiguous()
+    assert k.is_contiguous()
+    assert v.is_contiguous()
+    assert o.is_contiguous()
+    assert softmax_lse.is_contiguous()
+
+    # init delta
+    delta = torch.empty_like(softmax_lse)
+    if is_varlen:
+        stride_deltam, stride_deltah = delta.stride()
+        stride_deltaz = 0
+    else:
+        stride_deltaz, stride_deltah, stride_deltam = delta.stride()
+
+    _bwd_preprocess_use_o[(num_blocks_m, batch_headsize)](
+        o,
+        do,
+        delta,
+        stride_oz, stride_oh, stride_om, stride_ok,
+        stride_oz, stride_oh, stride_om, stride_ok,
+        stride_deltaz, stride_deltah, stride_deltam,
+        cu_seqlens_q,
+        cu_seqlens_k,
+        max_seqlen_q,
+        max_seqlen_k,
+        BLOCK_M=BLOCK_M,
+        BLOCK_DMODEL=BLOCK_DMODEL,
+        ACTUAL_BLOCK_DMODEL=ACTUAL_BLOCK_DMODEL,
+        N_CTX_Q=max_seqlen_q,
+        Z=batch,
+        H=nheads_q,
+        IS_VARLEN=is_varlen
+    )
+
+    _bwd_kernel[(batch_headsize, num_blocks_n if sequence_parallel else 1)](
+        q,
+        k,
+        v,
+        sm_scale,
+        o,
+        do,
+        dq,
+        dk,
+        dv,
+        softmax_lse,
+        delta,
+        stride_dq_all,
+        stride_qz, stride_qh, stride_qm, stride_qk,
+        stride_kz, stride_kh, stride_kn, stride_kk,
+        stride_vz, stride_vh, stride_vn, stride_vk,
+        stride_deltaz, stride_deltah, stride_deltam,
+        batch,
+        nheads_q,
+        num_blocks_m,
+        num_blocks_n,
+        cu_seqlens_q,
+        cu_seqlens_k,
+        max_seqlen_q,
+        max_seqlen_k,
+        BLOCK_M=BLOCK_M,
+        BLOCK_N=BLOCK_N,
+        BLOCK_DMODEL=BLOCK_DMODEL,
+        ACTUAL_BLOCK_DMODEL=ACTUAL_BLOCK_DMODEL,
+        SEQUENCE_PARALLEL=sequence_parallel,
+        CAUSAL=causal,
+        USE_EXP2=use_exp2,
+        num_warps=num_warps,
+        num_stages=num_stages,
+        waves_per_eu = waves_per_eu,
+        IS_VARLEN=is_varlen
+    )
+
+    if sequence_parallel:
+        dq = dq.sum(dim=0)
+
+    if copy_back["dq"]:
+        dq_og.copy_(dq)
+        dq = dq_og
+    if copy_back["dk"]:
+        dk_og.copy_(dk)
+        dk = dk_og
+    if copy_back["dv"]:
+        dv_og.copy_(dv)
+        dv = dv_og
+
+    return dq, dk, dv, delta, None, None
diff --git a/modules/flash_attn_triton_amd/bwd_ref.py b/modules/flash_attn_triton_amd/bwd_ref.py
new file mode 100644
index 000000000..2b1befd88
--- /dev/null
+++ b/modules/flash_attn_triton_amd/bwd_ref.py
@@ -0,0 +1,271 @@
+import math
+import torch
+
+
+def attention_backward_core_ref_impl(
+    do, q, k, v, o, softmax_lse, sm_scale, causal, use_exp2
+):
+    # cast to float32
+    do = do.to(torch.float32)
+    q = q.to(torch.float32)
+    k = k.to(torch.float32)
+    v = v.to(torch.float32)
+    o = o.to(torch.float32)
+    softmax_lse = softmax_lse.to(torch.float32)
+
+    # recompute attention_scores. Make sure it matches the forward impl. i.e. It use float32
+    attention_scores = torch.matmul(q.to(torch.float32), k.transpose(-2, -1).to(torch.float32))
+
+    # scale scores
+    attention_scaled_scores = sm_scale * attention_scores
+
+    # Apply causal mask if necessary
+    if causal:
+        L_q, L_k = q.shape[1], k.shape[1]
+        row_idx = torch.arange(L_q, device=q.device).unsqueeze(1)
+        col_idx = torch.arange(L_k, device=q.device).unsqueeze(0)
+        col_offset = L_q-L_k
+        causal_mask = row_idx >= (col_offset + col_idx)
+        # set -inf to places the causal mask is false
+        attention_scaled_scores = attention_scaled_scores.masked_fill(
+             torch.logical_not(causal_mask.unsqueeze(0)), float('-inf')
+        )
+
+    # compute probabilities using softmax_lse
+    if use_exp2:
+        RCP_LN = 1 / math.log(2)
+        attention_scaled_scores_base2 = attention_scaled_scores * RCP_LN
+        softmax_lse_base2 = softmax_lse * RCP_LN
+        softmax_lse_3d =  softmax_lse_base2.unsqueeze(-1)
+        p = torch.exp2(attention_scaled_scores_base2 - softmax_lse_3d)
+    else:
+        softmax_lse_3d =  softmax_lse.unsqueeze(-1)
+        p = torch.exp(attention_scaled_scores - softmax_lse_3d)
+
+    # compute gradient wrt v
+    dv = torch.matmul(p.transpose(-2, -1), do.to(torch.float32))
+
+    # compute dp
+    dp = torch.matmul(do, v.transpose(-2, -1))
+
+    # calculate ds using dp
+    delta = torch.sum(o * do, axis=-1).to(torch.float32)  # what OAI kernel uses
+    delta_3d = delta.unsqueeze(-1)
+    ds = (p * (dp - delta_3d)) * sm_scale
+
+    # compute gradient wrt k
+    dk = torch.matmul(ds.transpose(-2, -1), q.to(torch.float32))
+
+    # compute gradient wrt q
+    dq = torch.matmul(ds, k.to(torch.float32))
+
+    # cast back to original dtype
+    dq = dq.to(torch.float16)
+    dk = dk.to(torch.float16)
+    dv = dv.to(torch.float16)
+
+    # remove d dim with size 1
+    delta = delta_3d.squeeze(-1)
+
+    return dq, dk, dv, delta
+
+def attention_varlen_backward_pytorch_ref_impl(
+    do,
+    q,
+    k,
+    v,
+    o,
+    softmax_lse,
+    sm_scale,
+    causal,
+    layout,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q, max_seqlen_k, # pylint: disable=unused-argument
+    use_exp2,
+):
+    # Ensure the layout is 'thd'
+    if layout != 'thd':
+        raise ValueError(f"Unsupported layout {layout}. Expected 'thd'.")
+
+    batch_size = cu_seqlens_q.shape[0] - 1
+    num_heads = q.shape[1]
+    head_dim = q.shape[2] # pylint: disable=unused-variable
+
+    # Pre-allocate outputs
+    total_L_q = q.shape[0]
+    total_L_k = k.shape[0] # pylint: disable=unused-variable
+
+    dq = torch.zeros_like(q)
+    dk = torch.zeros_like(k)
+    dv = torch.zeros_like(v)
+    # delta has the same shape as softmax_lse: [total_L_q, num_heads]
+    delta = torch.zeros((total_L_q, num_heads), dtype=torch.float32, device=o.device)
+
+    for i in range(batch_size):
+        # Get the start and end indices for the current sequence
+        start_q = cu_seqlens_q[i].item()
+        end_q = cu_seqlens_q[i + 1].item()
+        start_k = cu_seqlens_k[i].item()
+        end_k = cu_seqlens_k[i + 1].item()
+
+        # Extract q_i, k_i, v_i, do_i, o_i, softmax_lse_i
+        q_i = q[start_q:end_q, :, :]      # [L_q_i, num_heads, head_dim]
+        k_i = k[start_k:end_k, :, :]      # [L_k_i, num_heads, head_dim]
+        v_i = v[start_k:end_k, :, :]      # [L_k_i, num_heads, head_dim]
+        do_i = do[start_q:end_q, :, :]    # [L_q_i, num_heads, head_dim]
+        o_i = o[start_q:end_q, :, :]      # [L_q_i, num_heads, head_dim]
+        # softmax_lse has shape [total_L_q, num_heads]
+        softmax_lse_i = softmax_lse[start_q:end_q, :]  # [L_q_i, num_heads]
+        softmax_lse_i = softmax_lse_i.transpose(0, 1)  # [num_heads, L_q_i]
+
+        # Permute to [num_heads, L_q_i, head_dim]
+        q_i = q_i.permute(1, 0, 2)
+        k_i = k_i.permute(1, 0, 2)
+        v_i = v_i.permute(1, 0, 2)
+        do_i = do_i.permute(1, 0, 2)
+        o_i = o_i.permute(1, 0, 2)
+        # softmax_lse_i is already in [num_heads, L_q_i]
+
+        # Call the core backward function for this sequence
+        dq_i, dk_i, dv_i, delta_i = attention_backward_core_ref_impl(
+            do_i,
+            q_i,
+            k_i,
+            v_i,
+            o_i,
+            softmax_lse_i,
+            sm_scale,
+            causal,
+            use_exp2
+        )
+
+        # Convert back to 'thd' layout
+        dq_i = dq_i.permute(1, 0, 2)  # [L_q_i, num_heads, head_dim]
+        dk_i = dk_i.permute(1, 0, 2)  # [L_k_i, num_heads, head_dim]
+        dv_i = dv_i.permute(1, 0, 2)  # [L_k_i, num_heads, head_dim]
+
+        # Place outputs in pre-allocated tensors
+        dq[start_q:end_q, :, :] = dq_i
+        dk[start_k:end_k, :, :] += dk_i  # Accumulate gradients for shared keys
+        dv[start_k:end_k, :, :] += dv_i  # Accumulate gradients for shared values
+        # delta_i has shape [num_heads, L_q_i]
+        delta_i = delta_i.transpose(1, 0)  # [L_q_i, num_heads]
+        delta[start_q:end_q, :] = delta_i
+
+    return dq, dk, dv, delta
+
+def attention_vanilla_backward_pytorch_ref_impl(
+    do,
+    q,
+    k,
+    v,
+    o,
+    softmax_lse,
+    sm_scale,
+    causal,
+    layout,
+    use_exp2,
+):
+    if layout == "bshd":
+        do = do.transpose(1, 2).contiguous()
+        q = q.transpose(1, 2).contiguous()
+        k = k.transpose(1, 2).contiguous()
+        v = v.transpose(1, 2).contiguous()
+        o = o.transpose(1, 2).contiguous()
+    elif layout == "bhsd":
+        pass
+    else:
+        raise ValueError(f"Unknown layout {layout}")
+
+    # Prepare tensors in [batch_size * num_heads, seq_len, head_dim] format
+    batch_size, num_heads, seq_len_q, head_dim = q.shape
+    seq_len_k = k.shape[2]
+
+    # Merge batch and heads dimensions
+    do = do.reshape(batch_size * num_heads, seq_len_q, head_dim)
+    q = q.reshape(batch_size * num_heads, seq_len_q, head_dim)
+    k = k.reshape(batch_size * num_heads, seq_len_k, head_dim)
+    v = v.reshape(batch_size * num_heads, seq_len_k, head_dim)
+    softmax_lse = softmax_lse.reshape(batch_size * num_heads, seq_len_q)
+    o = o.reshape(batch_size * num_heads, seq_len_q, head_dim)
+
+    dq, dk, dv, delta = attention_backward_core_ref_impl(
+        do,
+        q,
+        k,
+        v,
+        o,
+        softmax_lse,
+        sm_scale,
+        causal,
+        use_exp2
+    )
+
+    # Reshape outputs back to [batch_size, num_heads, seq_len, head_dim]
+    dq = dq.reshape(batch_size, num_heads, seq_len_q, head_dim)
+    dk = dk.reshape(batch_size, num_heads, seq_len_k, head_dim)
+    dv = dv.reshape(batch_size, num_heads, seq_len_k, head_dim)
+    delta = delta.reshape(batch_size, num_heads, seq_len_q)
+
+    # Go back to original layout
+    if layout == "bshd":
+        dq = dq.transpose(1, 2)
+        dk = dk.transpose(1, 2)
+        dv = dv.transpose(1, 2)
+    elif layout == "bhsd":
+        pass
+    else:
+        raise ValueError(f"Unknown layout {layout}")
+
+    return dq, dk, dv, delta
+
+
+def attention_backward_pytorch_ref_impl(
+    do,
+    q,
+    k,
+    v,
+    o,
+    softmax_lse,
+    sm_scale,
+    causal,
+    layout,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q,
+    max_seqlen_k,
+    use_exp2
+):
+    if layout == "thd":
+        dq, dk, dv, delta = attention_varlen_backward_pytorch_ref_impl(
+            do,
+            q,
+            k,
+            v,
+            o,
+            softmax_lse,
+            sm_scale,
+            causal,
+            layout,
+            cu_seqlens_q,
+            cu_seqlens_k,
+            max_seqlen_q,
+            max_seqlen_k,
+            use_exp2,
+        )
+    else:
+        dq, dk, dv, delta = attention_vanilla_backward_pytorch_ref_impl(
+            do,
+            q,
+            k,
+            v,
+            o,
+            softmax_lse,
+            sm_scale,
+            causal,
+            layout,
+            use_exp2,
+        )
+
+    return dq, dk, dv, delta
diff --git a/modules/flash_attn_triton_amd/fwd_decode.py b/modules/flash_attn_triton_amd/fwd_decode.py
new file mode 100644
index 000000000..7a2a234d6
--- /dev/null
+++ b/modules/flash_attn_triton_amd/fwd_decode.py
@@ -0,0 +1,700 @@
+import torch
+import triton
+import triton.language as tl
+from modules.flash_attn_triton_amd.utils import _strides, get_padded_headsize
+
+
+@triton.jit
+def _fwd_kernel_splitK(
+    Q,
+    K,
+    V,
+    sm_scale,
+    Out_splitK,  # [B, H, split_k, Mq, K]
+    Metadata,  # [B, H, 2, split_k, M_ceil] contains [mi, li]
+    K_new,
+    V_new,
+    Cache_seqlens,
+    Cache_batch_idx,
+    Alibi_slopes,
+    stride_qz,
+    stride_qm,
+    stride_qg,
+    stride_qh,
+    stride_qd,
+    stride_kz,
+    stride_kn,
+    stride_kg,
+    stride_kh,
+    stride_kd,
+    stride_vz,
+    stride_vn,
+    stride_vg,
+    stride_vh,
+    stride_vd,
+    stride_osk_zhg,
+    stride_osk_s,
+    stride_osk_m,
+    stride_osk_d, # pylint: disable=unused-argument
+    stride_mzhg,
+    stride_m2,
+    stride_ms,
+    stride_mm, # pylint: disable=unused-argument
+    stride_kn_z,
+    stride_kn_n,
+    stride_kn_g,
+    stride_kn_h,
+    stride_kn_d,
+    stride_vn_z,
+    stride_vn_n,
+    stride_vn_g,
+    stride_vn_h,
+    stride_vn_d,
+    stride_az,
+    stride_ah,
+    Z, # pylint: disable=unused-argument
+    N_CTX_Q,
+    N_CTX_K,
+    N_CTX_NEW,
+    BLOCK_N_PER_SPLIT,
+    H_q: tl.constexpr,
+    H_kv: tl.constexpr,
+    G_q: tl.constexpr,
+    BLOCK_M: tl.constexpr,
+    BLOCK_DMODEL: tl.constexpr,
+    ACTUAL_BLOCK_DMODEL: tl.constexpr,
+    BLOCK_N: tl.constexpr,
+    BOUNDS_CHECKS_N: tl.constexpr,
+    USE_CACHE_SEQLENs: tl.constexpr,
+    USE_CACHE_BATCH_IDX: tl.constexpr,
+    NEW_KV: tl.constexpr,
+    IS_GQA: tl.constexpr,
+    IS_CAUSAL: tl.constexpr,
+    USE_ALIBI: tl.constexpr,
+):
+    # Padding
+    PADDED_HEAD: tl.constexpr = ACTUAL_BLOCK_DMODEL != BLOCK_DMODEL
+    if PADDED_HEAD:
+        d_mask = tl.arange(0, BLOCK_DMODEL) < ACTUAL_BLOCK_DMODEL
+
+    start_m = tl.program_id(0)
+    off_zhg = tl.program_id(1)
+    off_z = off_zhg // (H_q * G_q)
+    off_h_q = (off_zhg // G_q) % H_q
+    off_g_q = off_zhg % G_q
+    splitk_idx = tl.program_id(2)
+
+    # pick batch index
+    if USE_CACHE_BATCH_IDX:
+        cache_batch_idx = tl.load(Cache_batch_idx + off_z)
+    else:
+        cache_batch_idx = off_z
+
+    # Load ALiBi slope if enabled
+    if USE_ALIBI:
+        a_offset = off_z * stride_az + off_h_q * stride_ah
+        alibi_slope = tl.load(Alibi_slopes + a_offset)
+    else:
+        alibi_slope = None
+
+    lo = splitk_idx * BLOCK_N_PER_SPLIT
+    if USE_CACHE_SEQLENs:
+        cache_seqlen_last_idx = tl.load(Cache_seqlens + off_z)
+        if NEW_KV:
+            kv_len = cache_seqlen_last_idx + N_CTX_NEW
+        else:
+            kv_len = cache_seqlen_last_idx
+    else:
+        kv_len = N_CTX_K
+    hi = tl.minimum((splitk_idx + 1) * BLOCK_N_PER_SPLIT, kv_len)
+
+    HEAD_RATIO: tl.constexpr = H_q // H_kv
+    if IS_GQA:
+        k_head_idx = off_h_q // HEAD_RATIO
+        v_head_idx = k_head_idx
+    else:
+        k_head_idx = off_h_q
+        v_head_idx = off_h_q
+
+    # calculate base offset
+    k_base = K + k_head_idx * stride_kh + cache_batch_idx * stride_kz + off_g_q * stride_kg
+    v_base = V + v_head_idx * stride_vh + cache_batch_idx * stride_vz + off_g_q * stride_vg
+
+    # Copy new Keys and Values into Cache
+    if NEW_KV:
+        knew_base = K_new + k_head_idx * stride_kn_h + off_z * stride_kn_z + off_g_q * stride_kn_g
+
+        # Determine the starting position for new data in the cache
+        if USE_CACHE_SEQLENs:
+            start_idx = tl.load(Cache_seqlens + off_z)
+        else:
+            start_idx = N_CTX_K - N_CTX_NEW
+
+        # Copy new Keys
+        for i in range(0, N_CTX_NEW, BLOCK_N):
+            # Load from K_new
+            k_new_block = tl.load(
+                knew_base +
+                tl.arange(0, BLOCK_DMODEL)[:, None] * stride_kn_d +
+                (tl.arange(0, BLOCK_N) + i)[None, :] * stride_kn_n,
+                 mask=(tl.arange(0, BLOCK_N)[None, :] + i < N_CTX_NEW) &
+                     (tl.arange(0, BLOCK_DMODEL)[:, None] < ACTUAL_BLOCK_DMODEL),
+                other=0
+            )
+
+            # Store to K
+            tl.store(
+                k_base +
+                tl.arange(0, BLOCK_DMODEL)[:, None] * stride_kd +
+                (tl.arange(0, BLOCK_N) + i + start_idx)[None, :] * stride_kn,
+                k_new_block,
+                 mask=(tl.arange(0, BLOCK_N)[None, :] + i < N_CTX_NEW) &
+                     (tl.arange(0, BLOCK_DMODEL)[:, None] < ACTUAL_BLOCK_DMODEL),
+            )
+
+        # Copy new Values
+        vnew_base = V_new + v_head_idx * stride_vn_h + off_z * stride_vn_z + off_g_q * stride_vn_g
+        for i in range(0, N_CTX_NEW, BLOCK_N):
+            # Load from V_new
+            v_new_block = tl.load(
+                vnew_base +
+                (tl.arange(0, BLOCK_N) + i)[:, None] * stride_vn_n +
+                tl.arange(0, BLOCK_DMODEL)[None, :] * stride_vn_d,
+                mask=(tl.arange(0, BLOCK_N)[:, None] + i < N_CTX_NEW) &
+                     (tl.arange(0, BLOCK_DMODEL)[None, :] < ACTUAL_BLOCK_DMODEL),
+                other=0
+            )
+
+            # Store to V
+            tl.store(
+                v_base +
+                (tl.arange(0, BLOCK_N) + i + start_idx)[:, None] * stride_vn +
+                tl.arange(0, BLOCK_DMODEL)[None, :] * stride_vd,
+                v_new_block,
+                 mask=(tl.arange(0, BLOCK_N)[:, None] + i < N_CTX_NEW) &
+                     (tl.arange(0, BLOCK_DMODEL)[None, :] < ACTUAL_BLOCK_DMODEL),
+            )
+
+    Q_block_ptr = tl.make_block_ptr(
+        base=Q + off_h_q * stride_qh + off_z * stride_qz + off_g_q * stride_qg,
+        shape=(N_CTX_Q, ACTUAL_BLOCK_DMODEL),
+        strides=(stride_qm, stride_qd),
+        offsets=(start_m * BLOCK_M, 0),
+        block_shape=(BLOCK_M, BLOCK_DMODEL),
+        order=(1, 0),
+    )
+
+    K_block_ptr = tl.make_block_ptr(
+        base=k_base,
+        shape=(ACTUAL_BLOCK_DMODEL, hi),
+        strides=(stride_kd, stride_kn),
+        offsets=(0, lo),
+        block_shape=(BLOCK_DMODEL, BLOCK_N),
+        order=(0, 1),
+    )
+    V_block_ptr = tl.make_block_ptr(
+        base=v_base,
+        shape=(hi, ACTUAL_BLOCK_DMODEL),
+        strides=(stride_vn, stride_vd),
+        offsets=(lo, 0),
+        block_shape=(BLOCK_N, BLOCK_DMODEL),
+        order=(1, 0),
+    )
+
+    K_scale_shift_block_ptr = None
+    V_scale_shift_block_ptr = None
+
+    # initialize pointer to m and l
+    m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32)
+    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
+
+    acc = tl.zeros([BLOCK_M, BLOCK_DMODEL], dtype=tl.float32)  # noqa: F821
+
+    # scale sm_scale by log_2(e) and use
+    # 2^x instead of exp in the loop because CSE and LICM
+    # don't work as expected with `exp` in the loop
+    qk_scale = sm_scale * 1.44269504
+    # load q: it will stay in SRAM throughout
+    q = tl.load(  # noqa: F821
+        tl.advance(Q_block_ptr, (0, 0)), boundary_check=(0, ))
+    q = (q * qk_scale).to(q.dtype)
+    if PADDED_HEAD:
+        q = tl.where(d_mask[None, :], q, 0.0)
+
+    # loop over k, v and update accumulator
+    for start_n in range(lo, hi, BLOCK_N):
+        k, v = load_k_v_group(
+            K_block_ptr,
+            V_block_ptr,
+            K_scale_shift_block_ptr,
+            V_scale_shift_block_ptr,
+            BOUNDS_CHECKS_N,
+            1,
+            BLOCK_DMODEL,
+            ACTUAL_BLOCK_DMODEL,
+            Q.dtype.element_ty,
+            0,
+        )
+        if PADDED_HEAD:
+            k = tl.where(d_mask[:, None], k, 0.0)
+            v = tl.where(d_mask[None, :], v, 0.0)
+
+        # -- compute qk ---
+        qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
+        qk += tl.dot(q, k)  # noqa: F821
+
+        if USE_ALIBI:
+            row_idx = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
+            col_idx = start_n + tl.arange(0, BLOCK_N)
+
+            # Compute relative positions
+            relative_pos = row_idx[:, None] + kv_len - (N_CTX_Q + col_idx[None, :])
+            relative_pos = tl.abs(relative_pos)
+
+            # Compute ALiBi bias
+            alibi_bias = -1 * alibi_slope * relative_pos
+            qk += (alibi_bias * 1.44269504)
+
+        # Apply causal mask if IS_CAUSAL is True
+        if IS_CAUSAL:
+            row_idx = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
+            col_idx = start_n + tl.arange(0, BLOCK_N)
+
+            # create a N_CTX_Q x kv_len causal mask
+            col_offset = N_CTX_Q - kv_len
+            causal_mask = row_idx[:, None] >= (col_offset + col_idx[None, :])
+
+            # Apply the mask
+            qk = tl.where(causal_mask, qk, float("-inf"))
+
+        # TODO: This is slow, and only needed at the last iteration.
+        # Maybe we can unroll the last iteration instead?
+        if BOUNDS_CHECKS_N:
+            qk = tl.where(tl.arange(0, BLOCK_N) < hi - start_n, qk, float("-inf"))
+
+        # -- compute scaling constant ---
+        m_i_new = tl.maximum(m_i, tl.max(qk, 1))
+        if IS_CAUSAL:
+            alpha = tl.math.exp2(tl.where(m_i > float("-inf"), m_i - m_i_new, float("-inf")))
+        else:
+            alpha = tl.math.exp2(m_i - m_i_new)
+        # cause of nan because subtracting infs
+        if IS_CAUSAL:
+            qk = tl.where(qk > float("-inf"), qk - m_i_new[:, None], float("-inf"))
+        else:
+            qk = qk - m_i_new[:, None]
+
+        p = tl.math.exp2(qk)
+
+        # -- update m_i and l_i --
+        l_i = l_i * alpha + tl.sum(p, 1)
+        m_i = m_i_new
+        p = p.to(Q.dtype.element_ty)
+
+        # -- scale and update acc --
+        acc *= alpha[:, None]
+        acc += tl.dot(p.to(v.dtype), v)
+
+        # update pointers
+        K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N))
+        V_block_ptr = tl.advance(V_block_ptr, (BLOCK_N, 0))
+
+    # write back O
+    O_block_ptr = tl.make_block_ptr(
+        base=Out_splitK + off_zhg * stride_osk_zhg + splitk_idx * stride_osk_s,
+        shape=(N_CTX_Q, BLOCK_DMODEL),
+        strides=(stride_osk_m, 1),
+        offsets=(start_m * BLOCK_M, 0),
+        block_shape=(BLOCK_M, BLOCK_DMODEL),
+        order=(1, 0),
+    )
+    tl.store(
+        tl.advance(O_block_ptr, (0, 0)),
+        acc,
+        boundary_check=(0, ),
+    )
+    # Write metadata for split-K reduction
+    Metadata_ptr = (Metadata + off_zhg * stride_mzhg + splitk_idx * stride_ms + start_m * BLOCK_M +
+                    tl.arange(0, BLOCK_M))
+    tl.store(Metadata_ptr, m_i)
+    tl.store(Metadata_ptr + stride_m2, l_i)
+
+
+@triton.jit
+def load_k_v_group(
+    K_block_ptr,
+    V_block_ptr,
+    K_scale_shift_block_ptr, V_scale_shift_block_ptr, # pylint: disable=unused-argument
+    BOUNDS_CHECKS_N: tl.constexpr,
+    PACKED_PER_VAL: tl.constexpr, BLOCK_DMODEL: tl.constexpr, # pylint: disable=unused-argument
+    ACTUAL_BLOCK_DMODEL: tl.constexpr,
+    dtype: tl.constexpr, # pylint: disable=unused-argument
+    group_id: tl.constexpr,
+):
+    # Load K/V for a given block
+    # Advance to the current quantization group
+    K_block_ptr = tl.advance(K_block_ptr, (ACTUAL_BLOCK_DMODEL * group_id, 0))
+    V_block_ptr = tl.advance(V_block_ptr, (0, ACTUAL_BLOCK_DMODEL * group_id))
+
+    # -- load k, v --
+    k = tl.load(K_block_ptr, boundary_check=(1, ) if BOUNDS_CHECKS_N else ())
+    v = tl.load(V_block_ptr, boundary_check=(0, ) if BOUNDS_CHECKS_N else ())
+
+    return k, v
+
+
+@triton.jit
+def cast_uint32_to_half2(scale_shift):
+    # Extract two float16 packed into one int32
+    scale = scale_shift & 0xFFFF
+    shift = scale_shift >> 16
+    scale = scale.to(tl.uint16).to(tl.float16, bitcast=True)
+    shift = shift.to(tl.uint16).to(tl.float16, bitcast=True)
+    return scale, shift
+
+
+@triton.jit
+def dequantize(
+    x_,
+    scale,
+    shift,
+    PACKED_PER_VAL: tl.constexpr = 8,
+):
+    # PACKED_PER_VAL is the number of values packed into
+    # each element x_. For example, for int4 quantization
+    #and x_ of type int32, PACKED_PER_VAL is 8.
+
+    BLOCK_N: tl.constexpr = x_.shape[0]
+    BLOCK_DMODEL_PACKED: tl.constexpr = x_.shape[1]
+    offsets = tl.arange(0, PACKED_PER_VAL) * 4
+    quant_offset = (x_[:, None, :] >> offsets[None, :, None])  # (BLOCK_N, PACKED_PER_VAL, D // PACKED_PER_VAL)
+
+    quant_offset = tl.view(quant_offset, (BLOCK_N, BLOCK_DMODEL_PACKED * PACKED_PER_VAL))
+    # Trick - instead of converting int4 to float16 we view it as float16
+    # and then multiply by 32768 * 512 == 2**24
+    quant_offset = (quant_offset & 0xF).to(tl.uint16).to(tl.float16, bitcast=True)
+    quant_offset = (quant_offset * 32768.0).to(tl.float16)
+    scale_512 = scale * 512
+
+    dequant = quant_offset * scale_512 + shift
+    return dequant
+
+
+@triton.jit
+def _splitK_reduce(
+    Out_splitK,  # [B, H, split_k, Mq, K]
+    Metadata,  # [B, H, 2, split_k, M_ceil] contains [mi, li]
+    Out,  # [B, H, M, K]
+    LSE,  # [B, H, M]
+    stride_osk_zhg,
+    stride_osk_s,
+    stride_osk_m,
+    stride_osk_k,
+    stride_mzhg,
+    stride_m2,
+    stride_ms,
+    stride_mm,
+    stride_oz,
+    stride_oh,
+    stride_og,
+    stride_om,
+    stride_ok, # pylint: disable=unused-argument
+    stride_lse_zhg,
+    stride_lse_m, M_ceil: tl.constexpr, # pylint: disable=unused-argument
+    BLOCK_SIZE: tl.constexpr,
+    H: tl.constexpr,
+    G: tl.constexpr,
+    split_k: tl.constexpr,
+    splitK_pow2: tl.constexpr,
+    use_mask: tl.constexpr,
+    IS_CAUSAL: tl.constexpr,
+):
+    off_zhg = tl.program_id(0)
+    off_z = off_zhg // (H * G)
+    off_h = (off_zhg // G) % H
+    off_g = off_zhg % G
+    off_m = tl.program_id(1)
+    off_k = tl.program_id(2)
+
+    # read  chunk
+    spk_idx = tl.arange(0, splitK_pow2)
+    kidx = tl.arange(0, BLOCK_SIZE)
+
+    Metadata_ptr = Metadata + stride_mzhg * off_zhg + spk_idx * stride_ms + off_m * stride_mm
+
+    o_ptr = (Out_splitK + off_zhg * stride_osk_zhg + stride_osk_m * off_m + off_k * BLOCK_SIZE +
+             stride_osk_s * spk_idx[:, None] + kidx[None, :] * stride_osk_k)
+
+    # read max values of each splitK
+    if use_mask:
+        spk_mask = spk_idx < split_k
+        l_m = tl.load(Metadata_ptr, mask=spk_mask, other=float("-inf"))
+        l_sum = tl.load(Metadata_ptr + stride_m2, mask=spk_mask, other=0.0)
+        acc = tl.load(o_ptr, mask=spk_mask[:, None], other=0.0)
+    else:
+        l_m = tl.load(Metadata_ptr)
+        l_sum = tl.load(Metadata_ptr + stride_m2)
+        acc = tl.load(o_ptr)
+
+    g_m = tl.max(l_m, axis=0)
+
+    if IS_CAUSAL:
+        l_m_offset = l_m - g_m
+        alpha = tl.where(l_m_offset > float("-inf"), tl.math.exp2(l_m_offset), 0.0)
+    else:
+        alpha = tl.math.exp2(l_m - g_m)
+
+    # read sum
+    l_sum *= alpha
+    g_sum = tl.sum(l_sum, axis=0)
+    acc = acc * alpha[:, None]
+
+    if IS_CAUSAL:
+        # Avoid division by zero
+        g_sum_safe = tl.where(g_sum > 0, g_sum, 1.0)
+        acc_out = tl.sum(acc, axis=0) / g_sum_safe
+    else:
+        acc_out = tl.sum(acc, axis=0) / g_sum
+
+    # Store output
+    Out_ptr = (Out + stride_oz * off_z + stride_oh * off_h + stride_og * off_g + stride_om * off_m +
+               off_k * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE))
+    tl.store(Out_ptr, acc_out)
+
+    # Store lse
+    l_ptrs = LSE + off_zhg * stride_lse_zhg + off_m
+    if IS_CAUSAL:
+        lse = tl.where(g_sum > 0, (g_m + tl.math.log2(g_sum)) / 1.44269504, g_m)
+        tl.store(l_ptrs, lse)
+    else:
+        tl.store(l_ptrs, (g_m + tl.math.log2(g_sum)) / 1.44269504)
+
+
+def quantize_kv_int4(k: torch.Tensor, num_groups: int = 1) -> torch.Tensor:
+    # Scale and shift are such that quantization linearly maps
+    # int4 values range [0..15] to input values range min(k)..max(k)
+    # individually for every row
+    k = k.reshape(*k.shape[:-1], num_groups, k.shape[-1] // num_groups)
+    max_vals = torch.max(k, dim=-1, keepdim=True).values
+    min_vals = torch.min(k, dim=-1, keepdim=True).values
+    scale_k: torch.Tensor = (max_vals - min_vals) / 15
+
+    shift_k = torch.min(k, dim=-1, keepdim=True).values
+    scale_k = scale_k.to(torch.float16)
+    shift_k = shift_k.to(torch.float16)
+
+    in_bytes = ((k - shift_k.expand(k.shape)) / scale_k.expand(k.shape)) + 0.5
+    in_bytes = in_bytes.to(torch.uint8)
+    in_int4 = in_bytes & 0xF
+    in_int4_packed = in_int4[..., ::2] + (in_int4[..., 1::2] << 4)
+    scale_shift = torch.concat([scale_k.view(torch.uint8), shift_k.view(torch.uint8)], dim=-1)
+    k_quant = torch.concat(
+        [
+            scale_shift.flatten(start_dim=-2),
+            in_int4_packed.flatten(start_dim=-2),
+        ],
+        dim=-1,
+    ).view(torch.int16)
+    return k_quant
+
+
+def dequantize_kv_fp16(quant_k: torch.Tensor, num_groups: int = 1) -> torch.Tensor:
+    k_i16 = quant_k.view(torch.int16)
+    k_ui8 = k_i16.view(torch.uint8)
+
+    ss_size = num_groups * 4
+    scale_shift_ui8 = k_ui8[..., 0:ss_size]
+    scale_shift_ui8 = scale_shift_ui8.reshape(*scale_shift_ui8.shape[:-1], num_groups, 4)
+    scale = scale_shift_ui8[..., 0:2].view(torch.float16)
+    shift = scale_shift_ui8[..., 2:4].view(torch.float16)
+
+    kv_ui8 = k_ui8[..., ss_size:]
+    k_ui8 = kv_ui8.reshape(*kv_ui8.shape[:-1], num_groups, -1)
+    k1_i4 = k_ui8 & 0xF
+    k2_i4 = (k_ui8 & 0xF0) >> 4
+    k_shape = k1_i4.shape
+    k1_f16 = k1_i4.to(torch.float16) * scale.expand(k_shape) + shift.expand(k_shape)
+    k2_f16 = k2_i4.to(torch.float16) * scale.expand(k_shape) + shift.expand(k_shape)
+
+    out = torch.empty((*k1_f16.shape[:-1], k1_f16.shape[-1] * 2), dtype=torch.float16, device=quant_k.device)
+    out[..., ::2] = k1_f16
+    out[..., 1::2] = k2_f16
+    out = out.reshape(*k_shape[:-2], -1)
+
+    return out
+
+
+def get_split_k(B: int, G: int, H: int, Mk: int) -> int:
+    """Heuristic for the number of splits"""
+    bh = max(B * H, 1)  # NOTE: Handle B*h=0 case
+    split_k = max(Mk, 1024) // bh
+    max_chunk_size = 64
+    while split_k > 0 and Mk / split_k < max_chunk_size:
+        split_k = split_k // 2
+    while B * H * G * split_k >= 1024:
+        split_k = split_k // 2
+    split_k = min(split_k, 512)
+    split_k = max(split_k, 1)
+    return split_k
+
+def attention_decode_forward_triton_impl(q, k, v, sm_scale, causal, alibi_slopes, layout, cache_seqlens, cache_batch_idx, new_kv, k_new, v_new):
+    # kernel config
+    BLOCK_M = 16
+    BLOCK_N = 64
+    SPLIT_K = None
+    NUM_QUANT_GROUPS = 1 # pylint: disable=unused-variable
+
+    # kernels expects "bsghd"
+    original_layout = layout
+    if layout == "bshd":
+        q = q.unsqueeze(2)
+        k = k.unsqueeze(2)
+        v = v.unsqueeze(2)
+        if new_kv:
+            k_new = k_new.unsqueeze(2)
+            v_new = v_new.unsqueeze(2)
+        layout = "bsghd"
+    elif layout == "bhsd":
+        q = q.permute(0, 2, 1, 3).unsqueeze(2)
+        k = k.permute(0, 2, 1, 3).unsqueeze(2)
+        v = v.permute(0, 2, 1, 3).unsqueeze(2)
+        if new_kv:
+            k_new = k_new.permute(0, 2, 1, 3).unsqueeze(2)
+            v_new = v_new.permute(0, 2, 1, 3).unsqueeze(2)
+        layout = "bsghd"
+    elif layout == "bsghd":
+        pass
+    elif layout is None:
+        raise ValueError("Layout not given")
+    assert layout == "bsghd"
+
+    # get dims
+    batch_size, seqlen_q, n_group_q, heads_per_group_q, dim_q = q.shape
+    _, seqlen_k, n_group_k, heads_per_group_k, dim_k = k.shape # pylint: disable=unused-variable
+    _, seqlen_v, n_group_v, heads_per_group_v, dim_v = v.shape # pylint: disable=unused-variable
+
+    assert dim_q == dim_k == dim_v, f"Dimensions must match: {dim_q}, {dim_k}, {dim_v}"
+
+    # get padded size
+    dim_padded  = get_padded_headsize(dim_k)
+
+    # Handle MQA/GQA case
+    if heads_per_group_q > heads_per_group_k:
+        is_gqa = True
+    elif heads_per_group_q < heads_per_group_k:
+        raise ValueError("heads_per_group_q < heads_per_group_k")
+    else:
+        is_gqa = False
+
+    assert dim_k == dim_q, f"Keys have head dim {dim_k} but queries have head dim {dim_q}"
+
+    if SPLIT_K is not None:
+        split_k = SPLIT_K
+    else:
+        # Use heuristics
+        split_k = get_split_k(batch_size, n_group_q, heads_per_group_q, seqlen_k) # NOTE: should the split think about seqlens?
+
+    seqlen_q_ceil = (seqlen_q + BLOCK_M - 1) // BLOCK_M * BLOCK_M
+    out_splitk = torch.empty([batch_size * n_group_q * heads_per_group_q, split_k, seqlen_q_ceil, dim_padded], dtype=torch.float32, device=q.device)
+    metadata = torch.empty([batch_size * n_group_q * heads_per_group_q, 2, split_k, seqlen_q_ceil], dtype=torch.float32, device=q.device)
+    lse = torch.empty((batch_size * n_group_q * heads_per_group_q, seqlen_q), device=q.device, dtype=torch.float32)
+    grid = (triton.cdiv(seqlen_q, BLOCK_M), batch_size * n_group_q * heads_per_group_q, split_k)
+
+    num_warps = 1
+    split_size = (seqlen_k + split_k - 1) // split_k
+    use_cache_seqlens = cache_seqlens is not None
+
+    # TODO: enable quantization
+    _fwd_kernel_splitK[grid](
+        Q=q,
+        K=k,
+        V=v,
+        sm_scale=sm_scale,
+        Out_splitK=out_splitk,
+        Metadata=metadata,
+        K_new = k_new,
+        V_new = v_new,
+        Cache_seqlens=cache_seqlens,
+        Cache_batch_idx=cache_batch_idx,
+        Alibi_slopes=alibi_slopes,
+        **_strides(q, "qz", "qm", "qg", "qh", "qd"),
+        **_strides(k, "kz", "kn", "kg", "kh", "kd"),
+        **_strides(v, "vz", "vn", "vg", "vh", "vd"),
+        **_strides(out_splitk, "osk_zhg", "osk_s", "osk_m", "osk_d"),
+        **_strides(metadata, "mzhg", "m2", "ms", "mm"),
+        **_strides(k_new, "kn_z", "kn_n", "kn_g", "kn_h", "kn_d"),
+        **_strides(v_new, "vn_z", "vn_n", "vn_g", "vn_h", "vn_d"),
+        **_strides(alibi_slopes, "az", "ah"),
+        Z=batch_size,
+        H_q=heads_per_group_q,
+        H_kv=heads_per_group_k,
+        G_q=n_group_q,
+        N_CTX_Q=seqlen_q,
+        N_CTX_K=seqlen_k,
+        N_CTX_NEW=k_new.shape[1] if new_kv else None,
+        BLOCK_N_PER_SPLIT=split_size,
+        BLOCK_M=BLOCK_M,
+        BLOCK_N=BLOCK_N,
+        BLOCK_DMODEL=dim_padded,
+        ACTUAL_BLOCK_DMODEL=dim_k,
+        BOUNDS_CHECKS_N=(split_size % BLOCK_N) > 0 or use_cache_seqlens,
+        USE_CACHE_SEQLENs=use_cache_seqlens,
+        USE_CACHE_BATCH_IDX=cache_batch_idx is not None,
+        NEW_KV=new_kv,
+        IS_GQA=is_gqa,
+        IS_CAUSAL=causal,
+        USE_ALIBI=False if alibi_slopes is None else True,
+        num_warps=num_warps,
+        num_stages=1,
+    )
+
+    out = torch.empty((batch_size, seqlen_q, n_group_q, heads_per_group_q, dim_padded), device=q.device, dtype=q.dtype)
+
+    # Merge together
+    splitK_pow2 = triton.next_power_of_2(split_k)
+    use_mask = splitK_pow2 > split_k
+    if batch_size * n_group_q * heads_per_group_q * seqlen_q >= 512:
+        k_block_num = 1
+    else:
+        k_block_num = 2
+    assert dim_padded % k_block_num == 0
+    k_block_size = dim_padded // k_block_num
+    grid = (batch_size * n_group_q * heads_per_group_q, seqlen_q, k_block_num)
+
+    _splitK_reduce[grid](
+        out_splitk,
+        metadata,
+        out,
+        lse,
+        **_strides(out_splitk, "osk_zhg", "osk_s", "osk_m", "osk_k"),
+        **_strides(metadata, "mzhg", "m2", "ms", "mm"),
+        **_strides(out, "oz", "om", "og", "oh", "ok"),
+        **_strides(lse, "lse_zhg", "lse_m"),
+        M_ceil=seqlen_q_ceil,
+        BLOCK_SIZE=k_block_size,
+        G=n_group_q,
+        H=heads_per_group_q,
+        # TODO: Tune num_warps
+        split_k=split_k,
+        splitK_pow2=splitK_pow2,
+        use_mask=use_mask,
+        IS_CAUSAL=causal,
+        num_warps=4)
+
+    lse = lse.reshape([batch_size, n_group_q, heads_per_group_q, seqlen_q])
+    if q.ndim == 4:
+        # BMGHK -> BMHK
+        assert n_group_q == 1
+        out = out[:, :, 0]
+        lse = lse[:, 0]
+    if seqlen_k == 0:
+        out.zero_()
+    out = out.reshape(batch_size, heads_per_group_q * n_group_q, -1, dim_padded).contiguous()
+
+    # output is batch_size, heads_per_group_q * group_q, seqlen_q, dim_q
+    if original_layout == "bshd":
+        # out=out.transpose(1, 2).contiguous() # this screws up heads and data.
+        # the data is laid out properly. Just need to reshape dims
+        out = out.reshape(batch_size, seqlen_q, -1, dim_padded)
+
+    return out.narrow(-1, 0, dim_k), lse
diff --git a/modules/flash_attn_triton_amd/fwd_prefill.py b/modules/flash_attn_triton_amd/fwd_prefill.py
new file mode 100644
index 000000000..3e2cd32af
--- /dev/null
+++ b/modules/flash_attn_triton_amd/fwd_prefill.py
@@ -0,0 +1,634 @@
+import torch
+import triton
+import triton.language as tl
+from modules.flash_attn_triton_amd.utils import get_shape_from_layout, get_strides_from_layout, is_cdna, is_rdna, AUTOTUNE
+
+
+@triton.jit
+def cdiv_fn(x, y):
+    return (x + y - 1) // y
+
+
+@triton.jit
+def dropout_offsets(philox_seed, philox_offset, dropout_p, m, n, stride): # pylint: disable=unused-argument
+    ms = tl.arange(0, m)
+    ns = tl.arange(0, n)
+    return philox_offset + ms[:, None] * stride + ns[None, :]
+
+
+@triton.jit
+def dropout_rng(philox_seed, philox_offset, dropout_p, m, n, stride):
+    rng_offsets = dropout_offsets(philox_seed, philox_offset, dropout_p, m, n, stride).to(tl.uint32)
+    # TODO: use tl.randint for better performance
+    return tl.rand(philox_seed, rng_offsets)
+
+
+@triton.jit
+def dropout_mask(philox_seed, philox_offset, dropout_p, m, n, stride):
+    rng_output = dropout_rng(philox_seed, philox_offset, dropout_p, m, n, stride)
+    rng_keep = rng_output > dropout_p
+    return rng_keep
+
+
+# Convenience function to load with optional boundary checks.
+# "First" is the major dim, "second" is the minor dim.
+@triton.jit
+def load_fn(ptrs, offset_first, offset_second, boundary_first, boundary_second):
+    if offset_first is not None and offset_second is not None:
+        mask = (offset_first[:, None] < boundary_first) & \
+               (offset_second[None, :] < boundary_second)
+        tensor = tl.load(ptrs, mask=mask, other=0.0)
+    elif offset_first is not None:
+        mask = offset_first[:, None] < boundary_first
+        tensor = tl.load(ptrs, mask=mask, other=0.0)
+    elif offset_second is not None:
+        mask = offset_second[None, :] < boundary_second
+        tensor = tl.load(ptrs, mask=mask, other=0.0)
+    else:
+        tensor = tl.load(ptrs)
+    return tensor
+
+
+@triton.jit
+def compute_alibi_block(alibi_slope, seqlen_q, seqlen_k, offs_m, offs_n, transpose=False):
+    # when seqlen_k and seqlen_q are different we want the diagonal to stick to the bottom right of the attention matrix
+    # for casual mask we want something like this where (1 is kept and 0 is masked)
+    # seqlen_q = 2 and seqlen_k = 5
+    #   1 1 1 1 0
+    #   1 1 1 1 1
+    # seqlen_q = 5 and seqlen_k = 2
+    #        0 0
+    #        0 0
+    #        0 0
+    #        1 0
+    #        1 1
+    # for alibi the diagonal is 0 indicating no penalty for attending to that spot and increasing penalty for attending further from the diagonal
+    # e.g. alibi_slope = 1, seqlen_q = 2, seqlen_k = 5, offs_m = [0, 1, 2, 3], offs_n = [0, 1, 2, 3, 4], transpose = False
+    # 1. offs_m[:,None] = [[0],
+    #                       [1],
+    # 2. offs_m[:,None] + seqlen_k = [[5],
+    #                                  [6],
+    # 3. offs_m[:,None] + seqlen_k - seqlen_q = [[3],
+    #                                             [4],
+    # 4. offs_m[:,None] + seqlen_k - seqlen_q - offs_n[None,:] = [[3], - [[0, 1, 2, 3, 4]] =  [[ 3, 2, 1, 0,-1],
+    #                                                            [4],                           [ 4, 3, 2, 1, 0]]
+    # 5. -1 * alibi_slope * tl.abs(relative_pos_block) = [[ -3, -2, -1, 0,-1],
+    #                                                     [ -4, -3, -2, -1, 0]],
+    relative_pos_block = offs_m[:, None] + seqlen_k - seqlen_q - offs_n[None, :]
+    alibi_block = -1 * alibi_slope * tl.abs(relative_pos_block)
+    if transpose:
+        return alibi_block.T
+    else:
+        return alibi_block
+
+
+@triton.jit
+def _attn_fwd_inner(acc, l_i, m_i, q, k_ptrs, v_ptrs, bias_ptrs, stride_kn, stride_vk, stride_bn, start_m,
+                    actual_seqlen_k, actual_seqlen_q, dropout_p, philox_seed, batch_philox_offset, exp_scores_ptrs,
+                    block_min, block_max, offs_n_causal, masked_blocks, n_extra_tokens, alibi_slope, score_ptrs, scores_scaled_shifted_ptrs, # pylint: disable=unused-argument
+                    IS_CAUSAL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_DMODEL: tl.constexpr, BLOCK_N: tl.constexpr,
+                    OFFS_M: tl.constexpr, OFFS_N: tl.constexpr, PRE_LOAD_V: tl.constexpr, MASK_STEPS: tl.constexpr,
+                    ENABLE_DROPOUT: tl.constexpr, PADDED_HEAD: tl.constexpr,
+                    ACTUAL_BLOCK_DMODEL: tl.constexpr, SM_SCALE: tl.constexpr, USE_EXP2: tl.constexpr,
+                    RETURN_SCORES: tl.constexpr):
+    if USE_EXP2:
+        RCP_LN2: tl.constexpr = 1.4426950408889634
+
+    # loop over k, v, and update accumulator
+    for start_n in range(block_min, block_max, BLOCK_N):
+        # For padded blocks, we will overrun the tensor size if
+        # we load all BLOCK_N. For others, the blocks are all within range.
+        if MASK_STEPS:
+            k_offs_n = start_n + tl.arange(0, BLOCK_N)
+        else:
+            k_offs_n = None
+        k_offs_k = None if not PADDED_HEAD else tl.arange(0, BLOCK_DMODEL)
+        k = load_fn(k_ptrs, k_offs_k, k_offs_n, ACTUAL_BLOCK_DMODEL, actual_seqlen_k)
+        if PRE_LOAD_V:
+            # We can use the same offsets as k, just with dims transposed.
+            v = load_fn(v_ptrs, k_offs_n, k_offs_k, actual_seqlen_k, ACTUAL_BLOCK_DMODEL)
+        qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
+        # We start from end of seqlen_k so only the first iteration would need
+        # to be checked for padding if it is not a multiple of block_n
+        # TODO: This can be optimized to only be true for the padded block.
+        if MASK_STEPS:
+            # If this is the last block / iteration, we want to
+            # mask if the sequence length is not a multiple of block size
+            # a solution is to always do BLOCK_M // BLOCK_N + 1 steps if not is_modulo_mn.
+            # last step might get wasted but that is okay. check if this masking works For
+            # that case.
+            if (start_n + BLOCK_N == block_max) and (n_extra_tokens != 0):
+                boundary_m = tl.full([BLOCK_M], actual_seqlen_k, dtype=tl.int32)
+                size_n = start_n + OFFS_N[None, :]
+                mask = size_n < boundary_m[:, None]
+                qk = tl.where(mask, qk, float("-inf"))
+
+        # -- compute qk ----
+        qk += tl.dot(q, k)
+        qk_scaled =  qk * SM_SCALE
+        if RETURN_SCORES:
+            score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ((start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k)
+            tl.store(score_ptrs, qk_scaled, mask=score_mask)
+
+        if IS_CAUSAL:
+            causal_boundary = start_n + offs_n_causal
+            causal_mask = OFFS_M[:, None] >= causal_boundary[None, :]
+            qk_scaled = tl.where(causal_mask, qk_scaled, float("-inf"))
+        if bias_ptrs is not None:
+            bias_offs_n = start_n + tl.arange(0, BLOCK_N) if MASK_STEPS else None
+            bias = load_fn(bias_ptrs, OFFS_M, bias_offs_n, actual_seqlen_q, actual_seqlen_k)
+            qk_scaled += bias
+
+        if alibi_slope is not None:
+            # Compute the global position of each token within the sequence
+            global_m_positions = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
+            global_n_positions = start_n + tl.arange(0, BLOCK_N)
+            alibi_block = compute_alibi_block(alibi_slope, actual_seqlen_q, actual_seqlen_k, global_m_positions,
+                                              global_n_positions)
+            qk_scaled += alibi_block
+        # get max scores so far
+        m_ij = tl.maximum(m_i, tl.max(qk_scaled, 1))
+
+        # scale and subtract max
+        q_shifted = qk_scaled - m_ij[:, None]
+        if RETURN_SCORES:
+            # NOTE: the returned score is not the same as the reference because we need to adjust as we find new maxes per block. We are not doing that
+            scores_scaled_shifted_mask = (OFFS_M[:, None] < actual_seqlen_q) & ((start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k)
+            tl.store(scores_scaled_shifted_ptrs, q_shifted, mask=scores_scaled_shifted_mask)
+
+        # Compute scaled QK and softmax probabilities
+        if USE_EXP2:
+            p = tl.math.exp2(q_shifted * RCP_LN2)
+        else:
+            p = tl.math.exp(q_shifted)
+
+        # CAVEAT: Must update l_ij before applying dropout
+        l_ij = tl.sum(p, 1)
+        if ENABLE_DROPOUT:
+            philox_offset = batch_philox_offset + start_m * BLOCK_M * actual_seqlen_k + start_n - BLOCK_N
+            keep = dropout_mask(philox_seed, philox_offset, dropout_p, BLOCK_M, BLOCK_N, actual_seqlen_k)
+            if RETURN_SCORES:
+                # NOTE: the returned score is not the same as the reference because we need to adjust as we find new maxes per block. We are not doing that
+                exp_score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ((start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k)
+                tl.store(exp_scores_ptrs, tl.where(keep, p, -p), mask=exp_score_mask)
+            p = tl.where(keep, p, 0.0)
+        elif RETURN_SCORES:
+            # NOTE: the returned score is not the same as the reference because we need to adjust as we find new maxes per block. We are not doing that
+            exp_score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ((start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k)
+            tl.store(exp_scores_ptrs, p, mask=exp_score_mask)
+
+        # -- update output accumulator --
+        # alpha is an adjustment factor for acc and li as we loop and find new maxes
+        # store the diff in maxes to adjust acc and li as we discover new maxes
+        m_diff = m_i - m_ij
+        if USE_EXP2:
+            alpha = tl.math.exp2(m_diff * RCP_LN2)
+        else:
+            alpha = tl.math.exp(m_diff)
+        acc = acc * alpha[:, None]
+        v = None
+        if not PRE_LOAD_V:
+            v = load_fn(v_ptrs, k_offs_n, k_offs_k, actual_seqlen_k, ACTUAL_BLOCK_DMODEL)
+        # -- update m_i and l_i
+        l_i = l_i * alpha + l_ij
+        # update m_i and l_i
+        m_i = m_ij
+        acc += tl.dot(p.to(v.type.element_ty), v)
+        k_ptrs += BLOCK_N * stride_kn
+        v_ptrs += BLOCK_N * stride_vk
+        if bias_ptrs is not None:
+            bias_ptrs += BLOCK_N * stride_bn
+        if RETURN_SCORES:
+            score_ptrs += BLOCK_N
+            scores_scaled_shifted_ptrs += BLOCK_N
+            exp_scores_ptrs += BLOCK_N
+    return acc, l_i, m_i
+
+
+def get_cdna_autotune_configs():
+    return [
+        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'waves_per_eu': 3, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'waves_per_eu': 1, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 32, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'waves_per_eu': 1, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+        # Fall-back config.
+        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 16, 'waves_per_eu': 1, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=4),
+    ], ['IS_CAUSAL', 'dropout_p', 'MAX_SEQLENS_Q', 'MAX_SEQLENS_K', 'ACTUAL_BLOCK_DMODEL', 'VARLEN', 'HQ', 'HK']
+
+
+def get_rdna_autotune_configs():
+    return [
+        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'waves_per_eu': 4, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 16, 'waves_per_eu': 4, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        triton.Config({'BLOCK_M': 32, 'BLOCK_N': 16, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 16, 'waves_per_eu': 4, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 16, 'waves_per_eu': 2, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+        # Fall-back config.
+        triton.Config({'BLOCK_M': 16, 'BLOCK_N': 16, 'waves_per_eu': 1, 'PRE_LOAD_V': False}, num_stages=1,
+                      num_warps=2),
+    ], ['IS_CAUSAL', 'dropout_p', 'MAX_SEQLENS_Q', 'MAX_SEQLENS_K', 'ACTUAL_BLOCK_DMODEL', 'VARLEN', 'HQ', 'HK']
+
+
+def get_autotune_configs():
+    if AUTOTUNE:
+        if is_rdna():
+            return get_rdna_autotune_configs()
+        elif is_cdna():
+            return get_cdna_autotune_configs()
+        else:
+            raise ValueError("Unknown Device Type")
+    else:
+        return [
+            triton.Config(
+                {"BLOCK_M": 64, "BLOCK_N": 64, "waves_per_eu": 1, "PRE_LOAD_V": False},
+                num_stages=1,
+                num_warps=4,
+            ),
+        ], [
+            "IS_CAUSAL",
+            "dropout_p",
+            "MAX_SEQLENS_Q",
+            "MAX_SEQLENS_K",
+            "ACTUAL_BLOCK_DMODEL",
+            "VARLEN",
+            "HQ",
+            "HK",
+        ]
+
+
+autotune_configs, autotune_keys = get_autotune_configs()
+
+@triton.autotune(
+    configs=autotune_configs,
+    key=autotune_keys,
+    # use_cuda_graph=True,
+)
+@triton.jit
+def attn_fwd(Q, K, V, bias, SM_SCALE: tl.constexpr, LSE, Out, stride_qz, stride_qh, stride_qm, stride_qk,
+             stride_kz, stride_kh, stride_kn, stride_kk, stride_vz, stride_vh, stride_vk, stride_vn,
+             stride_oz, stride_oh, stride_om, stride_on, stride_bz, stride_bh, stride_bm, stride_bn, stride_az, stride_ah, # pylint: disable=unused-argument
+             stride_sz, stride_sh, stride_sm, stride_sn, stride_lse_z, stride_lse_h, stride_lse_m, cu_seqlens_q, cu_seqlens_k,
+             dropout_p, philox_seed, philox_offset_base, scores, scores_scaled_shifted, exp_scores, alibi_slopes,  HQ: tl.constexpr,
+             HK: tl.constexpr, ACTUAL_BLOCK_DMODEL: tl.constexpr, MAX_SEQLENS_Q: tl.constexpr,
+             MAX_SEQLENS_K: tl.constexpr, VARLEN: tl.constexpr, IS_CAUSAL: tl.constexpr, BLOCK_M: tl.constexpr,
+             BLOCK_DMODEL: tl.constexpr, BLOCK_N: tl.constexpr, PRE_LOAD_V: tl.constexpr, USE_BIAS: tl.constexpr,
+             ENABLE_DROPOUT: tl.constexpr, RETURN_SCORES: tl.constexpr, USE_ALIBI: tl.constexpr, USE_EXP2: tl.constexpr):
+    start_m = tl.program_id(0)
+    off_h_q = tl.program_id(1)
+    off_z = tl.program_id(2)
+    offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
+    offs_n = tl.arange(0, BLOCK_N)
+    offs_d = tl.arange(0, BLOCK_DMODEL)
+    if VARLEN:
+        cu_seqlens_q_start = tl.load(cu_seqlens_q + off_z)
+        cu_seqlens_q_end = tl.load(cu_seqlens_q + off_z + 1)
+        # print("cu_seqlens_q_start:", cu_seqlens_q_start)
+
+        seqlen_q = cu_seqlens_q_end - cu_seqlens_q_start
+        # We have a one-size-fits-all grid in id(0). Some seqlens might be too
+        # small for all start_m so for those we return early.
+        if start_m * BLOCK_M > seqlen_q:
+            return
+        cu_seqlens_k_start = tl.load(cu_seqlens_k + off_z)
+        cu_seqlens_k_end = tl.load(cu_seqlens_k + off_z + 1)
+        seqlen_k = cu_seqlens_k_end - cu_seqlens_k_start
+    else:
+        cu_seqlens_q_start = 0
+        cu_seqlens_k_start = 0
+        seqlen_q = MAX_SEQLENS_Q
+        seqlen_k = MAX_SEQLENS_K
+
+    # Now we compute whether we need to exit early due to causal masking.
+    # This is because for seqlen_q > seqlen_k, M rows of the attn scores
+    # are completely masked, resulting in 0s written to the output, and
+    # inf written to LSE. We don't need to do any GEMMs in this case.
+    # This block of code determines what N is, and if this WG is operating
+    # on those M rows.
+    n_blocks = cdiv_fn(seqlen_k, BLOCK_N)
+    if IS_CAUSAL:
+        # If seqlen_q == seqlen_k, the attn scores are a square matrix.
+        # If seqlen_q != seqlen_k, attn scores are rectangular which means
+        # the causal mask boundary is bottom right aligned, and ends at either
+        # the top edge (seqlen_q < seqlen_k) or left edge.
+        # This captures the decrease in n_blocks if we have a rectangular attn matrix
+        n_blocks_seqlen = cdiv_fn((start_m + 1) * BLOCK_M + seqlen_k - seqlen_q, BLOCK_N)
+        # This is what adjusts the block_max for the current WG, only
+        # if IS_CAUSAL. Otherwise we want to always iterate through all n_blocks
+        n_blocks = min(n_blocks, n_blocks_seqlen)
+        # If we have no blocks after adjusting for seqlen deltas, this WG is part of
+        # the blocks that are all 0. We exit early.
+        if n_blocks <= 0:
+            o_offset = Out + off_z * stride_oz + off_h_q * stride_oh + cu_seqlens_q_start * stride_om
+            o_ptrs = o_offset + offs_m[:, None] * stride_om + offs_d[None, :] * stride_on
+            acc = tl.zeros([BLOCK_M, BLOCK_DMODEL], dtype=Out.type.element_ty)
+            o_ptrs_mask = offs_m[:, None] < seqlen_q
+            # We still need to write 0s to the result
+            tl.store(o_ptrs, acc, mask=o_ptrs_mask)
+            # The tensor allocated for L is based on MAX_SEQLENS_Q as that is
+            # statically known.
+            l_offset = LSE + off_z * stride_lse_z + off_h_q * stride_lse_h + cu_seqlens_q_start * stride_lse_m
+            l_ptrs = l_offset + offs_m * stride_lse_m 
+
+            l = tl.full([BLOCK_M], value=0.0, dtype=tl.float32)
+
+            # mask_m_offsets = start_m + tl.arange(0, BLOCK_M)
+            # lse_mask = mask_m_offsets < causal_start_idx
+            # softmax_lse = tl.where(lse_mask, 0.0, softmax_lse)
+            l_ptrs_mask = offs_m < MAX_SEQLENS_Q
+            tl.store(l_ptrs, l, mask=l_ptrs_mask)
+            # TODO: Should dropout and return encoded softmax be handled here too?
+            return
+
+    # If MQA / GQA, set the K and V head offsets appropriately.
+    GROUP_SIZE: tl.constexpr = HQ // HK
+    if GROUP_SIZE != 1:
+        off_h_k = off_h_q // GROUP_SIZE
+    else:
+        off_h_k = off_h_q
+
+    n_extra_tokens = 0
+    # print("n_extra_tokens:", n_extra_tokens)
+    # print("seqlen_k:", seqlen_k)
+    # print("BLOCK_N:", BLOCK_N)
+    # return
+    if seqlen_k < BLOCK_N:
+        n_extra_tokens = BLOCK_N - seqlen_k
+    elif seqlen_k % BLOCK_N:
+        n_extra_tokens = seqlen_k % BLOCK_N
+    PADDED_HEAD: tl.constexpr = ACTUAL_BLOCK_DMODEL != BLOCK_DMODEL
+
+    # Compute pointers for all the tensors used in this kernel.
+    q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + cu_seqlens_q_start * stride_qm
+    q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qk
+    k_offset = K + off_z * stride_kz + off_h_k * stride_kh + cu_seqlens_k_start * stride_kn
+    k_ptrs = k_offset + offs_d[:, None] * stride_kk + offs_n[None, :] * stride_kn
+    v_offset = V + off_z * stride_vz + off_h_k * stride_vh + cu_seqlens_k_start * stride_vk
+    v_ptrs = v_offset + offs_n[:, None] * stride_vk + offs_d[None, :] * stride_vn
+    if USE_BIAS:
+        # Note: this might get large enough to overflow on some configs
+        bias_offset = off_h_q * stride_bh
+        bias_ptrs = bias + bias_offset + offs_m[:, None] * stride_bm + offs_n[None, :] * stride_bn
+    else:
+        bias_ptrs = None
+
+    if USE_ALIBI:
+        a_offset = off_z * stride_az + off_h_q * stride_ah
+        alibi_slope = tl.load(alibi_slopes + a_offset)
+    else:
+        alibi_slope = None
+
+    if RETURN_SCORES:
+        scores_offset = scores + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm
+        score_ptrs = scores_offset + offs_m[:, None] * stride_sm + offs_n[None, :] * stride_sn
+
+        scores_scaled_shifted_offset = scores_scaled_shifted + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm
+        scores_scaled_shifted_ptrs = scores_scaled_shifted_offset + offs_m[:, None] * stride_sm + offs_n[None, :] * stride_sn
+
+        exp_scores_offset = exp_scores + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm
+        exp_scores_ptrs = exp_scores_offset + offs_m[:, None] * stride_sm + offs_n[None, :] * stride_sn
+    else:
+        score_ptrs = None
+        scores_scaled_shifted_ptrs = None
+        exp_scores_ptrs = None
+
+    if ENABLE_DROPOUT:
+        off_hz = off_z * HQ + off_h_q
+        batch_philox_offset = philox_offset_base + off_hz * seqlen_q * seqlen_k
+    else:
+        batch_philox_offset = 0
+    # initialize pointer to m and l
+    m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32)
+    l_i = tl.full([BLOCK_M], 1.0, dtype=tl.float32)
+    acc = tl.zeros([BLOCK_M, BLOCK_DMODEL], dtype=tl.float32)
+    # Q is loaded once at the beginning and shared by all N blocks.
+    q_ptrs_mask = offs_m[:, None] < seqlen_q
+    if PADDED_HEAD:
+        q_ptrs_mask = q_ptrs_mask & (offs_d[None, :] < ACTUAL_BLOCK_DMODEL)
+    q = tl.load(q_ptrs, mask=q_ptrs_mask, other=0.0)
+
+    # Here we compute how many full and masked blocks we have.
+    padded_block_k = n_extra_tokens != 0
+    is_modulo_mn = not padded_block_k and (seqlen_q % BLOCK_M == 0)
+    if IS_CAUSAL:
+        # There are always at least BLOCK_M // BLOCK_N masked blocks.
+        # Additionally there might be one more due to dissimilar seqlens.
+        masked_blocks = BLOCK_M // BLOCK_N + (not is_modulo_mn)
+    else:
+        # Padding on Q does not need to be masked in the FA loop.
+        masked_blocks = padded_block_k
+    # if IS_CAUSAL, not is_modulo_mn does not always result in an additional block.
+    # In this case we might exceed n_blocks so pick the min.
+    masked_blocks = min(masked_blocks, n_blocks)
+    n_full_blocks = n_blocks - masked_blocks
+    block_min = 0
+    block_max = n_blocks * BLOCK_N
+    # Compute for full blocks. Here we set causal to false regardless of its actual
+    # value because there is no masking. Similarly we do not need padding.
+    if n_full_blocks > 0:
+        block_max = (n_blocks - masked_blocks) * BLOCK_N
+        acc, l_i, m_i = _attn_fwd_inner(acc, l_i, m_i, q, k_ptrs, v_ptrs, bias_ptrs, stride_kn, stride_vk, stride_bn,
+                                        start_m, seqlen_k, seqlen_q, dropout_p, philox_seed, batch_philox_offset,
+                                        exp_scores_ptrs,
+                                        # _, _, offs_n_causal, masked_blocks, n_extra_tokens, _
+                                        block_min, block_max, 0, 0, 0, alibi_slope, score_ptrs, scores_scaled_shifted_ptrs,
+                                        # IS_CAUSAL, ....
+                                        False, BLOCK_M, BLOCK_DMODEL, BLOCK_N, offs_m, offs_n,
+                                        # _, MASK_STEPS, ...
+                                        PRE_LOAD_V, False, ENABLE_DROPOUT, PADDED_HEAD,
+                                        ACTUAL_BLOCK_DMODEL, SM_SCALE,  USE_EXP2=USE_EXP2, RETURN_SCORES=RETURN_SCORES)
+        block_min = block_max
+        block_max = n_blocks * BLOCK_N
+
+    tl.debug_barrier()
+    # Remaining blocks, if any, are full / not masked.
+    if masked_blocks > 0:
+        if IS_CAUSAL:
+            offs_n_causal = offs_n + (seqlen_q - seqlen_k)
+        else:
+            offs_n_causal = 0
+        k_ptrs += n_full_blocks * BLOCK_N * stride_kn
+        v_ptrs += n_full_blocks * BLOCK_N * stride_vk
+        if USE_BIAS:
+            bias_ptrs += n_full_blocks * BLOCK_N * stride_bn
+        if RETURN_SCORES:
+            score_ptrs += n_full_blocks * BLOCK_N
+            scores_scaled_shifted_ptrs += n_full_blocks * BLOCK_N
+            exp_scores_ptrs += n_full_blocks * BLOCK_N
+        acc, l_i, m_i = _attn_fwd_inner(acc, l_i, m_i, q, k_ptrs, v_ptrs, bias_ptrs, stride_kn, stride_vk, stride_bn,
+                                        start_m, seqlen_k, seqlen_q, dropout_p, philox_seed, batch_philox_offset,
+                                        exp_scores_ptrs, block_min, block_max, offs_n_causal, masked_blocks,
+                                        n_extra_tokens, alibi_slope, score_ptrs, scores_scaled_shifted_ptrs,
+                                        IS_CAUSAL, BLOCK_M, BLOCK_DMODEL, BLOCK_N, offs_m, offs_n,
+                                        # _, MASK_STEPS, ...
+                                        PRE_LOAD_V, True, ENABLE_DROPOUT, PADDED_HEAD,
+                                        ACTUAL_BLOCK_DMODEL, SM_SCALE, USE_EXP2=USE_EXP2, RETURN_SCORES=RETURN_SCORES)
+    # epilogue
+    # This helps the compiler do Newton Raphson on l_i vs on acc which is much larger.
+    l_recip = 1 / l_i[:, None]
+    acc = acc * l_recip
+    if ENABLE_DROPOUT:
+        acc = acc / (1 - dropout_p)
+    # If seqlen_q > seqlen_k but the delta is not a multiple of BLOCK_M,
+    # then we have one block with a row of all NaNs which come from computing
+    # softmax over a row of all -infs (-inf - inf = NaN). We check for that here
+    # and store 0s where there are NaNs as these rows should've been zeroed out.
+    end_m_idx = (start_m + 1) * BLOCK_M
+    start_m_idx = start_m * BLOCK_M
+    causal_start_idx = seqlen_q - seqlen_k
+    acc = acc.to(Out.type.element_ty)
+    if IS_CAUSAL:
+        if causal_start_idx > start_m_idx and causal_start_idx < end_m_idx:
+            out_mask_boundary = tl.full((BLOCK_DMODEL, ), causal_start_idx, dtype=tl.int32)
+            mask_m_offsets = start_m_idx + tl.arange(0, BLOCK_M)
+            out_ptrs_mask = mask_m_offsets[:, None] >= out_mask_boundary[None, :]
+            z: tl.tensor = 0.0
+            acc = tl.where(out_ptrs_mask, acc, z.to(acc.type.element_ty))
+
+    # write back LSE(Log Sum Exponents), the log of the normalization constant
+    l_offset = LSE + off_z * stride_lse_z + off_h_q * stride_lse_h + cu_seqlens_q_start * stride_lse_m
+    l_ptrs = l_offset + offs_m * stride_lse_m 
+    if USE_EXP2:
+        RCP_LN2: tl.constexpr = 1.4426950408889634
+        LN2: tl.constexpr = 0.6931471824645996
+        # compute log-sum-exp in base 2 units
+        mi_base2 = m_i * RCP_LN2
+        softmax_lse = mi_base2 + tl.math.log2(l_i)
+        # convert back to natural units
+        softmax_lse *= LN2
+    else:
+        softmax_lse = m_i + tl.math.log(l_i)
+
+    if IS_CAUSAL:
+        # zero out nans caused by -infs when doing causal
+        lse_mask = (start_m_idx + tl.arange(0, BLOCK_M)) < causal_start_idx
+        softmax_lse = tl.where(lse_mask, 0.0, softmax_lse)
+
+    # If seqlen_q not multiple of BLOCK_M, we need to mask out the last few rows.
+    # This is only true for the last M block. For others, overflow_size will be -ve
+    overflow_size = end_m_idx - seqlen_q
+    if overflow_size > 0:
+        boundary = tl.full((BLOCK_M, ), BLOCK_M - overflow_size, dtype=tl.int32)
+        l_ptrs_mask = tl.arange(0, BLOCK_M) < boundary
+        tl.store(l_ptrs, softmax_lse, mask=l_ptrs_mask) # the log of the normalization constant
+    else:
+        tl.store(l_ptrs, softmax_lse) # the log of the normalization constant
+
+    # write back O
+    o_offset = Out + off_z * stride_oz + off_h_q * stride_oh + cu_seqlens_q_start * stride_om
+    o_ptrs = o_offset + offs_m[:, None] * stride_om + offs_d[None, :] * stride_on
+    o_ptrs_mask = tl.full([BLOCK_M, BLOCK_DMODEL], 1, dtype=tl.int1)
+    if overflow_size > 0:
+        o_ptrs_mask = o_ptrs_mask & (offs_m[:, None] < seqlen_q)
+    if PADDED_HEAD:
+        o_ptrs_mask = o_ptrs_mask & (offs_d[None, :] < ACTUAL_BLOCK_DMODEL)
+    tl.store(o_ptrs, acc.to(Out.dtype.element_ty), mask=o_ptrs_mask)
+
+
+def attention_prefill_forward_triton_impl(
+                                        q,
+                                        k,
+                                        v,
+                                        o,
+                                        sm_scale,
+                                        alibi_slopes,
+                                        causal,
+                                        bias,
+                                        dropout_p,
+                                        layout,
+                                        cu_seqlens_q,
+                                        cu_seqlens_k,
+                                        max_seqlens_q,
+                                        max_seqlens_k,
+                                        return_scores,
+                                        use_exp2):
+    # check if varlen
+    is_varlen = layout == "thd"
+
+    # NOTE: a large bias tensor leads to overflow during pointer arithmetic
+    if bias is not None:
+        assert bias.numel() < 2**31
+
+    batch, nheads_q, nheads_k, head_size, seqlen_q, seqlen_k = get_shape_from_layout(q, k, layout, cu_seqlens_q, cu_seqlens_k, max_seqlens_q, max_seqlens_k) # pylint: disable=unused-variable
+    q_strides, k_strides, v_strides, o_strides = get_strides_from_layout(q, k, v, o, layout)
+
+    # Get closest power of 2 over or equal to 32.
+    padded_d_model = 1 << (head_size - 1).bit_length()
+    # Smallest head_dim supported is 16. If smaller, the tile in the
+    # kernel is padded - there is no padding in memory for any dims.
+    padded_d_model = max(padded_d_model, 16)
+
+    grid = lambda META: (triton.cdiv(max_seqlens_q, META['BLOCK_M']), nheads_q, batch) # pylint: disable=unnecessary-lambda-assignment
+
+    if return_scores:
+        scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device,
+                                        dtype=torch.float32)
+        scores_scaled_shifted = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device,
+                                        dtype=torch.float32)
+        scores_strides = (scores.stride(0), scores.stride(1), scores.stride(2), scores.stride(3))
+    else:
+        scores = None
+        scores_scaled_shifted = None
+        scores_strides = (0, 0 , 0 , 0)
+
+    # exp_scores is used to validate dropout behavior vs the PyTorch SDPA math backend reference.  We zero this out
+    # to give a consistent starting point and then populate it with the output of softmax with the sign bit set according
+    # to the dropout mask. The resulting return allows this mask to be fed into the reference implementation for testing
+    # only.  This return holds no useful output aside from debugging.
+    if return_scores:
+        exp_scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device,
+                                        dtype=torch.float32)
+    else:
+        exp_scores = None
+
+    # stores LSE the log of the normalization constant / sum of expoential score(unnormalzied probablities)
+    if is_varlen:
+        softmax_lse = torch.empty((q.shape[0], nheads_q), device=q.device, dtype=torch.float32)
+        stride_lse_m, stride_lse_h = softmax_lse.stride()
+        stride_lse_z = 0
+    else:
+        softmax_lse = torch.empty((batch, nheads_q, max_seqlens_q), device=q.device, dtype=torch.float32)
+        stride_lse_z, stride_lse_h, stride_lse_m = softmax_lse.stride()
+
+    # Seed the RNG so we get reproducible results for testing.
+    philox_seed = 0x1BF52
+    philox_offset = 0x1D4B42
+
+    if bias is not None:
+        bias_strides = (bias.stride(0), bias.stride(1),bias.stride(2),
+                        bias.stride(3))
+    else:
+        bias_strides = (0, 0, 0, 0)
+
+    if alibi_slopes is not None:
+        alibi_strides = (alibi_slopes.stride(0), alibi_slopes.stride(1))
+    else:
+        alibi_strides = (0, 0)
+
+
+    attn_fwd[grid](q, k, v, bias, sm_scale, softmax_lse, o, *q_strides, *k_strides, *v_strides, *o_strides,
+                    *bias_strides, *alibi_strides, *scores_strides, stride_lse_z, stride_lse_h, stride_lse_m, cu_seqlens_q, cu_seqlens_k,
+                    dropout_p=dropout_p, philox_seed=philox_seed, philox_offset_base=philox_offset, scores=scores,
+                    scores_scaled_shifted=scores_scaled_shifted, exp_scores=exp_scores, alibi_slopes=alibi_slopes,
+                    HQ=nheads_q, HK=nheads_k, ACTUAL_BLOCK_DMODEL=head_size, MAX_SEQLENS_Q=max_seqlens_q,
+                    MAX_SEQLENS_K=max_seqlens_k, IS_CAUSAL=causal, VARLEN=is_varlen,
+                    BLOCK_DMODEL=padded_d_model, USE_BIAS=False if bias is None else True,
+                    USE_ALIBI=False if alibi_slopes is None else True, ENABLE_DROPOUT=dropout_p
+                    > 0.0, USE_EXP2=use_exp2, RETURN_SCORES=return_scores)
+
+    return o, softmax_lse, exp_scores, grid, head_size, philox_seed, philox_offset, scores, scores_scaled_shifted
diff --git a/modules/flash_attn_triton_amd/fwd_ref.py b/modules/flash_attn_triton_amd/fwd_ref.py
new file mode 100644
index 000000000..03e53efde
--- /dev/null
+++ b/modules/flash_attn_triton_amd/fwd_ref.py
@@ -0,0 +1,258 @@
+import math
+import torch
+
+
+def attention_forward_core_ref_impl(q, k, v, sm_scale, causal, use_exp2):
+    # Compute attention scores
+    attention_scores = torch.matmul(q.to(torch.float32), k.transpose(-2, -1).to(torch.float32))
+
+    # Scale scores
+    attention_scaled_scores = sm_scale * attention_scores
+
+    # Apply causal mask if necessary
+    if causal:
+        L_q, L_k = q.shape[1], k.shape[1]
+        row_idx = torch.arange(L_q, device=q.device).unsqueeze(1)
+        col_idx = torch.arange(L_k, device=q.device).unsqueeze(0)
+        col_offset = L_q-L_k
+        causal_mask = row_idx >= (col_offset + col_idx)
+        # set -inf to places the causal mask is false
+        attention_scaled_scores = attention_scaled_scores.masked_fill(
+             torch.logical_not(causal_mask.unsqueeze(0)), float('-inf')
+        )
+
+
+    # Compute max for numerical stability
+    max_scores = torch.max(attention_scaled_scores, dim=-1, keepdim=True)[0]
+    if causal:
+        # Replace -inf in max_scores with zeros to avoid NaN in subtraction
+        max_scores = torch.where(
+            torch.isinf(max_scores), torch.zeros_like(max_scores), max_scores
+        )
+
+    # Shift scores
+    attention_shifted_scaled_scores = attention_scaled_scores - max_scores
+
+    # Exponentiate
+    if use_exp2:
+        RCP_LN = 1 / math.log(2)
+        exp_scores = torch.exp2(RCP_LN * attention_shifted_scaled_scores)
+    else:
+        exp_scores = torch.exp(attention_shifted_scaled_scores)
+
+    # Sum of exponentials
+    sum_exp_scores = torch.sum(exp_scores, dim=-1, keepdim=True)
+    if causal:
+        # if sum of exp scores is 0.0 it means scores where -inf, we cannot compute softmax and softmax_lse. Setting to 1 deals with -inf case cleanly
+        sum_exp_scores = torch.where(
+        sum_exp_scores == 0,
+        torch.ones_like(sum_exp_scores),
+        sum_exp_scores
+        )
+
+    # Compute softmax probabilities
+    softmax = exp_scores / sum_exp_scores
+
+    # Compute log-sum-exp
+    if use_exp2:
+        LN2 = math.log(2)
+        RCP_LN = 1 / math.log(2)
+        max_scores_base2 = max_scores * RCP_LN
+        softmax_lse_base2 = max_scores_base2 + torch.log2(sum_exp_scores)
+        softmax_lse = softmax_lse_base2 * LN2
+        softmax_lse.squeeze_(-1)
+    else:
+        softmax_lse = max_scores + torch.log(sum_exp_scores)
+        softmax_lse = softmax_lse.squeeze(-1)
+
+    # Compute output
+    o = torch.matmul(softmax, v.to(torch.float32)).to(torch.float16)
+
+    return o, softmax_lse, exp_scores, softmax, attention_shifted_scaled_scores, attention_scaled_scores, attention_scores
+
+def attention_vanilla_forward_pytorch_ref_impl(q, k, v, sm_scale, causal, layout, use_exp2):
+    """Compute reference output and softmax_lse using PyTorch's built-in function"""
+
+    # Ensure the layout is 'bhsd'
+    if layout == "bshd":
+        q = q.transpose(1, 2).contiguous()
+        k = k.transpose(1, 2).contiguous()
+        v = v.transpose(1, 2).contiguous()
+    elif layout != "bhsd":
+        raise ValueError(f"Unknown layout {layout}")
+
+    # Prepare tensors in [batch_size * num_heads, seq_len, head_dim] format
+    batch_size, num_heads, seq_len_q, head_dim = q.shape
+    seq_len_k = k.shape[2]
+
+    # Merge batch and heads dimensions
+    q = q.reshape(batch_size * num_heads, seq_len_q, head_dim)
+    k = k.reshape(batch_size * num_heads, seq_len_k, head_dim)
+    v = v.reshape(batch_size * num_heads, seq_len_k, head_dim)
+
+    # Call the core attention function
+    o, softmax_lse, exp_scores, softmax, attention_shifted_scaled_scores, attention_scaled_scores, attention_scores = attention_forward_core_ref_impl(
+        q, k, v, sm_scale, causal, use_exp2
+    )
+
+    # Reshape outputs back to [batch_size, num_heads, seq_len, head_dim]
+    o = o.reshape(batch_size, num_heads, seq_len_q, head_dim)
+    softmax_lse = softmax_lse.reshape(batch_size, num_heads, seq_len_q)
+    exp_scores = exp_scores.reshape(batch_size, num_heads, seq_len_q, seq_len_k)
+    softmax = softmax.reshape(batch_size, num_heads, seq_len_q, seq_len_k)
+    attention_shifted_scaled_scores = attention_shifted_scaled_scores.reshape(batch_size, num_heads, seq_len_q, seq_len_k)
+    attention_scaled_scores = attention_scaled_scores.reshape(batch_size, num_heads, seq_len_q, seq_len_k)
+    attention_scores = attention_scores.reshape(batch_size, num_heads, seq_len_q, seq_len_k)
+
+    # Restore original layout if necessary
+    if layout == "bshd":
+        o = o.transpose(1, 2)
+
+    return o, softmax_lse, exp_scores, softmax, attention_shifted_scaled_scores, attention_scaled_scores, attention_scores
+
+def attention_varlen_forward_pytorch_ref_impl(
+    q,
+    k,
+    v,
+    sm_scale,
+    causal,
+    layout,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q, max_seqlen_k, # pylint: disable=unused-argument
+    use_exp2
+):
+    # Ensure the layout is 'thd'
+    if layout != 'thd':
+        raise ValueError(f"Unsupported layout {layout}. Expected 'thd'.")
+
+    batch_size = cu_seqlens_q.shape[0] - 1
+    num_heads = q.shape[1]
+    head_dim = q.shape[2]
+
+    # Pre-allocate outputs
+    total_L_q = q.shape[0]
+    total_L_k = k.shape[0] # pylint: disable=unused-variable
+
+    o = torch.empty((total_L_q, num_heads, head_dim), dtype=q.dtype, device=q.device)
+    softmax_lse = torch.empty((total_L_q, num_heads), dtype=torch.float32, device=q.device)
+
+    for i in range(batch_size):
+        # Get the start and end indices for the current sequence
+        start_q = cu_seqlens_q[i].item()
+        end_q = cu_seqlens_q[i + 1].item()
+        start_k = cu_seqlens_k[i].item()
+        end_k = cu_seqlens_k[i + 1].item()
+
+        # Extract q_i, k_i, v_i
+        q_i = q[start_q:end_q, :, :]  # [L_q_i, num_heads, head_dim]
+        k_i = k[start_k:end_k, :, :]  # [L_k_i, num_heads, head_dim]
+        v_i = v[start_k:end_k, :, :]  # [L_k_i, num_heads, head_dim]
+
+        # Permute to [num_heads, L_q_i, head_dim]
+        q_i = q_i.permute(1, 0, 2)
+        k_i = k_i.permute(1, 0, 2)
+        v_i = v_i.permute(1, 0, 2)
+
+        # Call the core attention function for this sequence
+        (
+            o_i,
+            softmax_lse_i,
+            exp_scores_i,
+            softmax_i,
+            attention_shifted_scaled_scores_i,
+            attention_scaled_scores_i,
+            attention_scores_i,
+        ) = attention_forward_core_ref_impl(q_i, k_i, v_i, sm_scale, causal, use_exp2)
+
+        # Convert back to 'thd' layout and float16
+        o_i = o_i.permute(1, 0, 2).to(torch.float16)  # [L_q_i, num_heads, head_dim]
+
+        # Place outputs in pre-allocated tensors
+        o[start_q:end_q, :, :] = o_i
+        softmax_lse[start_q:end_q, :] = softmax_lse_i.transpose(0, 1)  # Transpose to [L_q_i, num_heads]
+
+        # For variable-sized outputs, map them into the preallocated tensors
+        # exp_scores_i: [num_heads, L_q_i, L_k_i] -> [L_q_i, num_heads, L_k_i]
+        exp_scores_i = exp_scores_i.permute(1, 0, 2)
+        softmax_i = softmax_i.permute(1, 0, 2)
+        attention_shifted_scaled_scores_i = attention_shifted_scaled_scores_i.permute(1, 0, 2)
+        attention_scaled_scores_i = attention_scaled_scores_i.permute(1, 0, 2)
+        attention_scores_i = attention_scores_i.permute(1, 0, 2)
+
+    return (
+        o,
+        softmax_lse,
+        None,
+        None,
+        None,
+        None,
+        None,
+    )
+
+
+def attention_forward_pytorch_ref_impl(
+    q,
+    k,
+    v,
+    sm_scale,
+    causal,
+    layout,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    max_seqlen_q,
+    max_seqlen_k,
+    use_exp2
+    ):
+     # compute reference
+    if layout == "thd":
+        (
+            o_ref,
+            softmax_lse_ref,
+            exp_scores_ref,
+            softmax_ref,
+            attention_shifted_scaled_scores_ref,
+            attention_scaled_scores_ref,
+            attention_scores_ref,
+        ) = attention_varlen_forward_pytorch_ref_impl(
+            q.clone(),
+            k.clone(),
+            v.clone(),
+            sm_scale,
+            causal,
+            layout,
+            cu_seqlens_q,
+            cu_seqlens_k,
+            max_seqlen_q,
+            max_seqlen_k,
+            use_exp2,
+        )
+    else:
+        (
+            o_ref,
+            softmax_lse_ref,
+            exp_scores_ref,
+            softmax_ref,
+            attention_shifted_scaled_scores_ref,
+            attention_scaled_scores_ref,
+            attention_scores_ref,
+        ) = attention_vanilla_forward_pytorch_ref_impl(
+            q.clone(), k.clone(), v.clone(), sm_scale, causal, layout, use_exp2
+        )
+
+    return (
+            o_ref,
+            softmax_lse_ref,
+            exp_scores_ref,
+            softmax_ref,
+            attention_shifted_scaled_scores_ref,
+            attention_scaled_scores_ref,
+            attention_scores_ref,
+    )
+
+
+def compute_alibi_tensor_ref(alibi_slopes, seqlen_q, seqlen_k):
+    q_idx = torch.arange(seqlen_q, dtype=torch.int32, device="cuda").unsqueeze(-1)  # (N_CTX_Q, 1)
+    k_idx = torch.arange(seqlen_k, dtype=torch.int32, device="cuda").unsqueeze(0)  # (1, N_CTX_K)
+    relative_pos = torch.abs(q_idx + seqlen_k - seqlen_q - k_idx)  # (N_CTX_Q, N_CTX_K)
+    return -1 * alibi_slopes.unsqueeze(-1).unsqueeze(-1) * relative_pos  # (Z, H, N_CTX_Q, N_CTX_K)
diff --git a/modules/flash_attn_triton_amd/interface_fa.py b/modules/flash_attn_triton_amd/interface_fa.py
new file mode 100644
index 000000000..72373d35f
--- /dev/null
+++ b/modules/flash_attn_triton_amd/interface_fa.py
@@ -0,0 +1,394 @@
+import os
+import torch
+from modules.flash_attn_triton_amd.fwd_prefill import attention_prefill_forward_triton_impl
+from modules.flash_attn_triton_amd.bwd_prefill import attention_prefill_backward_triton_impl
+from modules.flash_attn_triton_amd.fwd_decode import attention_decode_forward_triton_impl
+from modules.flash_attn_triton_amd.fwd_ref import attention_forward_pytorch_ref_impl
+from modules.flash_attn_triton_amd.bwd_ref import attention_backward_pytorch_ref_impl
+from modules.flash_attn_triton_amd.utils import MetaData, get_shape_from_layout
+
+
+USE_REF = os.environ.get('FLASH_ATTENTION_TRITON_AMD_REF', '0').lower() in ('1', 'true', 'yes')
+
+
+def fwd(q,
+    k,
+    v,
+    o,
+    alibi_slopes,
+    dropout_p,
+    softmax_scale,
+    causal,
+    window_size_left, window_size_right, softcap, # pylint: disable=unused-argument
+    return_softmax,
+    gen_ # pylint: disable=unused-argument
+):
+    if dropout_p != 0.0:
+        raise ValueError("dropout is not supported on AMD's Triton Backend yet")
+
+    if o is None:
+        o = torch.empty_like(q)
+
+    # Setup metadata
+    metadata = MetaData(sm_scale=softmax_scale)
+    metadata.max_seqlens_q = q.shape[1]
+    metadata.max_seqlens_k = k.shape[1]
+    metadata.layout = "bshd"
+    if return_softmax:
+        metadata.return_scores = True
+
+    batch, nheads_q, nheads_k, head_size, _, _ = get_shape_from_layout(q, k, metadata.layout) # pylint: disable=unused-variable
+
+    if causal:
+        metadata.need_causal()
+
+    if alibi_slopes is not None:
+        metadata.need_alibi(alibi_slopes, batch, nheads_q)
+
+    if dropout_p > 0.0:
+        metadata.need_dropout(dropout_p, return_softmax)
+
+    # Check arguments
+    metadata.check_args(q, k, v, o)
+    if USE_REF:
+        (output,
+        softmax_lse,
+        exp_scores,
+        _,
+        _,
+        _,
+        _) = attention_forward_pytorch_ref_impl(
+                                                q,
+                                                k,
+                                                v,
+                                                metadata.sm_scale,
+                                                metadata.causal,
+                                                metadata.layout,
+                                                metadata.cu_seqlens_q,
+                                                metadata.cu_seqlens_k,
+                                                metadata.max_seqlens_q,
+                                                metadata.max_seqlens_k,
+                                                metadata.use_exp2)
+        o.copy_(output)
+    else:
+        (_,
+        softmax_lse,
+        exp_scores,
+        _,
+        _,
+        _,
+        _,
+        _,
+        _) = attention_prefill_forward_triton_impl(
+                                                q,
+                                                k,
+                                                v,
+                                                o,
+                                                metadata.sm_scale,
+                                                metadata.alibi_slopes,
+                                                metadata.causal,
+                                                metadata.bias,
+                                                metadata.dropout_p,
+                                                metadata.layout,
+                                                metadata.cu_seqlens_q,
+                                                metadata.cu_seqlens_k,
+                                                metadata.max_seqlens_q,
+                                                metadata.max_seqlens_k,
+                                                metadata.return_scores,
+                                                metadata.use_exp2)
+
+    return o, softmax_lse, exp_scores, None
+
+
+def bwd(
+    dout,
+    q,
+    k,
+    v,
+    out,
+    softmax_lse,
+    dq,
+    dk,
+    dv,
+    alibi_slopes,
+    dropout_p,
+    softmax_scale,
+    causal,
+    window_size_left, window_size_right, softcap, deterministic, gen_, rng_state, # pylint: disable=unused-argument
+):
+    if dropout_p != 0.0:
+        raise ValueError("dropout is not supported on AMD yet")
+
+    if USE_REF:
+        dq_ref, dk_ref, dv_ref, delta_ref = attention_backward_pytorch_ref_impl(
+            dout,
+            q,
+            k,
+            v,
+            out,
+            softmax_lse,
+            softmax_scale,
+            causal,
+            "bshd",
+            None,
+            None,
+            None,
+            None,
+            False,
+        )
+        dq.copy_(dq_ref)
+        dk.copy_(dk_ref)
+        dv.copy_(dv_ref)
+        delta = delta_ref
+    else:
+        dq_triton, dk_triton, dv_triton, delta_triton, _, _ = attention_prefill_backward_triton_impl( # pylint: disable=unused-variable
+            dout,
+            q,
+            k,
+            v,
+            out,
+            softmax_lse,
+            dq,
+            dk,
+            dv,
+            softmax_scale,
+            alibi_slopes,
+            causal,
+            "bshd",
+            None,
+            None,
+            None,
+            None,
+            False,
+        )
+        delta = delta_triton
+
+    return dq, dk, dv, delta
+
+
+def varlen_fwd(
+    q,
+    k,
+    v,
+    o,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    seqused_k, leftpad_k, block_table_, # pylint: disable=unused-argument
+    alibi_slopes,\
+    max_seqlen_q,
+    max_seqlen_k,
+    dropout_p,
+    softmax_scale,
+    zero_tensors, # pylint: disable=unused-argument
+    causal,
+    window_size_left, window_size_right, softcap, # pylint: disable=unused-argument
+    return_softmax,
+    gen_ # pylint: disable=unused-argument
+):
+    if dropout_p != 0.0:
+        raise ValueError("dropout is not supported on AMD's Triton Backend yet")
+
+    if o is None:
+        o = torch.empty_like(q)
+
+    # Setup metadata
+    metadata = MetaData(sm_scale=softmax_scale)
+    if return_softmax:
+        metadata.return_scores = True
+    metadata.set_varlen_params(cu_seqlens_q, cu_seqlens_k)  # set layout to "thd" and other metdata
+
+    # get shapes
+    batch, nheads_q, nheads_k, head_size , seqlen_q, seqlen_k = get_shape_from_layout(q, k, metadata.layout, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k) # pylint: disable=unused-variable
+
+    if causal:
+        metadata.need_causal()
+
+    if alibi_slopes is not None:
+        metadata.need_alibi(alibi_slopes, batch, nheads_q)
+
+    if dropout_p > 0.0:
+        metadata.need_dropout(dropout_p, return_softmax)
+
+    # Check arguments
+    metadata.check_args(q, k, v, o)
+    if o is None:
+        o = torch.empty_like(q, dtype=v.dtype)
+
+    if USE_REF:
+        (output,
+        softmax_lse,
+        exp_scores,
+        _,
+        _,
+        _,
+        _) = attention_forward_pytorch_ref_impl(
+                                                q,
+                                                k,
+                                                v,
+                                                metadata.sm_scale,
+                                                metadata.causal,
+                                                metadata.layout,
+                                                metadata.cu_seqlens_q,
+                                                metadata.cu_seqlens_k,
+                                                metadata.max_seqlens_q,
+                                                metadata.max_seqlens_k,
+                                                metadata.use_exp2)
+        o.copy_(output)
+    else:
+        (_,
+        softmax_lse,
+        exp_scores,
+        _,
+        _,
+        _,
+        _,
+        _,
+        _) = attention_prefill_forward_triton_impl(
+                                                    q,
+                                                    k,
+                                                    v,
+                                                    o,
+                                                    metadata.sm_scale,
+                                                    metadata.alibi_slopes,
+                                                    metadata.causal,
+                                                    metadata.bias,
+                                                    metadata.dropout_p,
+                                                    metadata.layout,
+                                                    metadata.cu_seqlens_q,
+                                                    metadata.cu_seqlens_k,
+                                                    metadata.max_seqlens_q,
+                                                    metadata.max_seqlens_k,
+                                                    metadata.return_scores,
+                                                    metadata.use_exp2)
+
+    return o, softmax_lse, exp_scores, None
+
+
+def varlen_bwd(
+    dout,
+    q,
+    k,
+    v,
+    out,
+    softmax_lse,
+    dq,
+    dk,
+    dv,
+    cu_seqlens_q,
+    cu_seqlens_k,
+    alibi_slopes,
+    max_seqlen_q,
+    max_seqlen_k,
+    dropout_p,
+    softmax_scale,
+    zero_tensors, # pylint: disable=unused-argument
+    causal,
+    window_size_left, window_size_right, softcap, deterministic, gen_, rng_state, # pylint: disable=unused-argument
+):
+    if dropout_p != 0.0:
+        raise ValueError("dropout is not supported on AMD yet")
+
+    if USE_REF:
+        dq_ref, dk_ref, dv_ref, delta_ref = attention_backward_pytorch_ref_impl(
+            dout,
+            q,
+            k,
+            v,
+            out,
+            softmax_lse,
+            softmax_scale,
+            causal,
+            "thd",
+            cu_seqlens_q,
+            cu_seqlens_k,
+            max_seqlen_q,
+            max_seqlen_k,
+            False,
+        )
+        dq.copy_(dq_ref)
+        dk.copy_(dk_ref)
+        dv.copy_(dv_ref)
+        delta = delta_ref
+    else:
+        dq_triton, dk_triton, dv_triton, delta_triton, _, _ = attention_prefill_backward_triton_impl( # pylint: disable=unused-variable
+            dout,
+            q,
+            k,
+            v,
+            out,
+            softmax_lse,
+            dq,
+            dk,
+            dv,
+            softmax_scale,
+            alibi_slopes,
+            causal,
+            "thd",
+            cu_seqlens_q,
+            cu_seqlens_k,
+            max_seqlen_q,
+            max_seqlen_k,
+            False,
+        )
+        delta = delta_triton
+
+    return dq, dk, dv, delta
+
+
+def fwd_kvcache(
+    q,
+    k_cache,
+    v_cache,
+    k,
+    v,
+    cache_seqlens,
+    rotary_cos, rotary_sin, # pylint: disable=unused-argument
+    cache_batch_idx,
+    cache_leftpad, block_table, # pylint: disable=unused-argument
+    alibi_slopes,
+    out,
+    softmax_scale,
+    causal,
+    window_size_left, window_size_right, softcap, rotary_interleaved, num_splits, # pylint: disable=unused-argument
+):
+    if out is None:
+        out = torch.empty_like(q)
+
+    # fill metadata
+    metadata = MetaData(sm_scale=softmax_scale)
+    metadata.layout = "bshd"
+    metadata.max_seqlens_q = q.shape[1]
+    metadata.max_seqlens_k = k_cache.shape[1]
+    metadata.cache_seqlens = cache_seqlens
+    metadata.cache_batch_idx = cache_batch_idx
+
+    if k is not None and v is not None:
+        metadata.new_kv = True
+        metadata.seqlen_new = k.shape[1]
+        metadata.k_new = k
+        metadata.v_new = v
+
+    if causal:
+        metadata.need_causal()
+
+    if alibi_slopes is not None:
+        batch, _ , nheads_q, _= q.shape
+        metadata.need_alibi(alibi_slopes, batch, nheads_q)
+
+    # launch kernel
+    # TODO: pass output as an arg. Maybe we are copying output which is causing slow down
+    output, softmax_lse = attention_decode_forward_triton_impl(
+        q,
+        k_cache,
+        v_cache,
+        metadata.sm_scale,
+        metadata.causal,
+        metadata.alibi_slopes,
+        metadata.layout,
+        metadata.cache_seqlens,
+        metadata.cache_batch_idx,
+        metadata.new_kv,
+        metadata.k_new,
+        metadata.v_new,
+    )
+    return output, softmax_lse
diff --git a/modules/flash_attn_triton_amd/utils.py b/modules/flash_attn_triton_amd/utils.py
new file mode 100644
index 000000000..77384cff6
--- /dev/null
+++ b/modules/flash_attn_triton_amd/utils.py
@@ -0,0 +1,280 @@
+import os
+import torch
+import triton
+
+
+AUTOTUNE = os.environ.get('FLASH_ATTENTION_TRITON_AMD_AUTOTUNE', '0').lower() in ('1', 'true', 'yes')
+PERF = os.environ.get('FLASH_ATTENTION_TRITON_AMD_PERF', '0').lower() in ('1', 'true', 'yes')
+
+
+class MetaData():
+    cu_seqlens_q = None
+    cu_seqlens_k = None
+    max_seqlens_q = 0
+    max_seqlens_k = 0
+    bias = None
+    alibi_slopes = None
+    causal = False
+    num_contexts = 0
+    varlen = False
+    layout = None
+    cache_seqlens = None
+    cache_batch_idx = None
+    new_kv = False
+    seqlen_new = None
+    k_new = None
+    v_new = None
+    dropout_p, return_scores= 0.0, False
+    # NOTE: scale sm_scale by log_2(e) and use 2^x in the loop as we do not have native e^x support in HW.
+    use_exp2 = False
+
+    def __repr__(self) -> str:
+        return (f"MetaData(\n"
+                f"  sm_scale={self.sm_scale},\n"
+                f"  cu_seqlens_q={self.cu_seqlens_q},\n"
+                f"  cu_seqlens_k={self.cu_seqlens_k},\n"
+                f"  max_seqlens_q={self.max_seqlens_q},\n"
+                f"  max_seqlens_k={self.max_seqlens_k},\n"
+                f"  bias={self.bias},\n"
+                f"  alibi_slopes={self.alibi_slopes},\n"
+                f"  causal={self.causal},\n"
+                f"  num_contexts={self.num_contexts},\n"
+                f"  varlen={self.varlen},\n"
+                f"  layout={self.layout},\n"
+                f"  cache_seqlens={self.cache_seqlens},\n"
+                f"  cache_batch_idx={self.cache_batch_idx},\n"
+                f"  new_kv={self.new_kv},\n"
+                f"  seqlen_new={self.seqlen_new},\n"
+                f"  k_new={self.k_new},\n"
+                f"  v_new={self.v_new},\n"
+                f"  dropout_p={self.dropout_p},\n"
+                f"  return_scores={self.return_scores}\n"
+                f")")
+
+    def __init__(self, sm_scale=1.0):
+        self.sm_scale = sm_scale
+
+    def set_varlen_params(self, cu_seqlens_q, cu_seqlens_k):
+        self.varlen = True
+        self.layout = 'thd'
+        self.cu_seqlens_q = cu_seqlens_q
+        self.cu_seqlens_k = cu_seqlens_k
+        # Without "varlen", there should still be one sequence.
+        assert len(cu_seqlens_q) >= 2
+        assert len(cu_seqlens_q) == len(cu_seqlens_k)
+        self.num_contexts = len(cu_seqlens_q) - 1
+        for i in range(0, self.num_contexts):
+            self.max_seqlens_q = max(cu_seqlens_q[i + 1].item() - cu_seqlens_q[i].item(), self.max_seqlens_q)
+            self.max_seqlens_k = max(cu_seqlens_k[i + 1].item() - cu_seqlens_k[i].item(), self.max_seqlens_k)
+
+    def need_bias(self, bias, batch, nheads, seqlen_q, seqlen_k): # pylint: disable=unused-argument
+        assert bias.is_cuda
+        assert bias.dim() == 4
+        assert bias.shape[0] == 1
+        assert bias.shape[2:] == (seqlen_q, seqlen_k)
+        self.bias = bias
+
+    def need_alibi(self, alibi_slopes, batch, nheads):
+        assert alibi_slopes.is_cuda
+        assert alibi_slopes.dim() == 2
+        assert alibi_slopes.shape[0] == batch
+        assert alibi_slopes.shape[1] == nheads
+        self.alibi_slopes = alibi_slopes
+
+    def need_causal(self):
+        self.causal = True
+
+    def need_dropout(self, dropout_p, return_scores):
+        self.dropout_p = dropout_p
+        self.return_scores = return_scores
+
+    def check_args(self, q, k, v, o):
+        assert q.dim() == k.dim() and q.dim() == v.dim()
+
+        batch, nheads_q, nheads_k, head_size, _, _ = get_shape_from_layout(q, k, self.layout, self.cu_seqlens_q, self.cu_seqlens_k, self.max_seqlens_q, self.max_seqlens_k) # pylint: disable=unused-variable
+        if self.varlen:
+            assert q.dim() == 3
+            assert self.cu_seqlens_q is not None
+            assert self.cu_seqlens_k is not None
+            assert len(self.cu_seqlens_q) == len(self.cu_seqlens_k)
+            # TODO: Remove once bias is supported with varlen
+            assert self.bias is None
+            # TODO:Remove once dropout is supported with varlen
+            assert self.dropout_p == 0.0
+            # assert not self.return_scores
+        else:
+            assert q.dim() == 4
+            assert self.max_seqlens_q > 0 and self.max_seqlens_k > 0
+            assert self.cu_seqlens_q is None and self.cu_seqlens_k is None
+        assert k.shape == v.shape
+        assert q.shape[-1] == k.shape[-1] and q.shape[-1] == v.shape[-1]
+        # TODO: Change assert if we support qkl f8 and v f16
+        assert q.dtype == k.dtype and q.dtype == v.dtype
+        assert head_size <= 256
+        assert o.shape == q.shape
+        assert (nheads_q % nheads_k) == 0
+        assert self.layout is not None
+        assert self.layout == 'thd' or not self.varlen
+
+def input_helper(Z, HQ, HK, N_CTX_Q, N_CTX_K, D_HEAD, dtype, layout, device="cuda", DEBUG_INPUT=False):
+    torch.manual_seed(20)
+
+    # Initialize q, k, v
+    if layout == 'bhsd':
+        q_tensor_shape = (Z, HQ, N_CTX_Q, D_HEAD)
+        k_tensor_shape = (Z, HK, N_CTX_K, D_HEAD)
+    elif layout == 'bshd':
+        q_tensor_shape = (Z, N_CTX_Q, HQ, D_HEAD)
+        k_tensor_shape = (Z, N_CTX_K, HK, D_HEAD)
+    else:
+        assert False, f'Got unsupported tensor layout: {layout}'
+
+    q = None
+    k = None
+    v = None
+
+    if DEBUG_INPUT:
+        if layout == "bhsd":
+            q = torch.arange(N_CTX_Q, dtype=dtype, device=device).view(1, 1, N_CTX_Q, 1).expand(*q_tensor_shape).contiguous().requires_grad_()
+            k = torch.arange(N_CTX_K, dtype=dtype, device=device).view(1, 1, N_CTX_K, 1).expand(*k_tensor_shape).contiguous().requires_grad_()
+            v = torch.arange(N_CTX_K, dtype=dtype, device=device).view(1, 1, N_CTX_K, 1).expand(*k_tensor_shape).contiguous().requires_grad_()
+        elif layout == "bshd":
+            q = torch.arange(N_CTX_Q, dtype=dtype, device=device).view(1, N_CTX_Q, 1, 1).expand(*q_tensor_shape).contiguous().requires_grad_()
+            k = torch.arange(N_CTX_K, dtype=dtype, device=device).view(1, N_CTX_K, 1, 1).expand(*k_tensor_shape).contiguous().requires_grad_()
+            v = torch.arange(N_CTX_K, dtype=dtype, device=device).view(1, N_CTX_K, 1, 1).expand(*k_tensor_shape).contiguous().requires_grad_()
+    else:
+        q = torch.randn(q_tensor_shape, dtype=dtype, device=device, requires_grad=True)
+        k = torch.randn(k_tensor_shape, dtype=dtype, device=device, requires_grad=True)
+        v = torch.randn(k_tensor_shape, dtype=dtype, device=device, requires_grad=True)
+
+    if DEBUG_INPUT:
+        sm_scale = 1
+    else:
+        sm_scale = D_HEAD**-0.5
+    input_metadata = MetaData(sm_scale=sm_scale)
+    input_metadata.max_seqlens_q = N_CTX_Q
+    input_metadata.max_seqlens_k = N_CTX_K
+    input_metadata.layout = layout
+    return q, k, v, input_metadata
+
+
+def varlen_input_helper(Z, HQ, HK, N_CTX_Q, N_CTX_K, D_HEAD, dtype, device="cuda", equal_seqlens=False, DEBUG_INPUT=False):
+    torch.manual_seed(20)
+
+    # Random or equal sequence lengths based on 'equal_seqlens' flag
+    if not equal_seqlens:
+        max_seqlens_q = N_CTX_Q // Z
+        max_seqlens_k = N_CTX_K // Z
+        seqlens_q = torch.randint(1, max_seqlens_q + 1, (Z,), dtype=torch.int32)
+        seqlens_k = torch.randint(1, max_seqlens_k + 1, (Z,), dtype=torch.int32)
+    else:
+        seqlens_q = torch.full((Z,), N_CTX_Q // Z, dtype=torch.int32)
+        seqlens_k = torch.full((Z,), N_CTX_K // Z, dtype=torch.int32)
+
+    # Calculate cumulative sequence lengths
+    cu_seqlens_q = torch.cat([torch.tensor([0], dtype=torch.int32), seqlens_q.cumsum(dim=0)])
+    cu_seqlens_k = torch.cat([torch.tensor([0], dtype=torch.int32), seqlens_k.cumsum(dim=0)])
+    cu_seqlens_q = cu_seqlens_q.to(device=device).to(torch.int32)
+    cu_seqlens_k = cu_seqlens_k.to(device=device).to(torch.int32)
+
+    # Total lengths
+    total_q = cu_seqlens_q[-1].item()
+    total_k = cu_seqlens_k[-1].item()
+
+    if DEBUG_INPUT:
+        # Initialize q, k, v with deterministic values
+        q = torch.arange(total_q, dtype=dtype, device=device).view(total_q, 1, 1)
+        q = q.expand(total_q, HQ, D_HEAD).contiguous().requires_grad_()
+        k = torch.arange(total_k, dtype=dtype, device=device).view(total_k, 1, 1)
+        k = k.expand(total_k, HK, D_HEAD).contiguous().requires_grad_()
+        v = torch.arange(total_k, dtype=dtype, device=device).view(total_k, 1, 1)
+        v = v.expand(total_k, HK, D_HEAD).contiguous().requires_grad_()
+        sm_scale = 1
+    else:
+        # Initialize q, k, v with random values
+        q = torch.randn((total_q, HQ, D_HEAD), dtype=dtype, device=device).requires_grad_()
+        k = torch.randn((total_k, HK, D_HEAD), dtype=dtype, device=device).requires_grad_()
+        v = torch.randn((total_k, HK, D_HEAD), dtype=dtype, device=device).requires_grad_()
+        sm_scale = D_HEAD ** -0.5
+
+    input_metadata = MetaData(sm_scale=sm_scale)
+    input_metadata.set_varlen_params(cu_seqlens_q, cu_seqlens_k)
+    return q, k, v, input_metadata
+
+
+def get_shape_from_layout(q, k, layout, cu_seqlens_q = None, cu_seqlens_k = None, max_seqlen_q=None, max_seqlen_k=None):
+    if layout == 'bhsd':
+        batch_q, nheads_q, max_seqlen_q, head_size_q = q.shape
+        batch_k, nheads_k, max_seqlen_k, head_size_k = k.shape
+    elif layout == 'bshd':
+        batch_q, max_seqlen_q, nheads_q, head_size_q = q.shape
+        batch_k, max_seqlen_k, nheads_k, head_size_k = k.shape
+    elif  layout == 'thd':
+        batch_q, max_seqlen_q, nheads_q, head_size_q = len(cu_seqlens_q) - 1, max_seqlen_q, q.shape[1], q.shape[2] # pylint: disable=self-assigning-variable
+        batch_k, max_seqlen_k, nheads_k, head_size_k = len(cu_seqlens_k) - 1, max_seqlen_k, k.shape[1], k.shape[2] # pylint: disable=self-assigning-variable
+    else:
+        assert False, "Got unsupported layout."
+
+    # assert
+    assert batch_q == batch_k
+    assert head_size_q == head_size_k
+
+    return batch_q, nheads_q, nheads_k, head_size_q, max_seqlen_q, max_seqlen_k
+
+
+def get_strides_from_layout(q, k, v, o, layout):
+    if layout == 'thd':
+        q_strides = (0, q.stride(1), q.stride(0), q.stride(2))
+        k_strides = (0, k.stride(1), k.stride(0), k.stride(2))
+        v_strides = (0, v.stride(1), v.stride(0), v.stride(2))
+        o_strides = (0, o.stride(1), o.stride(0), o.stride(2))
+    elif layout == 'bhsd':
+        q_strides = (q.stride(0), q.stride(1), q.stride(2), q.stride(3))
+        k_strides = (k.stride(0), k.stride(1), k.stride(2), k.stride(3))
+        v_strides = (v.stride(0), v.stride(1), v.stride(2), v.stride(3))
+        o_strides = (o.stride(0), o.stride(1), o.stride(2), o.stride(3))
+    elif layout == 'bshd':
+        q_strides = (q.stride(0), q.stride(2), q.stride(1), q.stride(3))
+        k_strides = (k.stride(0), k.stride(2), k.stride(1), k.stride(3))
+        v_strides = (v.stride(0), v.stride(2), v.stride(1), v.stride(3))
+        o_strides = (o.stride(0), o.stride(2), o.stride(1), o.stride(3))
+    else:
+        assert False, 'Got unsupported layout.'
+    return q_strides, k_strides, v_strides, o_strides
+
+
+def get_padded_headsize(size):
+    # Get closest power of 2 over or equal to 32.
+    padded_d_model = 1 << (size - 1).bit_length()
+    # Smallest head_dim supported is 16. If smaller, the tile in the
+    # kernel is padded - there is no padding in memory for any dims.
+    padded_d_model = max(padded_d_model, 16)
+    return padded_d_model
+
+
+def _strides(x: torch.Tensor, *stride_names: str):
+    if x is None:
+        return {f"stride_{s}": 0 for i, s in enumerate(stride_names)}
+
+    assert x.ndim == len(stride_names)
+    return {f"stride_{s}": x.stride(i) for i, s in enumerate(stride_names)}
+
+
+def get_input_shapes():
+    cases = [(max(1, 2**(16 - i)), 1, 2**i, 16, 1, 128)
+             for i in range(8, 18)] + [(max(1, 2**(16 - i)), 1, 2**i, 16, 2, 128) for i in range(8, 18)]
+    return cases
+
+
+def is_hip():
+    return triton.runtime.driver.active.get_current_target().backend == "hip"
+
+
+def is_cdna():
+    return is_hip() and triton.runtime.driver.active.get_current_target().arch in ('gfx940', 'gfx941', 'gfx942',
+                                                                                   'gfx90a', 'gfx908')
+
+
+def is_rdna():
+    return is_hip() and triton.runtime.driver.active.get_current_target().arch in ("gfx1030", "gfx1100", "gfx1101",
+                                                                                   "gfx1102", "gfx1200", "gfx1201")
diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py
index 1b3f750a9..4f1224923 100644
--- a/modules/zluda_hijacks.py
+++ b/modules/zluda_hijacks.py
@@ -1,13 +1,21 @@
+from functools import wraps
 import torch
 import torch._dynamo.device_interface
-from modules import rocm, zluda
+from modules import rocm, zluda, shared
 
 
-_topk = torch.topk
-def topk(input: torch.Tensor, *args, **kwargs): # pylint: disable=redefined-builtin
-    device = input.device
-    values, indices = _topk(input.cpu(), *args, **kwargs)
-    return torch.return_types.topk((values.to(device), indices.to(device),))
+MEM_BUS_WIDTH = {
+    "AMD Radeon RX 9070 XT": 256,
+    "AMD Radeon RX 9070": 256,
+    "AMD Radeon RX 9060 XT": 192,
+    "AMD Radeon RX 7900 XTX": 384,
+    "AMD Radeon RX 7900 XT": 320,
+    "AMD Radeon RX 7900 GRE": 256,
+    "AMD Radeon RX 7800 XT": 256,
+    "AMD Radeon RX 7700 XT": 192,
+    "AMD Radeon RX 7600 XT": 128,
+    "AMD Radeon RX 7600": 128,
+}
 
 
 class DeviceProperties:
@@ -35,20 +43,60 @@ def torch__C__cuda_getCurrentRawStream(device):
 
 def do_hijack():
     torch.version.hip = rocm.version
-    torch.topk = topk
 
     if zluda.default_agent is not None:
         DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = zluda.default_agent.name
     torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access
     torch._C._cuda_getCurrentRawStream = torch__C__cuda_getCurrentRawStream # pylint: disable=protected-access
     torch._dynamo.device_interface.CudaInterface.get_raw_stream = staticmethod(torch__C__cuda_getCurrentRawStream) # pylint: disable=protected-access
+
+    # Triton
     try:
         import triton
         _get_device_properties = triton.runtime.driver.active.utils.get_device_properties
         def triton_runtime_driver_active_utils_get_device_properties(device):
             props = _get_device_properties(device)
-            props["mem_bus_width"] = 384
+            name = torch.cuda.get_device_name()[:-8]
+            if name in MEM_BUS_WIDTH:
+                props["mem_bus_width"] = MEM_BUS_WIDTH[name]
+            else:
+                props["mem_bus_width"] = 128
+                shared.log.warning(f'[TRITON] defaulting mem_bus_width=128 for device "{name}".')
             return props
         triton.runtime.driver.active.utils.get_device_properties = triton_runtime_driver_active_utils_get_device_properties
+
+        if 'Flash attention' in shared.opts.sdp_options:
+            from modules.flash_attn_triton_amd import interface_fa
+            sdpa_pre_flash_atten = torch.nn.functional.scaled_dot_product_attention
+            @wraps(sdpa_pre_flash_atten)
+            def sdpa_flash_atten(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None):
+                if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32:
+                    if scale is None:
+                        scale = query.shape[-1] ** (-0.5)
+                    head_size_og = query.size(3)
+                    if head_size_og % 8 != 0:
+                        query = torch.nn.functional.pad(query, [0, 8 - head_size_og % 8])
+                        key = torch.nn.functional.pad(key, [0, 8 - head_size_og % 8])
+                        value = torch.nn.functional.pad(value, [0, 8 - head_size_og % 8])
+                    out_padded, _, _, _ = interface_fa.fwd(
+                        query.transpose(1, 2),
+                        key.transpose(1, 2),
+                        value.transpose(1, 2),
+                        None,
+                        None,
+                        dropout_p,
+                        scale,
+                        is_causal,
+                        -1,
+                        -1,
+                        0.0,
+                        False,
+                        None,
+                    )
+                    return out_padded[..., :head_size_og].transpose(1, 2)
+                else:
+                    return sdpa_pre_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale)
+            torch.nn.functional.scaled_dot_product_attention = sdpa_flash_atten
+            shared.log.debug('Torch attention: type="triton flash attention"')
     except Exception:
         pass

From 2c77a1619508a5a060358a018ad1dbc969a19092 Mon Sep 17 00:00:00 2001
From: Vladimir Mandic 
Date: Sat, 22 Mar 2025 10:34:02 -0400
Subject: [PATCH 048/122] video updates

Signed-off-by: Vladimir Mandic 
---
 TODO.md                                 | 19 +++++++++----------
 modules/video_models/video_load.py      |  6 +++++-
 modules/video_models/video_overrides.py | 25 +++++++++++++++++++++++++
 modules/video_models/video_run.py       |  5 ++++-
 modules/video_models/video_vae.py       |  1 +
 wiki                                    |  2 +-
 6 files changed, 45 insertions(+), 13 deletions(-)
 create mode 100644 modules/video_models/video_overrides.py

diff --git a/TODO.md b/TODO.md
index 95426f72b..267b1822b 100644
--- a/TODO.md
+++ b/TODO.md
@@ -10,21 +10,20 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
 - VAE Remote encode: SD15 and Flux.1 issues:   
 - Video: API support is TBD  
 - Video: Hunyuan Video I2V: transformers incompatibility   
-- Video: Hunyuan Video I2V: 16ch vs 33ch processing 
-- Video: WAN 2.1 14B I2V 480p: broken offload
-- Video: WAN 2.1 14B I2V 720p: broken offload
-- Video: CogVideoX 1.5 5B T2V/I2V: requires pipeline update
-- Video: CogVideoX 1.5 5B I2V: requires pipeline update
-- Video: LTXVideo 0.9.5 T2V/I2V: broken offload, new pipeline
-- Video: LTXVideo 0.9.5 T2V/I2V: set preset params
-- Video: LTXVideo 0.9.5 T2V/I2V: support for conditioned input  
-- Video: LTXVideo 0.9.1 I2V: generator list mismatch
+- Video: Hunyuan Video I2V: add 16ch vs 33ch processing   
+- Video: WAN 2.1 14B I2V 480p/720p: broken offload  
+- Video: CogVideoX 1.5 5B T2V/I2V: requires pipeline update  
+- Video: LTXVideo 0.9.5 T2V/I2V: broken offload  
+- Video: LTXVideo 0.9.5 T2V/I2V: requires different params  
+- Video: LTXVideo 0.9.5 T2V/I2V: add support for conditioned input  
+- Video: LTXVideo 0.9.1 I2V: generator list mismatch  
+- Video: Latte 1 T2V: dtype mismatch   
+- Video: Allegro T2V: all-gray output, requires vae-fp32  
 - Video: FasterCache and PyramidAttentionBroadcast granular config  
 - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN   
 
 ## Future Candidates
 
-- Redesign postprocessing  
 - Flux NF4 loader:   
 - IPAdapter negative:   
 - Control API enhance scripts compatibility  
diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py
index f380a8d88..9254c4ed6 100644
--- a/modules/video_models/video_load.py
+++ b/modules/video_models/video_load.py
@@ -1,7 +1,7 @@
 import os
 import time
 from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices
-from modules.video_models import models_def, video_utils, video_vae
+from modules.video_models import models_def, video_utils, video_vae, video_overrides
 
 
 loaded_model = None
@@ -49,6 +49,9 @@ def load_model(selected: models_def.Model):
         errors.display(e, 'video')
         transformer = None
 
+    # overrides
+    kwargs = video_overrides.load_override(selected)
+
     # model
     try:
         debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}')
@@ -58,6 +61,7 @@ def load_model(selected: models_def.Model):
             text_encoder=text_encoder,
             cache_dir=shared.opts.hfcache_dir,
             torch_dtype=devices.dtype,
+            **kwargs,
         )
     except Exception as e:
         shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}')
diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py
new file mode 100644
index 000000000..2dbb2409e
--- /dev/null
+++ b/modules/video_models/video_overrides.py
@@ -0,0 +1,25 @@
+import os
+import torch
+import diffusers
+from modules import shared, processing
+from modules.video_models.models_def import Model
+
+
+debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
+
+
+def load_override(selected: Model):
+    kwargs = {}
+    if selected.name == 'Allegro T2V':
+        kwargs['vae'] = diffusers.AutoencoderKLAllegro.from_pretrained(selected.repo,
+                                                                       subfolder="vae",
+                                                                       torch_dtype=torch.float32,
+                                                                       cache_dir=shared.opts.hfcache_dir)
+        debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}')
+    return kwargs
+
+
+def set_overrides(p: processing.StableDiffusionProcessingVideo, selected: Model):
+    if selected.name == 'Latte 1 T2V':
+        p.task_args['enable_temporal_attentions'] = False
+        debug(f'Video overrides: model="{selected.name}" args={p.task_args}')
diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py
index 2172fe053..57c2bee53 100644
--- a/modules/video_models/video_run.py
+++ b/modules/video_models/video_run.py
@@ -1,7 +1,7 @@
 import os
 import time
 from modules import shared, errors, sd_models, processing, devices, images, ui_common
-from modules.video_models import models_def, video_utils, video_load, video_vae, video_cache
+from modules.video_models import models_def, video_utils, video_load, video_vae, video_cache, video_overrides
 
 
 debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -65,12 +65,15 @@ def generate(*args, **kwargs):
     video_vae.set_vae_params(p)
     video_cache.set_cache(faster_cache=faster_cache, pyramid_attention_broadcast=pyramid_attention)
     video_utils.set_prompt(p)
+    p.task_args['width'] = p.width
+    p.task_args['height'] = p.height
     p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil'
     p.ops.append('video')
     orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
     orig_sampler_shift = shared.opts.schedulers_shift
     shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
     shared.opts.data['schedulers_shift'] = sampler_shift
+    video_overrides.set_overrides(p, selected)
     debug(f'Video: task_args={p.task_args}')
 
     # run processing
diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py
index bd552f66c..b609265cc 100644
--- a/modules/video_models/video_vae.py
+++ b/modules/video_models/video_vae.py
@@ -56,6 +56,7 @@ def hijack_vae_decode(*args, **kwargs):
     if res is None:
         shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae'])
         res = shared.sd_model.vae.orig_decode(*args, **kwargs)
+        print('HERE', shared.sd_model.vae.dtype)
     t1 = time.time()
     timer.process.add('vae', t1-t0)
     debug(f'Video decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}')
diff --git a/wiki b/wiki
index 00145b30c..f58dcaaf6 160000
--- a/wiki
+++ b/wiki
@@ -1 +1 @@
-Subproject commit 00145b30c5ed318423487f8aa6336d834b578db7
+Subproject commit f58dcaaf6a98aca2a6adf0a509ebb70a9bd3516a

From d7bab01df0069445529ab55fce38c3a641e5cc42 Mon Sep 17 00:00:00 2001
From: Vladimir Mandic 
Date: Sun, 23 Mar 2025 13:19:12 -0400
Subject: [PATCH 049/122] update video

Signed-off-by: Vladimir Mandic 
---
 CHANGELOG.md                            |  1 +
 TODO.md                                 | 19 ++++++-------
 modules/processing_args.py              |  2 ++
 modules/sd_models.py                    |  7 +++--
 modules/ui_sections.py                  |  2 +-
 modules/ui_video.py                     |  3 +-
 modules/video_models/models_def.py      | 38 ++-----------------------
 modules/video_models/video_overrides.py | 21 +++++++++++++-
 modules/video_models/video_run.py       |  3 +-
 modules/video_models/video_utils.py     |  9 ++++--
 modules/video_models/video_vae.py       | 16 +++++++++--
 wiki                                    |  2 +-
 12 files changed, 66 insertions(+), 57 deletions(-)

diff --git a/CHANGELOG.md b/CHANGELOG.md
index c839cf403..3119b0ef4 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -12,6 +12,7 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio
 ### Details for 2025-03-22
 
 - **Video tab**
+  - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details!  
   - new top-level tab, replaces previous *video* script in text/image tabs  
     old scripts are still present, but will be removed in the future  
   - support for all latest models:  
diff --git a/TODO.md b/TODO.md
index 267b1822b..02768136d 100644
--- a/TODO.md
+++ b/TODO.md
@@ -11,22 +11,21 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
 - Video: API support is TBD  
 - Video: Hunyuan Video I2V: transformers incompatibility   
 - Video: Hunyuan Video I2V: add 16ch vs 33ch processing   
-- Video: WAN 2.1 14B I2V 480p/720p: broken offload  
-- Video: CogVideoX 1.5 5B T2V/I2V: requires pipeline update  
-- Video: LTXVideo 0.9.5 T2V/I2V: broken offload  
-- Video: LTXVideo 0.9.5 T2V/I2V: requires different params  
-- Video: LTXVideo 0.9.5 T2V/I2V: add support for conditioned input  
-- Video: LTXVideo 0.9.1 I2V: generator list mismatch  
 - Video: Latte 1 T2V: dtype mismatch   
-- Video: Allegro T2V: all-gray output, requires vae-fp32  
+- Video: WAN 2.1 14B I2V 480p/720p: broken offload  
+- Video: CogVideoX 1.5 5B T2V/I2V: all-gray output  
+- Video: Allegro T2V: all-gray output
 - Video: FasterCache and PyramidAttentionBroadcast granular config  
 - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN   
 
 ## Future Candidates
 
-- Flux NF4 loader:   
-- IPAdapter negative:   
-- Control API enhance scripts compatibility  
+- Flux: NF4 loader:   
+- IPAdapter: negative guidance:   
+- Control: API enhance scripts compatibility  
+- Video: OponSora v2 https://huggingface.co/hpcai-tech/Open-Sora-v2
+- Video: STG: https://github.com/huggingface/diffusers/blob/main/examples/community/README.md#spatiotemporal-skip-guidance
+- Video: add generate context menu
 
 ## Code TODO
 
diff --git a/modules/processing_args.py b/modules/processing_args.py
index db8794214..058c8bdf1 100644
--- a/modules/processing_args.py
+++ b/modules/processing_args.py
@@ -21,6 +21,7 @@ disable_pbar = os.environ.get('SD_DISABLE_PBAR', None) is not None
 def task_specific_kwargs(p, model):
     task_args = {}
     is_img2img_model = bool('Zero123' in shared.sd_model.__class__.__name__)
+    print('HERE', sd_models.get_diffusers_task(model))
     if len(getattr(p, 'init_images', [])) > 0:
         if isinstance(p.init_images[0], str):
             p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images]
@@ -34,6 +35,7 @@ def task_specific_kwargs(p, model):
                 'height': 8 * math.ceil(p.height / 8),
             }
     elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.IMAGE_2_IMAGE or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0:
+        print('HERE1', p.denoising_strength)
         if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'):
             model.register_to_config(requires_aesthetics_score = False)
         if 'hires' not in p.ops:
diff --git a/modules/sd_models.py b/modules/sd_models.py
index 8c4356793..9f66d9814 100644
--- a/modules/sd_models.py
+++ b/modules/sd_models.py
@@ -631,9 +631,12 @@ class DiffusersTaskType(Enum):
 
 
 def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType:
-    if pipe.__class__.__name__ in ["StableVideoDiffusionPipeline", "LEditsPPPipelineStableDiffusion", "LEditsPPPipelineStableDiffusionXL", "OmniGenPipeline"]:
+    cls = pipe.__class__.__name__
+    if cls in ["LEditsPPPipelineStableDiffusion", "LEditsPPPipelineStableDiffusionXL", "OmniGenPipeline"]: # special case
         return DiffusersTaskType.IMAGE_2_IMAGE
-    elif pipe.__class__.__name__ == "StableDiffusionXLInstructPix2PixPipeline":
+    elif 'ImageToVideo' in cls or cls in ['LTXConditionPipeline', 'StableVideoDiffusionPipeline']: # i2v pipelines
+        return DiffusersTaskType.IMAGE_2_IMAGE
+    elif 'Instruct' in cls:
         return DiffusersTaskType.INSTRUCT
     elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING.values():
         return DiffusersTaskType.IMAGE_2_IMAGE
diff --git a/modules/ui_sections.py b/modules/ui_sections.py
index d2e73bdcc..8d038ab13 100644
--- a/modules/ui_sections.py
+++ b/modules/ui_sections.py
@@ -219,7 +219,7 @@ def create_sampler_and_steps_selection(choices, tabname):
         sd_samplers.set_samplers()
         choices = [x for x in sd_samplers.samplers if not x.name == 'Same as primary']
     with gr.Row(elem_classes=['flex-break']):
-        steps = gr.Slider(minimum=1, maximum=99, step=1, label="Steps", elem_id=f"{tabname}_steps", value=20)
+        steps = gr.Slider(minimum=1, maximum=100, step=1, label="Steps", elem_id=f"{tabname}_steps", value=20)
         sampler_index = gr.Dropdown(label='Sampling method', elem_id=f"{tabname}_sampling", choices=[x.name for x in choices], value='Default', type="index")
     return steps, sampler_index
 
diff --git a/modules/ui_video.py b/modules/ui_video.py
index cff8a54b0..c0495defe 100644
--- a/modules/ui_video.py
+++ b/modules/ui_video.py
@@ -110,6 +110,7 @@ def create_ui():
                 with gr.Accordion(open=False, label="Init image", elem_id='video_init_accordion'):
                     gr.HTML("
  Init image") init_image = gr.Image(elem_id="video_image", show_label=False, type="pil", image_mode="RGB", height=512) + init_strength = gr.Slider(label='Init strength', minimum=0.0, maximum=1.0, step=0.01, value=0.5, elem_id="video_denoising_strength") with gr.Accordion(open=False, label="Accelerate", elem_id='video_accelerate_accordion'): faster_cache = gr.Checkbox(label='FasterCache', value=False, elem_id="video_faster_cache") pyramid_attention = gr.Checkbox(label='PyramidAttention', value=False, elem_id="video_pyramid_attention") @@ -163,7 +164,7 @@ def create_ui(): sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, - init_image, + init_image, init_strength, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index f5a7804d1..f7181b150 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -3,38 +3,6 @@ import diffusers import transformers -""" -# Model tests: download/load/generate - -- Hunyuan Video T2V: pass/pass/pass -- Hunyuan Video I2V: pass/pass/fail, transformers incompatibility -- SkyReels Hunyuan T2V: pass/pass/pass -- SkyReels Hunyuan I2V: pass/pass/pass -- Fast Hunyuan T2V: pass/pass/pass - -- LTXVideo 0.9.5 T2V: pass/pass/fail, completely broken offload, new pipeline -- LTXVideo 0.9.5 I2V: pass/pass/fail, completely broken offload, new pipeline -- LTXVideo 0.9.1 T2V: pass/pass/pass -- LTXVideo 0.9.1 I2V: pass/pass/fail, generator list mismatch -- LTXVideo 0.9.0 T2V: pass/pass/pass -- LTXVideo 0.9.0 I2V: pass/pass/pass - -- WAN 2.1 1.3B T2V: pass/pass/pass -- WAN 2.1 14B T2V: pass/pass/pass -- WAN 2.1 14B I2V 480p: pass/pass/fail, offloading cpu vs cuda -- WAN 2.1 14B I2V 720p: pass/pass/fail, offloading cpu vs cuda - -- CogVideoX 1.0 2B T2V: pass/pass/pass -- CogVideoX 1.0 5B T2V: pass/pass/pass -- CogVideoX 1.0 5B I2V: pass/pass/pass -- CogVideoX 1.5 5B T2V: download/load/fail, pipeline is tbd -- CogVideoX 1.5 5B I2V: download/load/fail, pipeline is tbd - -- Mochi 1 T2V: pass/pass/pass -- Latte 1 T2V: pass/pass/fail, float vs bfloat during generate -- Allegro T2V: pass/pass/fail, output is pure gray -""" - @dataclass class Model(): name: str @@ -104,13 +72,13 @@ models = { Model(name='None'), Model(name='LTXVideo 0.9.5 T2V', # https://github.com/huggingface/diffusers/pull/10968 url='https://huggingface.co/Lightricks/LTX-Video-0.9.5', - repo='YiYiXu/ltx-95', + repo='Lightricks/LTX-Video-0.9.5', repo_cls=diffusers.LTXPipeline, te_cls=transformers.T5EncoderModel, dit_cls=diffusers.LTXVideoTransformer3DModel), Model(name='LTXVideo 0.9.5 I2V', url='https://huggingface.co/Lightricks/LTX-Video-0.9.5', - repo='YiYiXu/ltx-95', + repo='Lightricks/LTX-Video-0.9.5', repo_cls=diffusers.LTXConditionPipeline, te_cls=transformers.T5EncoderModel, dit_cls=diffusers.LTXVideoTransformer3DModel), @@ -214,7 +182,7 @@ models = { te_cls=transformers.T5EncoderModel, dit_cls=diffusers.CogVideoXTransformer3DModel), Model(name='CogVideoX 1.5 5B T2V', - url='https://huggingface.co/THUDM/THUDM/CogVideoX1.5-5B', + url='https://huggingface.co/THUDM/CogVideoX1.5-5B', repo='THUDM/CogVideoX1.5-5B', repo_cls=diffusers.CogVideoXPipeline, te_cls=transformers.T5EncoderModel, diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index 2dbb2409e..e771f5254 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -16,10 +16,29 @@ def load_override(selected: Model): torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir) debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}') + if selected.name == 'LTXVideo 0.9.5 I2V': + kwargs['vae'] = diffusers.AutoencoderKLLTXVideo.from_pretrained(selected.repo, + subfolder="vae", + torch_dtype=torch.float32, + cache_dir=shared.opts.hfcache_dir) return kwargs def set_overrides(p: processing.StableDiffusionProcessingVideo, selected: Model): + cls = shared.sd_model.__class__.__name__ + # Allegro + if selected.name == 'Allegro T2V': + shared.sd_model.vae.enable_tiling() + # Latte if selected.name == 'Latte 1 T2V': p.task_args['enable_temporal_attentions'] = False - debug(f'Video overrides: model="{selected.name}" args={p.task_args}') + p.task_args['video_length'] = p.frames + # LTX + if cls == 'LTXImageToVideoPipeline' or cls == 'LTXConditionPipeline': + p.task_args['generator'] = None + if cls == 'LTXConditionPipeline': + print('HERE2', p.denoising_strength) + p.task_args['strength'] = p.denoising_strength + if 'LTX' in shared.sd_model.__class__.__name__: + p.task_args['width'] = 32 * (p.width // 32) + p.task_args['height'] = 32 * (p.height // 32) diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 57c2bee53..95d75e149 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -8,7 +8,7 @@ debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None e def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, faster_cache, pyramid_attention, override_settings = args + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, init_strength, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, faster_cache, pyramid_attention, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') found = [model.name for model in models_def.models.get(engine, [])] @@ -36,6 +36,7 @@ def generate(*args, **kwargs): width=16 * int(width // 16), height=16 * int(height // 16), frames=int(frames), + denoising_strength=float(init_strength), init_image=init_image, cfg_scale=float(guidance_scale), diffusers_guidance_rescale=float(guidance_true), diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 8a69e19ec..bb971d9d5 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -1,6 +1,6 @@ import os import time -from modules import shared, sd_models, timer +from modules import shared, sd_models, timer, errors debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -30,7 +30,12 @@ def set_prompt(p): def hijack_encode_prompt(*args, **kwargs): t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + try: + res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + except Exception as e: + shared.log.error(f'Video encode: {e}') + errors.display(e, 'Video encode') + res = None t1 = time.time() timer.process.add('te', t1-t0) debug(f'Video encode: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py index b609265cc..8adbc939e 100644 --- a/modules/video_models/video_vae.py +++ b/modules/video_models/video_vae.py @@ -1,6 +1,7 @@ import os import time -from modules import shared, sd_models, devices, timer +import torch +from modules import shared, sd_models, devices, timer, errors debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -55,8 +56,17 @@ def hijack_vae_decode(*args, **kwargs): pass if res is None: shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - print('HERE', shared.sd_model.vae.dtype) + try: + if torch.is_tensor(args[0]): + latent = args[0] + latent = latent.to(device=devices.device, dtype=shared.sd_model.vae.dtype) # upcast to vae dtype + res = shared.sd_model.vae.orig_decode(latent, *args[1:], **kwargs) + else: + res = shared.sd_model.vae.orig_decode(*args, **kwargs) + except Exception as e: + shared.log.error(f'Video VAE: type={vae_type} {e}') + errors.display(e, 'Video VAE') + res = None t1 = time.time() timer.process.add('vae', t1-t0) debug(f'Video decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}') diff --git a/wiki b/wiki index f58dcaaf6..ef3c65cb0 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit f58dcaaf6a98aca2a6adf0a509ebb70a9bd3516a +Subproject commit ef3c65cb0eb3d023daae68d0817c1828cf16ea39 From d4c8de99f59c90f2f1a7c13b3faaa815c9a276ca Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 23 Mar 2025 14:26:57 -0400 Subject: [PATCH 050/122] cleanup Signed-off-by: Vladimir Mandic --- modules/processing_args.py | 2 -- modules/sd_samplers.py | 5 +++-- modules/video_models/video_overrides.py | 1 - 3 files changed, 3 insertions(+), 5 deletions(-) diff --git a/modules/processing_args.py b/modules/processing_args.py index 058c8bdf1..db8794214 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -21,7 +21,6 @@ disable_pbar = os.environ.get('SD_DISABLE_PBAR', None) is not None def task_specific_kwargs(p, model): task_args = {} is_img2img_model = bool('Zero123' in shared.sd_model.__class__.__name__) - print('HERE', sd_models.get_diffusers_task(model)) if len(getattr(p, 'init_images', [])) > 0: if isinstance(p.init_images[0], str): p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images] @@ -35,7 +34,6 @@ def task_specific_kwargs(p, model): 'height': 8 * math.ceil(p.height / 8), } elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.IMAGE_2_IMAGE or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0: - print('HERE1', p.denoising_strength) if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'): model.register_to_config(requires_aesthetics_score = False) if 'hires' not in p.ops: diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 67ed8e8ee..be375dbff 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -12,6 +12,8 @@ samplers = all_samplers samplers_for_img2img = all_samplers samplers_map = {} loaded_config = None +flow_models = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow', 'Sana', 'CogView4'] +flow_models += ['Hunyuan', 'LTX', 'Mochi'] def list_samplers(): @@ -79,10 +81,9 @@ def create_sampler(name, model): shared.log.debug(f'Sampler: "{name}" config={config.options}') return sampler elif shared.native: - FlowModels = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow', 'Sana', 'HunyuanVideoPipeline', 'CogView4Pipeline'] if 'KDiffusion' in model.__class__.__name__: return None - if not any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' in name: + if not any(x in model.__class__.__name__ for x in flow_models) and 'FlowMatch' in name: shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} flow-match scheduler unsupported') return None sampler = config.constructor(model) diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index e771f5254..7af00f6be 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -37,7 +37,6 @@ def set_overrides(p: processing.StableDiffusionProcessingVideo, selected: Model) if cls == 'LTXImageToVideoPipeline' or cls == 'LTXConditionPipeline': p.task_args['generator'] = None if cls == 'LTXConditionPipeline': - print('HERE2', p.denoising_strength) p.task_args['strength'] = p.denoising_strength if 'LTX' in shared.sd_model.__class__.__name__: p.task_args['width'] = 32 * (p.width // 32) From ec488d2a369b0c94d3c3637378bbdf5a3e1a6a81 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 23 Mar 2025 14:34:55 -0400 Subject: [PATCH 051/122] cleanup Signed-off-by: Vladimir Mandic --- modules/sd_samplers_diffusers.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/modules/sd_samplers_diffusers.py b/modules/sd_samplers_diffusers.py index 667683eb0..523354e66 100644 --- a/modules/sd_samplers_diffusers.py +++ b/modules/sd_samplers_diffusers.py @@ -50,6 +50,10 @@ try: from modules.schedulers.scheduler_bdia import BDIA_DDIMScheduler # pylint: disable=ungrouped-imports from modules.schedulers.scheduler_ufogen import UFOGenScheduler # pylint: disable=ungrouped-imports from modules.perflow import PeRFlowScheduler # pylint: disable=ungrouped-imports + # from modules.schedulers.scheduler_kohaku import KohakuLoNyuYogScheduler # pylint: disable=ungrouped-imports + # from modules.schedulers.scheduler_smea import SMEAScheduler # pylint: disable=ungrouped-imports + # from modules.schedulers.scheduler_dy import DYScheduler # pylint: disable=ungrouped-imports + # from modules.schedulers.scheduler_negative import EulerNegativeScheduler # pylint: disable=ungrouped-imports except Exception as e: shared.log.error(f'Sampler import: version={diffusers.__version__} error: {e}') if os.environ.get('SD_SAMPLER_DEBUG', None) is not None: @@ -68,6 +72,10 @@ config = { 'Euler SGM': { 'steps_offset': 0, 'interpolation_type': "linear", 'rescale_betas_zero_snr': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'trailing', 'use_beta_sigmas': False, 'use_exponential_sigmas': False, 'use_karras_sigmas': False, 'prediction_type': "sample" }, 'Euler EDM': { 'sigma_schedule': "karras" }, 'Euler FlowMatch': { 'timestep_spacing': "linspace", 'shift': 1, 'use_dynamic_shifting': False, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False }, + # 'Euler SMEA': {}, + # 'Euler DY': {}, + # 'Euler Negative': {}, + # 'Kohaku LoNyu': {}, 'DPM++': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_flow_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 1 }, 'DPM++ 2M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_flow_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 }, @@ -124,6 +132,9 @@ samplers_data_diffusers = [ SamplerData('Euler SGM', lambda model: DiffusionSampler('Euler SGM', EulerDiscreteScheduler, model), [], {}), SamplerData('Euler EDM', lambda model: DiffusionSampler('Euler EDM', EDMEulerScheduler, model), [], {}), SamplerData('Euler FlowMatch', lambda model: DiffusionSampler('Euler FlowMatch', FlowMatchEulerDiscreteScheduler, model), [], {}), + # SamplerData('Euler SMEA', lambda model: DiffusionSampler('Euler SMEA', SMEAScheduler, model), [], {}), + # SamplerData('Euler DY', lambda model: DiffusionSampler('Euler DY', DYScheduler, model), [], {}), + # SamplerData('Euler Negative', lambda model: DiffusionSampler('Euler Negative', EulerNegativeScheduler, model), [], {}), SamplerData('DPM++', lambda model: DiffusionSampler('DPM++', DPMSolverMultistepScheduler, model), [], {}), SamplerData('DPM++ 2M', lambda model: DiffusionSampler('DPM++ 2M', DPMSolverMultistepScheduler, model), [], {}), @@ -169,6 +180,7 @@ samplers_data_diffusers = [ SamplerData('TDD', lambda model: DiffusionSampler('TDD', TDDScheduler, model), [], {}), SamplerData('PeRFlow', lambda model: DiffusionSampler('PeRFlow', PeRFlowScheduler, model), [], {}), SamplerData('UFOGen', lambda model: DiffusionSampler('UFOGen', UFOGenScheduler, model), [], {}), + # SamplerData('Kohaku LoNyu', lambda model: DiffusionSampler('Kohaku LoNyu', KohakuLoNyuYogScheduler, model), [], {}), SamplerData('Same as primary', None, [], {}), ] From 6b1d83d2e69f3c049f4cabcd9f9ea0ecff8eddf0 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Mon, 24 Mar 2025 21:20:51 +0900 Subject: [PATCH 052/122] zluda catch attributeerror --- modules/zluda_installer.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index ae801b4b8..4b915c724 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -47,8 +47,11 @@ class Core(ZLUDALibrary): internal.zluda_get_hip_object.restype = ZLUDAResult internal.zluda_get_hip_object.argtypes = [ctypes.c_void_p, ctypes.c_int] - internal.zluda_get_nightly_flag.restype = ctypes.c_int - internal.zluda_get_nightly_flag.argtypes = [] + try: + internal.zluda_get_nightly_flag.restype = ctypes.c_int + internal.zluda_get_nightly_flag.argtypes = [] + except AttributeError: + internal.zluda_get_nightly_flag = lambda: 0 super().__init__(internal) From 93b731e0cf31172688a5710320b95949cdc6595f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 24 Mar 2025 17:27:46 -0400 Subject: [PATCH 053/122] add sana 1.5 support Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 23 ++++++++++++++++------- TODO.md | 8 +++++--- html/reference.json | 26 ++++++++++++++++++++++---- installer.py | 2 +- modules/model_sana.py | 19 ++++++++++++------- wiki | 2 +- 6 files changed, 57 insertions(+), 23 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 3119b0ef4..525a34fac 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,15 +1,15 @@ # Change Log for SD.Next -## Update for 2025-03-22 +## Update for 2025-03-24 -### Highlights for 2025-03-22 +### Highlights for 2025-03-24 -Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** +Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! -Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides +Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to remote VAE, additional docs/guides -### Details for 2025-03-22 +### Details for 2025-03-24 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -46,20 +46,29 @@ Plus support for CogView-4, new CLiP models, improvements to remote VAE, additio and may require specific settings - see model links for details - see *ToDo/Limitations* section for additional notes - **Models** - - [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B) + - [THUDM CogView 4](https://huggingface.co/THUDM/CogView4-6B) **6B** variant new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer fully supports offloading and on-the-fly quantization simply select from *networks -> models -> reference* *note* cogview4 is compatible with flowmatching samplers + - [NVLabs SANA 1.5](https://huggingface.co/Efficient-Large-Model/SANA1.5_4.8B_1024px_diffusers) in **1.6B**, **4.8B** and [Sprint](https://huggingface.co/Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers) variations + big update to previous SANA model + fully supports offloading and on-the-fly quantization + simply select from *networks -> models -> reference* - New [zer0int CLiP-L](https://huggingface.co/zer0int/CLIP-Registers-Gated_MLP-ViT-L-14) models: download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui - **Wiki/Docs** + - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info + - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide - new [Caption](https://github.com/vladmandic/sdnext/wiki/Caption) guide - new [VAE](https://github.com/vladmandic/sdnext/wiki/VAE) guide - - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - updated [SD3](https://github.com/vladmandic/sdnext/wiki/SD3) guide + - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide + - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide + - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide + - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide - **Remote VAE** - add support for remote vae encode in addition to remote vae decode - used by *img2img, inpaint, hires, detailer* diff --git a/TODO.md b/TODO.md index 02768136d..de0c7c8be 100644 --- a/TODO.md +++ b/TODO.md @@ -15,17 +15,19 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: WAN 2.1 14B I2V 480p/720p: broken offload - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output - Video: Allegro T2V: all-gray output -- Video: FasterCache and PyramidAttentionBroadcast granular config -- Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN ## Future Candidates - Flux: NF4 loader: - IPAdapter: negative guidance: - Control: API enhance scripts compatibility +- Video: add generate context menu +- Video: FasterCache and PyramidAttentionBroadcast granular config +- Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN - Video: OponSora v2 https://huggingface.co/hpcai-tech/Open-Sora-v2 - Video: STG: https://github.com/huggingface/diffusers/blob/main/examples/community/README.md#spatiotemporal-skip-guidance -- Video: add generate context menu +- Video SmoothCache: https://github.com/huggingface/diffusers/issues/11135 +- FasterCache, PyramidAttentionBroadcast, SmoothCache general support ## Code TODO diff --git a/html/reference.json b/html/reference.json index 7b6589215..b18818750 100644 --- a/html/reference.json +++ b/html/reference.json @@ -187,25 +187,43 @@ "extras": "sampler: Default, cfg_scale: 3.5" }, - "NVLabs Sana 1.6B 4k": { + "NVLabs Sana 1.5 1.6B 1k": { + "path": "Efficient-Large-Model/SANA1.5_1.6B_1024px_diffusers", + "desc": "Sana is an efficient model with scaling of training-time and inference time techniques. SANA-1.5 delivers: efficient model growth from 1.6B Sana-1.0 model to 4.8B, achieving similar or better performance than training from scratch and saving 60% training cost; efficient model depth pruning, slimming any model size as you want; powerful VLM selection based inference scaling, smaller model+inference scaling > larger model.", + "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", + "skip": true + }, + "NVLabs Sana 1.5 4.8B 1k": { + "path": "Efficient-Large-Model/SANA1.5_4.8B_1024px_diffusers", + "desc": "Sana is an efficient model with scaling of training-time and inference time techniques. SANA-1.5 delivers: efficient model growth from 1.6B Sana-1.0 model to 4.8B, achieving similar or better performance than training from scratch and saving 60% training cost; efficient model depth pruning, slimming any model size as you want; powerful VLM selection based inference scaling, smaller model+inference scaling > larger model.", + "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", + "skip": true + }, + "NVLabs Sana 1.5 1.6B 1k Sprint": { + "path": "Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers", + "desc": "SANA-Sprint is an ultra-efficient diffusion model for text-to-image (T2I) generation, reducing inference steps from 20 to 1-4 while achieving state-of-the-art performance.", + "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", + "skip": true + }, + "NVLabs Sana 1.0 1.6B 4k": { "path": "Efficient-Large-Model/Sana_1600M_4Kpx_BF16_diffusers", "desc": "Sana is a text-to-image framework that can efficiently generate images up to 4096 × 4096 resolution. Sana can synthesize high-resolution, high-quality images with strong text-image alignment at a remarkably fast speed, deployable on laptop GPU.", "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", "skip": true }, - "NVLabs Sana 1.6B 2k": { + "NVLabs Sana 1.0 1.6B 2k": { "path": "Efficient-Large-Model/Sana_1600M_2Kpx_BF16_diffusers", "desc": "Sana is a text-to-image framework that can efficiently generate images up to 4096 × 4096 resolution. Sana can synthesize high-resolution, high-quality images with strong text-image alignment at a remarkably fast speed, deployable on laptop GPU.", "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", "skip": true }, - "NVLabs Sana 1.6B 1k": { + "NVLabs Sana 1.0 1.6B 1k": { "path": "Efficient-Large-Model/Sana_1600M_1024px_diffusers", "desc": "Sana is a text-to-image framework that can efficiently generate images up to 4096 × 4096 resolution. Sana can synthesize high-resolution, high-quality images with strong text-image alignment at a remarkably fast speed, deployable on laptop GPU.", "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", "skip": true }, - "NVLabs Sana 0.6B 0.5k": { + "NVLabs Sana 1.0 0.6B 0.5k": { "path": "Efficient-Large-Model/Sana_600M_512px_diffusers", "desc": "Sana is a text-to-image framework that can efficiently generate images up to 4096 × 4096 resolution. Sana can synthesize high-resolution, high-quality images with strong text-image alignment at a remarkably fast speed, deployable on laptop GPU.", "preview": "Efficient-Large-Model--Sana_1600M_1024px_diffusers.jpg", diff --git a/installer.py b/installer.py index 1ae99f737..12ec654b2 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git: return - sha = 'a7d53a59394d5d8367826663601b69828e9f74fc' # diffusers commit hash + sha = '5dbe4f5de6398159f8c2bedd371bc116683edbd3' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/model_sana.py b/modules/model_sana.py index 54f2681fa..509db4c88 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -33,8 +33,6 @@ def load_sana(checkpoint_info, kwargs={}): if not repo_id.endswith('_diffusers'): repo_id = f'{repo_id}_diffusers' - if devices.dtype == torch.bfloat16 and 'BF16' not in repo_id: - repo_id = repo_id.replace('_diffusers', '_BF16_diffusers') if 'Sana_1600M' in repo_id: if devices.dtype == torch.bfloat16 or 'BF16' in repo_id: @@ -47,13 +45,20 @@ def load_sana(checkpoint_info, kwargs={}): if 'Sana_600M' in repo_id: kwargs['variant'] = 'fp16' - if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)): - # TODO sana: fails when quantized - # kwargs = load_quants(kwargs, repo_id, cache_dir=shared.opts.diffusers_dir) - pass + kwargs = load_quants(kwargs, repo_id, cache_dir=shared.opts.diffusers_dir) shared.log.debug(f'Load model: type=Sana repo="{repo_id}" args={list(kwargs)}') t0 = time.time() - pipe = diffusers.SanaPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **kwargs) + if devices.dtype == torch.bfloat16 or devices.dtype == torch.float32: + kwargs['torch_dtype'] = devices.dtype + if 'Sprint' in repo_id: + cls = diffusers.SanaSprintPipeline + else: + cls = diffusers.SanaPipeline + pipe = cls.from_pretrained( + repo_id, + cache_dir=shared.opts.diffusers_dir, + **kwargs, + ) if devices.dtype == torch.bfloat16 or devices.dtype == torch.float32: if 'transformer' not in kwargs: pipe.transformer = pipe.transformer.to(dtype=devices.dtype) diff --git a/wiki b/wiki index ef3c65cb0..abc426fa7 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit ef3c65cb0eb3d023daae68d0817c1828cf16ea39 +Subproject commit abc426fa76531e1c626b2e7275bd8925e4beef9b From 6fe4b0604a5450c428555f353409b6698e324a7b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 24 Mar 2025 17:45:50 -0400 Subject: [PATCH 054/122] fix en cover and inline displays Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + javascript/extraNetworks.js | 8 ++++++++ 2 files changed, 9 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 525a34fac..09396f9e6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -118,6 +118,7 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - fix legacy diffusion latent upscalers - fix upscaler selection in postprocessing - fix sd35 with batch processing + - fix extra networks cover and inline views ## Update for 2025-02-28 diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index b9f045a6b..90701cdff 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -418,6 +418,9 @@ function setupExtraNetworksForTab(tabname) { if (h <= 0) return; const vh = opts.logmonitor_show ? '55vh' : '68vh'; if (window.opts.extra_networks_card_cover === 'sidebar' && window.opts.theme_type === 'Standard') el.style.height = `max(${vh}, ${h - 90}px)`; + else if (window.opts.extra_networks_card_cover === 'inline' && window.opts.theme_type === 'Standard') el.style.height = '25vh'; + else if (window.opts.extra_networks_card_cover === 'cover' && window.opts.theme_type === 'Standard') el.style.height = '50vh'; + else el.style.height = 'unset'; // log(`${tabname} height: ${entry.target.id}=${h} ${el.id}=${el.clientHeight}`); } } @@ -457,6 +460,8 @@ function setupExtraNetworksForTab(tabname) { en.style.height = 'unset'; en.style.width = 'unset'; en.style.right = 'unset'; + en.style.maxWidth = 'unset'; + en.style.maxHeight = '58vh'; en.style.top = '13em'; en.style.transition = ''; en.style.zIndex = 100; @@ -466,6 +471,7 @@ function setupExtraNetworksForTab(tabname) { en.style.height = 'auto'; en.style.width = `${window.opts.extra_networks_sidebar_width}vw`; en.style.maxWidth = '50vw'; + en.style.maxHeight = 'unset'; en.style.right = '0'; en.style.top = '13em'; en.style.transition = 'width 0.3s ease'; @@ -477,6 +483,8 @@ function setupExtraNetworksForTab(tabname) { en.style.height = 'unset'; en.style.width = 'unset'; en.style.right = 'unset'; + en.style.maxWidth = 'unset'; + en.style.maxHeight = '33vh'; en.style.top = 0; en.style.transition = ''; en.style.zIndex = 0; From 06b0abc4f496757728a314dd87e4c56ff2c18cdc Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 24 Mar 2025 18:08:33 -0400 Subject: [PATCH 055/122] fix token counter error style with modernui Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + extensions-builtin/sdnext-modernui | 2 +- javascript/base.css | 10 +++++----- javascript/promptChecker.js | 2 ++ javascript/sdnext.css | 10 +++++----- modules/sd_models.py | 3 ++- 6 files changed, 16 insertions(+), 12 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 09396f9e6..cfe49ffee 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -119,6 +119,7 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - fix upscaler selection in postprocessing - fix sd35 with batch processing - fix extra networks cover and inline views + - fix token counter error style with modernui ## Update for 2025-02-28 diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 4d0bde42e..baee2e439 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 4d0bde42e95c801dba3a48b9161e12b8ddc5d5bd +Subproject commit baee2e4397537554e6db67e7e2496258f024880f diff --git a/javascript/base.css b/javascript/base.css index 8f89685c2..fe06e6732 100644 --- a/javascript/base.css +++ b/javascript/base.css @@ -4,11 +4,11 @@ .gradio-button.tool { max-width: min-content; min-width: min-content !important; align-self: end; font-size: 1.4em; color: var(--body-text-color) !important; } /* token counters */ -.block.token-counter { position: absolute; display: inline-block; right: 0; min-width: 0 !important; width: auto; z-index: 100; top: 0; } -.block.token-counter span { background: var(--input-background-fill) !important; box-shadow: 0 0 0.0 0.3em rgba(192,192,192,0.15), inset 0 0 0.6em rgba(192,192,192,0.075); border: 2px solid rgba(192,192,192,0.4) !important; } -.block.token-counter.error span{ box-shadow: 0 0 0.0 0.3em rgba(255,0,0,0.15), inset 0 0 0.6em rgba(255,0,0,0.075); border: 2px solid rgba(255,0,0,0.4) !important; } -.block.token-counter div { display: inline; } -.block.token-counter span { padding: 0.1em 0.75em; } +.token-counter { position: absolute; display: inline-block; right: 0; min-width: 0 !important; width: auto; z-index: 100; top: 0; } +.token-counter span { background: var(--input-background-fill) !important; box-shadow: 0 0 0.0 0.3em rgba(192,192,192,0.15), inset 0 0 0.6em rgba(192,192,192,0.075); border: 2px solid rgba(192,192,192,0.4) !important; } +.token-counter.error span{ box-shadow: 0 0 0.0 0.3em rgba(255,0,0,0.15), inset 0 0 0.6em rgba(255,0,0,0.075); border: 2px solid rgba(255,0,0,0.4) !important; } +.token-counter div { display: inline; } +.token-counter span { padding: 0.1em 0.75em; } /* tooltips and statuses */ .infotext { overflow-wrap: break-word; } diff --git a/javascript/promptChecker.js b/javascript/promptChecker.js index a02119e97..e6a269d78 100644 --- a/javascript/promptChecker.js +++ b/javascript/promptChecker.js @@ -34,4 +34,6 @@ async function initPromptChecker() { setupBracketChecking('img2img_neg_prompt', 'img2img_negative_token_counter'); setupBracketChecking('control_prompt', 'control_token_counter'); setupBracketChecking('control_neg_prompt', 'control_negative_token_counter'); + setupBracketChecking('video_prompt', 'video_token_counter'); + setupBracketChecking('video_neg_prompt', 'video_negative_token_counter'); } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 10f1f66d4..872f2b12f 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -84,11 +84,11 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- .theme-preview { display: none; position: fixed; border: var(--spacing-sm) solid var(--neutral-600); box-shadow: 2px 2px 2px 2px var(--neutral-700); top: 0; bottom: 0; left: 0; right: 0; margin: auto; max-width: 75vw; z-index: 999; } /* txt2img/img2img specific */ -.block.token-counter{ position: absolute; right: 1em; min-width: 0 !important; width: auto; z-index: 100; top: -0.5em; } -.block.token-counter span{ background: var(--input-background-fill) !important; box-shadow: 0 0 0.0 0.3em rgba(192,192,192,0.15), inset 0 0 0.6em rgba(192,192,192,0.075); border: 2px solid rgba(192,192,192,0.4) !important; } -.block.token-counter.error span{ box-shadow: 0 0 0.0 0.3em rgba(255,0,0,0.15), inset 0 0 0.6em rgba(255,0,0,0.075); border: 2px solid rgba(255,0,0,0.4) !important; } -.block.token-counter div{ display: inline; } -.block.token-counter span{ padding: 0.1em 0.75em; } +.token-counter { position: absolute; right: 1em; min-width: 0 !important; width: auto; z-index: 100; top: -0.5em; } +.token-counter span { background: var(--input-background-fill) !important; box-shadow: 0 0 0.0 0.3em rgba(192,192,192,0.15), inset 0 0 0.6em rgba(192,192,192,0.075); border: 2px solid rgba(192,192,192,0.4) !important; } +.token-counter.error span { box-shadow: 0 0 0.0 0.3em rgba(255,0,0,0.15), inset 0 0 0.6em rgba(255,0,0,0.075); border: 2px solid rgba(255,0,0,0.4) !important; } +.token-counter div { display: inline; } +.token-counter span { padding: 0.1em 0.75em; } .performance { font-size: var(--text-xs); color: #444; } .performance p { display: inline-block; color: var(--primary-500) !important } .performance .time { margin-right: 0; } diff --git a/modules/sd_models.py b/modules/sd_models.py index 9f66d9814..43e926924 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -764,7 +764,6 @@ def set_diffuser_pipe(pipe, new_pipe_type): 'InstantIRPipeline', 'FluxFillPipeline', 'FluxControlPipeline', - 'StableVideoDiffusionPipeline', 'PixelSmithXLPipeline', 'PhotoMakerStableDiffusionXLPipeline', 'StableDiffusionXLInstantIDPipeline', @@ -781,6 +780,8 @@ def set_diffuser_pipe(pipe, new_pipe_type): cls = pipe.__class__.__name__ if cls in exclude: return pipe + if 'Video' in cls: + return pipe if 'Onnx' in cls: return pipe From 7a84a5e8ef70771b1cb05ebcc124ab040483f6d8 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 24 Mar 2025 19:17:44 -0400 Subject: [PATCH 056/122] cleanup video code Signed-off-by: Vladimir Mandic --- TODO.md | 8 +++----- installer.py | 4 ++-- modules/video_models/video_run.py | 1 + 3 files changed, 6 insertions(+), 7 deletions(-) diff --git a/TODO.md b/TODO.md index de0c7c8be..1ccc0a2dc 100644 --- a/TODO.md +++ b/TODO.md @@ -8,9 +8,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` - VAE Remote encode: SD15 and Flux.1 issues: -- Video: API support is TBD -- Video: Hunyuan Video I2V: transformers incompatibility -- Video: Hunyuan Video I2V: add 16ch vs 33ch processing +- Video: Hunyuan Video I2V: requires `transformers==4.47.1` - Video: Latte 1 T2V: dtype mismatch - Video: WAN 2.1 14B I2V 480p/720p: broken offload - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output @@ -24,8 +22,8 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: add generate context menu - Video: FasterCache and PyramidAttentionBroadcast granular config - Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN -- Video: OponSora v2 https://huggingface.co/hpcai-tech/Open-Sora-v2 -- Video: STG: https://github.com/huggingface/diffusers/blob/main/examples/community/README.md#spatiotemporal-skip-guidance +- Video: API support +- Video: STG: - Video SmoothCache: https://github.com/huggingface/diffusers/issues/11135 - FasterCache, PyramidAttentionBroadcast, SmoothCache general support diff --git a/installer.py b/installer.py index 12ec654b2..006603abe 100644 --- a/installer.py +++ b/installer.py @@ -536,9 +536,9 @@ def check_python(supported_minors=[9, 10, 11, 12], reason=None): # check diffusers version def check_diffusers(): t_start = time.time() - if args.skip_all or args.skip_git: + if args.skip_all or args.skip_git or args.experimental: return - sha = '5dbe4f5de6398159f8c2bedd371bc116683edbd3' # diffusers commit hash + sha = '1ddf3f3a19095344166ad7207ebc5be7a862d17e' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 95d75e149..5f60efd3c 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -66,6 +66,7 @@ def generate(*args, **kwargs): video_vae.set_vae_params(p) video_cache.set_cache(faster_cache=faster_cache, pyramid_attention_broadcast=pyramid_attention) video_utils.set_prompt(p) + p.task_args['num_inference_steps'] = p.steps p.task_args['width'] = p.width p.task_args['height'] = p.height p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil' From 90f887ac4af4984bd7433e9e5b0026c55dc1cd33 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Tue, 25 Mar 2025 03:50:21 +0300 Subject: [PATCH 057/122] Add dim checks to ck flash atten and fix dim check on dyn atten --- modules/devices.py | 21 ++++++++++++++++++--- modules/intel/ipex/attention.py | 12 ++++++------ modules/intel/ipex/hijacks.py | 2 +- modules/sd_hijack_dynamic_atten.py | 12 ++++++------ 4 files changed, 31 insertions(+), 16 deletions(-) diff --git a/modules/devices.py b/modules/devices.py index b2a488d7b..c35c8b909 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -387,8 +387,6 @@ def set_cudnn_params(): torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = True torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = True - if hasattr(torch.backends.cuda, "allow_fp16_bf16_reduction_math_sdp"): # only valid for torch >= 2.5 - torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(True) except Exception as e: log.warning(f'Torch matmul: {e}') if torch.backends.cudnn.is_available(): @@ -411,6 +409,7 @@ def override_ipex_math(): try: if hasattr(torch.xpu, "set_fp32_math_mode"): # not available with pure torch+xpu, requires ipex torch.xpu.set_fp32_math_mode(mode=torch.xpu.FP32MathMode.TF32) + torch.backends.mkldnn.allow_tf32 = True except Exception as e: log.warning(f'Torch ipex: {e}') @@ -432,6 +431,8 @@ def set_sdpa_params(): torch.backends.cuda.enable_flash_sdp('Flash attention' in opts.sdp_options) torch.backends.cuda.enable_mem_efficient_sdp('Memory attention' in opts.sdp_options) torch.backends.cuda.enable_math_sdp('Math attention' in opts.sdp_options) + if hasattr(torch.backends.cuda, "allow_fp16_bf16_reduction_math_sdp"): # only valid for torch >= 2.5 + torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(True) log.debug(f'Torch attention: type="sdpa" flash={"Flash attention" in opts.sdp_options} memory={"Memory attention" in opts.sdp_options} math={"Math attention" in opts.sdp_options}') except Exception as err: log.warning(f'Torch attention: type="sdpa" {err}') @@ -462,7 +463,21 @@ def set_sdpa_params(): @wraps(sdpa_pre_flash_atten) def sdpa_flash_atten(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None): if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: - return flash_attn_func(q=query.transpose(1, 2), k=key.transpose(1, 2), v=value.transpose(1, 2), dropout_p=dropout_p, causal=is_causal, softmax_scale=scale).transpose(1, 2) + is_unsqueezed = False + if query.dim() == 3: + query = query.unsqueeze(0) + is_unsqueezed = True + if key.dim() == 3: + key = key.unsqueeze(0) + if value.dim() == 3: + value = value.unsqueeze(0) + query = query.transpose(1, 2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + attn_output = flash_attn_func(q=query, k=key, v=value, dropout_p=dropout_p, causal=is_causal, softmax_scale=scale).transpose(1, 2) + if is_unsqueezed: + attn_output = attn_output.squeeze(0) + return attn_output else: return sdpa_pre_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale) torch.nn.functional.scaled_dot_product_attention = sdpa_flash_atten diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py index 400b59b66..177f5bc5e 100644 --- a/modules/intel/ipex/attention.py +++ b/modules/intel/ipex/attention.py @@ -61,13 +61,13 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop if query.device.type != "xpu": return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) is_unsqueezed = False - if len(query.shape) == 3: + if query.dim() == 3: query = query.unsqueeze(0) is_unsqueezed = True - if len(key.shape) == 3: - key = key.unsqueeze(0) - if len(value.shape) == 3: - value = value.unsqueeze(0) + if key.dim() == 3: + key = key.unsqueeze(0) + if value.dim() == 3: + value = value.unsqueeze(0) do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size = find_sdpa_slice_sizes(query.shape, key.shape, query.element_size(), slice_rate=attention_slice_rate, trigger_rate=sdpa_slice_trigger_rate) # Slice SDPA @@ -115,5 +115,5 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop else: hidden_states = original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) if is_unsqueezed: - hidden_states.squeeze(0) + hidden_states = hidden_states.squeeze(0) return hidden_states diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index e47065a62..b9630afd3 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -118,7 +118,7 @@ original_torch_bmm = torch.bmm @wraps(torch.bmm) def torch_bmm(input, mat2, *, out=None): if input.dtype != mat2.dtype: - mat2 = mat2.to(input.dtype) + mat2 = mat2.to(dtype=input.dtype) return original_torch_bmm(input, mat2, out=out) # Diffusers FreeU diff --git a/modules/sd_hijack_dynamic_atten.py b/modules/sd_hijack_dynamic_atten.py index 6c3e69e3a..7a394772a 100644 --- a/modules/sd_hijack_dynamic_atten.py +++ b/modules/sd_hijack_dynamic_atten.py @@ -56,13 +56,13 @@ if devices.sdpa_pre_dyanmic_atten is None: @wraps(devices.sdpa_pre_dyanmic_atten) def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, **kwargs): is_unsqueezed = False - if len(query.shape) == 3: + if query.dim() == 3: query = query.unsqueeze(0) is_unsqueezed = True - if len(key.shape) == 3: - key = key.unsqueeze(0) - if len(value.shape) == 3: - value = value.unsqueeze(0) + if key.dim() == 3: + key = key.unsqueeze(0) + if value.dim() == 3: + value = value.unsqueeze(0) do_batch_split, do_head_split, do_query_split, split_batch_size, split_head_size, split_query_size = find_sdpa_slice_sizes(query.shape, key.shape, query.element_size(), slice_rate=shared.opts.dynamic_attention_slice_rate, trigger_rate=shared.opts.dynamic_attention_trigger_rate) # Slice SDPA @@ -111,7 +111,7 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop else: hidden_states = devices.sdpa_pre_dyanmic_atten(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) if is_unsqueezed: - hidden_states.squeeze(0) + hidden_states = hidden_states.squeeze(0) return hidden_states From c0a2ac9c23e68e8a658d7e844d4ec646b6908f27 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 25 Mar 2025 10:16:46 -0400 Subject: [PATCH 058/122] update requirements Signed-off-by: Vladimir Mandic --- installer.py | 2 +- modules/sd_offload.py | 2 +- requirements.txt | 6 +++--- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/installer.py b/installer.py index 006603abe..cb80e3962 100644 --- a/installer.py +++ b/installer.py @@ -1302,7 +1302,7 @@ def check_ui(ver): def same(ver): core = ver['branch'] if ver is not None and 'branch' in ver else 'unknown' ui = ver['ui'] if ver is not None and 'ui' in ver else 'unknown' - return core == ui or (core == 'master' and ui == 'main') + return (core == ui) or (core == 'master' and ui == 'main') or (core == 'dev' and ui == 'dev') t_start = time.time() if not same(ver): diff --git a/modules/sd_offload.py b/modules/sd_offload.py index cadfd3c02..8ab114e93 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -127,7 +127,7 @@ class OffloadHook(accelerate.hooks.ModelHook): self.cpu = int(shared.cpu_memory * shared.opts.diffusers_offload_max_cpu_memory * 1024*1024*1024) self.offload_map = {} self.param_map = {} - gpu = f'{shared.gpu_memory * shared.opts.diffusers_offload_min_gpu_memory:.3f}-{shared.gpu_memory * shared.opts.diffusers_offload_max_gpu_memory}:{shared.gpu_memory}' + gpu = f'{(shared.gpu_memory * shared.opts.diffusers_offload_min_gpu_memory):.2f}-{(shared.gpu_memory * shared.opts.diffusers_offload_max_gpu_memory):.2f}:{shared.gpu_memory:.2f}' shared.log.info(f'Offload: type=balanced op=init watermark={self.min_watermark}-{self.max_watermark} gpu={gpu} cpu={shared.cpu_memory:.3f} limit={shared.opts.cuda_mem_fraction:.2f}') self.validate() super().__init__() diff --git a/requirements.txt b/requirements.txt index e64ada91b..28726d758 100644 --- a/requirements.txt +++ b/requirements.txt @@ -32,7 +32,7 @@ pi-heif # versioned rich==13.9.4 -safetensors==0.5.2 +safetensors==0.5.3 tensordict==0.1.2 peft==0.14.0 httpx==0.24.1 @@ -51,8 +51,8 @@ numpy==1.26.4 numba==0.59.1 protobuf==4.25.3 pytorch_lightning==1.9.4 -tokenizers==0.21.0 -transformers==4.49.0 +tokenizers==0.21.1 +transformers==4.50.0 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 From aa15ed1fa71c4ce0d1eaa565dd571c0c295ad870 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 25 Mar 2025 14:09:34 -0400 Subject: [PATCH 059/122] reapply offload after applying lora Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 7 ++- modules/lora/extra_networks_lora.py | 3 +- modules/lora/lora_extract.py | 4 +- modules/lora/networks.py | 89 ++++++++++++++++------------- modules/sd_models_utils.py | 4 -- 5 files changed, 57 insertions(+), 50 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index cfe49ffee..64262450b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,15 +1,15 @@ # Change Log for SD.Next -## Update for 2025-03-24 +## Update for 2025-03-25 -### Highlights for 2025-03-24 +### Highlights for 2025-03-25 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to remote VAE, additional docs/guides -### Details for 2025-03-24 +### Details for 2025-03-25 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -120,6 +120,7 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - fix sd35 with batch processing - fix extra networks cover and inline views - fix token counter error style with modernui + - improve lora compatibility with balanced offload ## Update for 2025-02-28 diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 2167f97ac..8aeb3f473 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -3,7 +3,7 @@ import os import re import numpy as np from modules.lora import networks, network_overrides -from modules import extra_networks, shared +from modules import extra_networks, shared, sd_models debug = os.environ.get('SD_SCRIPT_DEBUG', None) is not None @@ -171,6 +171,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if len(networks.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or force_diffusers) and step == 0: infotext(p) prompt(p) + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) # TODO lora: required for flux to reapply offload after lora has been applied, but fails with oom if (has_changed or force_diffusers) and len(include) == 0: # print only once shared.log.info(f'Network load: type=LoRA apply={[n.name for n in networks.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') diff --git a/modules/lora/lora_extract.py b/modules/lora/lora_extract.py index 21351d9e3..5050187a2 100644 --- a/modules/lora/lora_extract.py +++ b/modules/lora/lora_extract.py @@ -5,7 +5,7 @@ import datetime import torch from safetensors.torch import save_file import gradio as gr -from rich import progress as p +from rich import progress as rp from modules import shared, devices from modules.ui_common import create_refresh_button from modules.call_queue import wrap_gradio_gpu_call @@ -134,7 +134,7 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): shared.log.debug(f'LoRA extract: modules={modules} maxrank={maxrank} auto={auto_rank} ratio={rank_ratio} fn="{fn}"') shared.state.begin('LoRA extract') - with p.Progress(p.TextColumn('[cyan]LoRA extract'), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TextColumn('[cyan]{task.description}'), console=shared.console) as progress: + with rp.Progress(rp.TextColumn('[cyan]LoRA extract'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) as progress: if 'te' in modules and getattr(shared.sd_model, 'text_encoder', None) is not None: modules = shared.sd_model.text_encoder.named_modules() diff --git a/modules/lora/networks.py b/modules/lora/networks.py index f6e8acadb..b22825639 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -387,16 +387,20 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. weight = self.weight.to(devices.device) except Exception: weight = self.weight + updown, ex_bias = module.calc_updown(weight) - if batch_updown is not None and updown is not None: - batch_updown += updown.to(batch_updown.device) - else: - batch_updown = updown - if batch_ex_bias is not None and ex_bias is not None: - batch_ex_bias += ex_bias.to(batch_ex_bias.device) - else: - batch_ex_bias = ex_bias + if updown is not None: + if batch_updown is not None: + batch_updown += updown.to(batch_updown.device) + else: + batch_updown = updown.to(devices.device) + if ex_bias is not None: + if batch_ex_bias: + batch_ex_bias += ex_bias.to(batch_ex_bias.device) + else: + batch_ex_bias = ex_bias.to(devices.device) timer.calc += time.time() - t0 + if shared.opts.diffusers_offload_mode == "sequential": t0 = time.time() if batch_updown is not None: @@ -418,17 +422,18 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False): if lora_weights is None: - return self.weight + return None if deactivate: lora_weights *= -1 if model_weights is None: # weights are used if provided-from-backup else use self.weight model_weights = self.weight # TODO lora: add other quantization types + weight = None if self.__class__.__name__ == 'Linear4bit' and bnb is not None: try: dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(devices.device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) - self.weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) + weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) except Exception as e: shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') else: @@ -436,12 +441,13 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G new_weight = model_weights.to(devices.device) + lora_weights.to(devices.device) except Exception: new_weight = model_weights + lora_weights # try without device cast - self.weight = torch.nn.Parameter(new_weight, requires_grad=False) + weight = torch.nn.Parameter(new_weight, requires_grad=False) try: - self.weight = self.weight.to(device=devices.device) # required since quantization happens only during .to call, not during params creation + # weight = weight.to(device=devices.device) # required since quantization happens only during .to call, not during params creation + pass except Exception: pass # may fail if weights is meta tensor - return self.weight + return weight def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False): @@ -452,24 +458,27 @@ def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if not isinstance(bias_backup, bool): bias_backup = True if not weights_backup and not bias_backup: - return None, None + return t0 = time.time() if weights_backup: - if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 + if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - self.weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate) + weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate) + if weight is not None: + self.weight = weight if bias_backup: if ex_bias is not None: - self.bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate) + bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate) + if bias is not None: + self.bias = bias if hasattr(self, "qweight") and hasattr(self, "freeze"): self.freeze() timer.apply += time.time() - t0 - return self.weight.device, self.weight.dtype def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, orig_device: torch.device, deactivate: bool = False): @@ -484,14 +493,18 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn if updown is not None and len(weights_backup.shape) == 4 and weights_backup.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - self.weight = network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate) + weight = network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate) + if weight is not None: + self.weight = weight else: self.weight = torch.nn.Parameter(weights_backup.to(device=orig_device), requires_grad=False) if bias_backup is not None: self.bias = None if ex_bias is not None: - self.weight = network_add_weights(self, model_weights=weights_backup, lora_weights=ex_bias, deactivate=deactivate) + bias = network_add_weights(self, model_weights=weights_backup, lora_weights=ex_bias, deactivate=deactivate) + if bias: + self.weight = bias else: self.bias = torch.nn.Parameter(bias_backup.to(device=orig_device), requires_grad=False) @@ -499,7 +512,6 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn self.freeze() timer.apply += time.time() - t0 - return self.weight.device, self.weight.dtype def network_deactivate(include=[], exclude=[]): @@ -529,8 +541,6 @@ def network_deactivate(include=[], exclude=[]): pbar = nullcontext() with devices.inference_context(), pbar: applied_layers.clear() - weights_devices = [] - weights_dtypes = [] for component in modules.keys(): orig_device = getattr(sd_model, component, None).device for _, module in modules[component]: @@ -541,11 +551,9 @@ def network_deactivate(include=[], exclude=[]): continue batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, use_previous=True) if shared.opts.lora_fuse_diffusers: - weights_device, weights_dtype = network_apply_direct(module, batch_updown, batch_ex_bias, deactivate=True) + network_apply_direct(module, batch_updown, batch_ex_bias, deactivate=True) else: - weights_device, weights_dtype = network_apply_weights(module, batch_updown, batch_ex_bias, orig_device, deactivate=True) - weights_devices.append(weights_device) - weights_dtypes.append(weights_dtype) + network_apply_weights(module, batch_updown, batch_ex_bias, orig_device, deactivate=True) if batch_updown is not None or batch_ex_bias is not None: applied_layers.append(network_layer_name) del batch_updown, batch_ex_bias @@ -555,8 +563,7 @@ def network_deactivate(include=[], exclude=[]): timer.deactivate = time.time() - t0 if debug and len(previously_loaded_networks) > 0: - weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403 # pylint: disable=R1718 - shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() if shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") @@ -584,12 +591,12 @@ def network_activate(include=[], exclude=[]): else: task = None pbar = nullcontext() + applied_weight = 0 + applied_bias = 0 with devices.inference_context(), pbar: wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) if len(loaded_networks) > 0 else () applied_layers.clear() backup_size = 0 - weights_devices = [] - weights_dtypes = [] for component in modules.keys(): orig_device = getattr(sd_model, component, None).device for _, module in modules[component]: @@ -602,24 +609,26 @@ def network_activate(include=[], exclude=[]): backup_size += network_backup_weights(module, network_layer_name, wanted_names) batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name) if shared.opts.lora_fuse_diffusers: - weights_device, weights_dtype = network_apply_direct(module, batch_updown, batch_ex_bias) + network_apply_direct(module, batch_updown, batch_ex_bias) else: - weights_device, weights_dtype = network_apply_weights(module, batch_updown, batch_ex_bias, orig_device) - weights_devices.append(weights_device) - weights_dtypes.append(weights_dtype) + network_apply_weights(module, batch_updown, batch_ex_bias, orig_device) if batch_updown is not None or batch_ex_bias is not None: applied_layers.append(network_layer_name) + if batch_updown is not None: + applied_weight += 1 + if batch_ex_bias is not None: + applied_bias += 1 del batch_updown, batch_ex_bias module.network_current_names = wanted_names if task is not None: - pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={active_components} layers={total} apply={len(applied_layers)} backup={backup_size}') + pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size}') if task is not None and len(applied_layers) == 0: pbar.remove_task(task) # hide progress bar for no action timer.activate += time.time() - t0 if debug and len(loaded_networks) > 0: - weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403 # pylint: disable=R1718 - shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() - if shared.opts.diffusers_offload_mode == "sequential": - sd_models.set_diffuser_offload(sd_model, op="model") + if len(loaded_networks) > 0 and (applied_weight > 0 or applied_bias > 0): + if shared.opts.diffusers_offload_mode == "sequential": + sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/sd_models_utils.py b/modules/sd_models_utils.py index 93367197c..cfbf38136 100644 --- a/modules/sd_models_utils.py +++ b/modules/sd_models_utils.py @@ -67,19 +67,15 @@ def read_state_dict(checkpoint_file, map_location=None, what:str='model'): # pyl return None if shared.opts.stream_load: if extension.lower() == ".safetensors": - # shared.log.debug('Model weights loading: type=safetensors mode=buffered') buffer = f.read() pl_sd = safetensors.torch.load(buffer) else: - # shared.log.debug('Model weights loading: type=checkpoint mode=buffered') buffer = io.BytesIO(f.read()) pl_sd = torch.load(buffer, map_location='cpu') else: if extension.lower() == ".safetensors": - # shared.log.debug('Model weights loading: type=safetensors mode=mmap') pl_sd = safetensors.torch.load_file(checkpoint_file, device='cpu') else: - # shared.log.debug('Model weights loading: type=checkpoint mode=direct') pl_sd = torch.load(f, map_location='cpu') sd = get_state_dict_from_checkpoint(pl_sd) del pl_sd From cd71fe51ff0a36794f9c85a6584eb67bdce19ed7 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 25 Mar 2025 14:56:53 -0400 Subject: [PATCH 060/122] fix lora change with balanced offload Signed-off-by: Vladimir Mandic --- modules/lora/extra_networks_lora.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 8aeb3f473..526926c9a 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -166,12 +166,12 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if has_changed: networks.network_deactivate(include, exclude) networks.network_activate(include, exclude) + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) # TODO lora: required for flux to reapply offload after lora has been applied, but fails with oom debug_log(f'Network load: type=LoRA previous={[n.name for n in networks.previously_loaded_networks]} current={[n.name for n in networks.loaded_networks]} changed') if len(networks.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or force_diffusers) and step == 0: infotext(p) prompt(p) - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) # TODO lora: required for flux to reapply offload after lora has been applied, but fails with oom if (has_changed or force_diffusers) and len(include) == 0: # print only once shared.log.info(f'Network load: type=LoRA apply={[n.name for n in networks.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') From d74e06cba3c1d2bc66b05624dff808cccdef8891 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Wed, 26 Mar 2025 17:13:43 +0900 Subject: [PATCH 061/122] zluda v3.9.2 --- modules/zluda_installer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 4b915c724..236ca85f5 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -100,7 +100,7 @@ def install() -> None: return platform = "windows" - commit = os.environ.get("ZLUDA_HASH", "ae0540beb129ffd140226ce956b386619b38f84c") + commit = os.environ.get("ZLUDA_HASH", "dba64c0966df2c71e82255e942c96e2e1cea3a2d") if nightly: platform = "nightly-" + platform urllib.request.urlretrieve(f'https://github.com/lshqqytiger/ZLUDA/releases/download/rel.{commit}/ZLUDA-{platform}-rocm{rocm.version[0]}-amd64.zip', '_zluda') From 16ed031781948821b3e9d9a7d43ca1c47fe5a834 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Wed, 26 Mar 2025 17:19:34 +0900 Subject: [PATCH 062/122] update changelog --- CHANGELOG.md | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 64262450b..ffbc2922a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,8 @@ And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to remote VAE, additional docs/guides +**Flash Attention 2** and Sage Attention is now available on ZLUDA backend! + ### Details for 2025-03-25 - **Video tab** @@ -103,6 +105,9 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - **ROCm** - add `--upgrade` to torch_command when using `--use-nightly` - disable fp16 for gfx1102 (rx 7600 and rx 7500 series) gpus +- **ZLUDA** + - add `torch.compile` support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) + - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled From f2ba630f601f184a8f61eb8b52ac4bdff544275e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 26 Mar 2025 07:09:10 -0400 Subject: [PATCH 063/122] fix transformers Signed-off-by: Vladimir Mandic --- requirements.txt | 2 +- wiki | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/requirements.txt b/requirements.txt index 28726d758..30a4e6065 100644 --- a/requirements.txt +++ b/requirements.txt @@ -52,7 +52,7 @@ numba==0.59.1 protobuf==4.25.3 pytorch_lightning==1.9.4 tokenizers==0.21.1 -transformers==4.50.0 +transformers==4.50.1 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 diff --git a/wiki b/wiki index abc426fa7..b654db617 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit abc426fa76531e1c626b2e7275bd8925e4beef9b +Subproject commit b654db61752a4d1925397fd08bb9f576943015a2 From 6a5e253ecf64867c9477f863857cdb94458f6a18 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 26 Mar 2025 10:50:09 -0400 Subject: [PATCH 064/122] add infiniteyou Signed-off-by: Vladimir Mandic --- .pylintrc | 2 + .ruff.toml | 2 + CHANGELOG.md | 13 +- TODO.md | 2 - modules/infiniteyou/__init__.py | 2 + .../infiniteyou/pipeline_flux_infusenet.py | 612 ++++++++++++++++++ modules/infiniteyou/pipeline_infu_flux.py | 325 ++++++++++ modules/infiniteyou/resampler.py | 121 ++++ modules/lora/networks.py | 2 +- modules/processing_args.py | 10 +- scripts/infiniteyou_ext.py | 121 ++++ 11 files changed, 1200 insertions(+), 12 deletions(-) create mode 100644 modules/infiniteyou/__init__.py create mode 100644 modules/infiniteyou/pipeline_flux_infusenet.py create mode 100644 modules/infiniteyou/pipeline_infu_flux.py create mode 100644 modules/infiniteyou/resampler.py create mode 100644 scripts/infiniteyou_ext.py diff --git a/.pylintrc b/.pylintrc index b472d26a0..bbe29b197 100644 --- a/.pylintrc +++ b/.pylintrc @@ -38,6 +38,8 @@ ignore-paths=/usr/lib/.*$, modules/todo, modules/unipc, modules/xadapter, + modules/infiniteyou, + modules/flash_attn_triton_amd, repositories, extensions-builtin/Lora, extensions-builtin/sd-webui-agent-scheduler, diff --git a/.ruff.toml b/.ruff.toml index 0fc9de8b3..48f2e9026 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -33,6 +33,8 @@ exclude = [ "modules/todo", "modules/unipc", "modules/xadapter", + "modules/infiniteyou", + "modules/flash_attn_triton_amd", "repositories", "extensions-builtin/Lora", "extensions-builtin/sd-extension-chainner/nodes", diff --git a/CHANGELOG.md b/CHANGELOG.md index ffbc2922a..072af4c91 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2025-03-25 +## Update for 2025-03-26 -### Highlights for 2025-03-25 +### Highlights for 2025-03-26 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! @@ -11,7 +11,7 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r **Flash Attention 2** and Sage Attention is now available on ZLUDA backend! -### Details for 2025-03-25 +### Details for 2025-03-26 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -47,7 +47,7 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - different video models support different video resolutions, frame counts, etc. and may require specific settings - see model links for details - see *ToDo/Limitations* section for additional notes -- **Models** +- **Models & Pipelines** - [THUDM CogView 4](https://huggingface.co/THUDM/CogView4-6B) **6B** variant new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer fully supports offloading and on-the-fly quantization @@ -57,6 +57,11 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r big update to previous SANA model fully supports offloading and on-the-fly quantization simply select from *networks -> models -> reference* + - [ByteDance InfiniteYou](https://github.com/bytedance/InfiniteYou/): Flexible Photo Recrafting While Preserving Your Identity + face-transfer model for FLUX.1 + select from *Scripts -> InfiniteYou* + its large, ~12GB on top of FLUX.1 base model so make sure you have offloading and quantization setup + *note* model will be auto-downloaded on first use - New [zer0int CLiP-L](https://huggingface.co/zer0int/CLIP-Registers-Gated_MLP-ViT-L-14) models: download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) load using *settings -> text encoder* diff --git a/TODO.md b/TODO.md index 1ccc0a2dc..d9debcb5d 100644 --- a/TODO.md +++ b/TODO.md @@ -6,8 +6,6 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ### Issues/Limitations -- VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3` -- VAE Remote encode: SD15 and Flux.1 issues: - Video: Hunyuan Video I2V: requires `transformers==4.47.1` - Video: Latte 1 T2V: dtype mismatch - Video: WAN 2.1 14B I2V 480p/720p: broken offload diff --git a/modules/infiniteyou/__init__.py b/modules/infiniteyou/__init__.py new file mode 100644 index 000000000..142921909 --- /dev/null +++ b/modules/infiniteyou/__init__.py @@ -0,0 +1,2 @@ +from .pipeline_flux_infusenet import FluxInfuseNetPipeline +from .pipeline_infu_flux import InfUFluxPipeline diff --git a/modules/infiniteyou/pipeline_flux_infusenet.py b/modules/infiniteyou/pipeline_flux_infusenet.py new file mode 100644 index 000000000..38fa186f9 --- /dev/null +++ b/modules/infiniteyou/pipeline_flux_infusenet.py @@ -0,0 +1,612 @@ +# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates. +# Copyright (c) 2024 Black Forest Labs, The HuggingFace Team and The InstantX Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import inspect +from typing import Any, Callable, Dict, List, Optional, Union + +import numpy as np +import torch +from diffusers import FluxControlNetPipeline +from diffusers.models.controlnet_flux import FluxControlNetModel, FluxMultiControlNetModel +from diffusers.image_processor import PipelineImageInput +from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput +from diffusers.utils import is_torch_xla_available, logging + + +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 + + +# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.16, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + sigmas: Optional[List[float]] = None, + **kwargs, +): + r""" + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, + `num_inference_steps` and `sigmas` must be `None`. + sigmas (`List[float]`, *optional*): + Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, + `num_inference_steps` and `timesteps` must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None and sigmas is not None: + raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + elif sigmas is not None: + accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accept_sigmas: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" sigmas schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +class FluxInfuseNetPipeline(FluxControlNetPipeline): + @torch.no_grad() + def __call__( + self, + prompt: Union[str, List[str]] = None, + prompt_2: Optional[Union[str, List[str]]] = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 28, + timesteps: List[int] = None, + guidance_scale: float = 3.5, + id_image: PipelineImageInput = None, + controlnet_guidance_scale: float = 1.0, + control_guidance_start: Union[float, List[float]] = 0.0, + control_guidance_end: Union[float, List[float]] = 1.0, + control_image: PipelineImageInput = None, + control_mode: Optional[Union[int, List[int]]] = None, + controlnet_conditioning_scale: Union[float, List[float]] = 1.0, + 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, + pooled_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 = 512, + + # ID-specific parameters + controlnet_prompt_embeds: Optional[torch.FloatTensor] = None, + + # True CFG parameters + true_guidance_scale: float = 1.0, + negative_prompt: Optional[Union[str, List[str]]] = None, + negative_prompt_2: Optional[Union[str, List[str]]] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + ): + 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. + prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is + will be used 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 7.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. + controlnet_guidance_scale (`float`, *optional*, defaults to 7.0): + Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). + `controlnet_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. + control_guidance_start (`float` or `List[float]`, *optional*, defaults to 0.0): + The percentage of total steps at which the ControlNet starts applying. + control_guidance_end (`float` or `List[float]`, *optional*, defaults to 1.0): + The percentage of total steps at which the ControlNet stops applying. + control_image (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,: + `List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`): + The ControlNet input condition to provide guidance to the `unet` for generation. If the type is + specified as `torch.Tensor`, it is passed to ControlNet as is. `PIL.Image.Image` can also be accepted + as an image. The dimensions of the output image defaults to `image`'s dimensions. If height and/or + width are passed, `image` is resized accordingly. If multiple ControlNets are specified in `init`, + images must be passed as a list such that each element of the list can be correctly batched for input + to a single ControlNet. + controlnet_conditioning_scale (`float` or `List[float]`, *optional*, defaults to 1.0): + The outputs of the ControlNet are multiplied by `controlnet_conditioning_scale` before they are added + to the residual in the original `unet`. If multiple ControlNets are specified in `init`, you can set + the corresponding scale as a list. + control_mode (`int` or `List[int]`,, *optional*, defaults to None): + The control mode when applying ControlNet-Union. + 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. + pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. + If not provided, pooled text embeddings will be generated from `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.flux.FluxPipelineOutput`] 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 512): Maximum sequence length to use with the `prompt`. + controlnet_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated embeddings for the InfuseNet. Can be used to easily tweak inputs, *e.g.* image embeddings. + If not provided, embeddings will be generated from `prompt` or `prompt_embeds` input arguments. + true_guidance_scale (`float`, *optional*, defaults to 1.0): + True CFG scale as defined in [Classifier-Free Diffusion Guidance]((https://arxiv.org/abs/2207.12598). + negative_prompt (`str` or `List[str]`, *optional*): + The negative prompt or negative prompts to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds`. instead. + negative_prompt_2 (`str` or `List[str]`, *optional*): + The negative prompt or negative prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, + `negative_prompt` is will be used instead. + 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 text embeddings will be generated from `negative_prompt` input + argument. + negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative pooled text embeddings will be generated from + `negative_prompt` input argument. + + 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 + + if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list): + control_guidance_start = len(control_guidance_end) * [control_guidance_start] + elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list): + control_guidance_end = len(control_guidance_start) * [control_guidance_end] + elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list): + mult = len(self.controlnet.nets) if isinstance(self.controlnet, FluxMultiControlNetModel) else 1 + control_guidance_start, control_guidance_end = ( + mult * [control_guidance_start], + mult * [control_guidance_end], + ) + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + prompt_2, + height, + width, + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_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._controlnet_guidance_scale = controlnet_guidance_scale + self._true_guidance_scale = true_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 + dtype = self.transformer.dtype + + lora_scale = ( + self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None + ) + ( + prompt_embeds, + pooled_prompt_embeds, + text_ids, + ) = self.encode_prompt( + prompt=prompt, + prompt_2=prompt_2, + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + if negative_prompt is not None or (negative_prompt_embeds is not None and negative_pooled_prompt_embeds is not None): + ( + negative_prompt_embeds, + negative_pooled_prompt_embeds, + negative_text_ids, + ) = self.encode_prompt( + prompt=negative_prompt, + prompt_2=negative_prompt_2, + prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=negative_pooled_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + + if controlnet_prompt_embeds is None: + controlnet_prompt_embeds = prompt_embeds + ( + controlnet_prompt_embeds, + pooled_prompt_embeds, + controlnet_text_ids, + ) = self.encode_prompt( + prompt=prompt, + prompt_2=prompt_2, + prompt_embeds=controlnet_prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + + # 3. Prepare control image + num_channels_latents = self.transformer.config.in_channels // 4 + if isinstance(self.controlnet, FluxControlNetModel) or True: + control_image = self.prepare_image( + image=control_image, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=self.vae.dtype, + ) + height, width = control_image.shape[-2:] + + # xlab controlnet has a input_hint_block and instantx controlnet does not + controlnet_blocks_repeat = False if self.controlnet.input_hint_block is None else True + if self.controlnet.input_hint_block is None: + # vae encode + control_image = self.vae.encode(control_image).latent_dist.sample() + control_image = (control_image - self.vae.config.shift_factor) * self.vae.config.scaling_factor + + # pack + height_control_image, width_control_image = control_image.shape[2:] + control_image = self._pack_latents( + control_image, + batch_size * num_images_per_prompt, + num_channels_latents, + height_control_image, + width_control_image, + ) + + # Here we ensure that `control_mode` has the same length as the control_image. + if control_mode is not None: + if not isinstance(control_mode, int): + raise ValueError(" For `FluxControlNet`, `control_mode` should be an `int` or `None`") + control_mode = torch.tensor(control_mode).to(device, dtype=torch.long) + control_mode = control_mode.view(-1, 1).expand(control_image.shape[0], 1) + + elif isinstance(self.controlnet, FluxMultiControlNetModel): + control_images = [] + # xlab controlnet has a input_hint_block and instantx controlnet does not + controlnet_blocks_repeat = False if self.controlnet.nets[0].input_hint_block is None else True + for _i, control_image_ in enumerate(control_image): + control_image_ = self.prepare_image( + image=control_image_, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=self.vae.dtype, + ) + height, width = control_image_.shape[-2:] + + if self.controlnet.nets[0].input_hint_block is None: + # vae encode + control_image_ = self.vae.encode(control_image_).latent_dist.sample() + control_image_ = (control_image_ - self.vae.config.shift_factor) * self.vae.config.scaling_factor + + # pack + height_control_image, width_control_image = control_image_.shape[2:] + control_image_ = self._pack_latents( + control_image_, + batch_size * num_images_per_prompt, + num_channels_latents, + height_control_image, + width_control_image, + ) + control_images.append(control_image_) + + control_image = control_images + + # Here we ensure that `control_mode` has the same length as the control_image. + if isinstance(control_mode, list) and len(control_mode) != len(control_image): + raise ValueError("For Multi-ControlNet, `control_mode` must be a list of the same length as the number of controlnets (control images) specified") + if not isinstance(control_mode, list): + control_mode = [control_mode] * len(control_image) + # set control mode + control_modes = [] + for cmode in control_mode: + if cmode is None: + cmode = -1 + control_mode = torch.tensor(cmode).expand(control_images[0].shape[0]).to(device, dtype=torch.long) + control_modes.append(control_mode) + control_mode = control_modes + + # 4. Prepare latent variables + num_channels_latents = self.transformer.config.in_channels // 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, + ) + + # 5. Prepare timesteps + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) + image_seq_len = latents.shape[1] + 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, + ) + + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self._num_timesteps = len(timesteps) + + # 6. Create tensor stating which controlnets to keep + controlnet_keep = [] + for i in range(len(timesteps)): + keeps = [ + 1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e) + for s, e in zip(control_guidance_start, control_guidance_end) + ] + controlnet_keep.append(keeps[0] if isinstance(self.controlnet, FluxControlNetModel) else keeps) + + # 7. Denoising loop + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + if self.interrupt: + continue + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latents.shape[0]).to(latents.dtype) + + if isinstance(self.controlnet, FluxMultiControlNetModel): + use_guidance = self.controlnet.nets[0].config.guidance_embeds + else: + use_guidance = self.controlnet.config.guidance_embeds + + guidance = torch.tensor([controlnet_guidance_scale], device=device) if use_guidance else None + guidance = guidance.expand(latents.shape[0]) if guidance is not None else None + + if isinstance(controlnet_keep[i], list): + if not isinstance(controlnet_conditioning_scale, list): + controlnet_conditioning_scale = len(controlnet_keep) * [controlnet_conditioning_scale] + cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])] + controlnet_conditioning_scale = controlnet_conditioning_scale[0] + else: + controlnet_cond_scale = controlnet_conditioning_scale + if isinstance(controlnet_cond_scale, list): + controlnet_cond_scale = controlnet_cond_scale[0] + cond_scale = controlnet_cond_scale * controlnet_keep[i] + + # controlnet + controlnet_block_samples, controlnet_single_block_samples = self.controlnet( + hidden_states=latents, + controlnet_cond=control_image, + controlnet_mode=control_mode, + conditioning_scale=cond_scale[0], + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=controlnet_prompt_embeds, + txt_ids=controlnet_text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + ) + + guidance = ( + torch.tensor([guidance_scale], device=device) if self.transformer.config.guidance_embeds else None + ) + guidance = guidance.expand(latents.shape[0]) if guidance is not None else None + + noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + controlnet_block_samples=controlnet_block_samples, + controlnet_single_block_samples=controlnet_single_block_samples, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + controlnet_blocks_repeat=controlnet_blocks_repeat, + )[0] + + # perform true CFG + if negative_prompt_embeds is not None and negative_pooled_prompt_embeds is not None and negative_text_ids is not None: + noise_pred_uncond = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + controlnet_block_samples=None, + controlnet_single_block_samples=None, + txt_ids=negative_text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + controlnet_blocks_repeat=controlnet_blocks_repeat, + )[0] + + noise_pred = noise_pred_uncond + true_guidance_scale * (noise_pred - noise_pred_uncond) + + # 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) + + # 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 / self.vae.config.scaling_factor) + self.vae.config.shift_factor + + image = self.vae.decode(latents, 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) diff --git a/modules/infiniteyou/pipeline_infu_flux.py b/modules/infiniteyou/pipeline_infu_flux.py new file mode 100644 index 000000000..8ae9f6e95 --- /dev/null +++ b/modules/infiniteyou/pipeline_infu_flux.py @@ -0,0 +1,325 @@ +# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates. All rights reserved. + +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at + +# http://www.apache.org/licenses/LICENSE-2.0 + +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +import os +import random +from typing import Optional + +import cv2 +import numpy as np +import torch +from diffusers.models import FluxControlNetModel +from facexlib.recognition import init_recognition_model +from huggingface_hub import snapshot_download +from insightface.app import FaceAnalysis +from insightface.utils import face_align +from PIL import Image + +from modules import shared, devices, model_quant +from .pipeline_flux_infusenet import FluxInfuseNetPipeline +from .resampler import Resampler + + +def seed_everything(seed, deterministic=False): + """Set random seed. + + Args: + seed (int): Seed to be used. + deterministic (bool): Whether to set the deterministic option for + CUDNN backend, i.e., set `torch.backends.cudnn.deterministic` + to True and `torch.backends.cudnn.benchmark` to False. + Default: False. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + os.environ['PYTHONHASHSEED'] = str(seed) + if deterministic: + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + + +def retrieve_latents( + encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample" +): + if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": + return encoder_output.latent_dist.sample(generator) + elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": + return encoder_output.latent_dist.mode() + elif hasattr(encoder_output, "latents"): + return encoder_output.latents + else: + raise AttributeError("Could not access latents of provided encoder_output") + + +# modified from https://github.com/instantX-research/InstantID/blob/main/pipeline_stable_diffusion_xl_instantid.py +def draw_kps(image_pil, kps, color_list=[(255,0,0), (0,255,0), (0,0,255), (255,255,0), (255,0,255)]): + stickwidth = 4 + limbSeq = np.array([[0, 2], [1, 2], [3, 2], [4, 2]]) + kps = np.array(kps) + + w, h = image_pil.size + out_img = np.zeros([h, w, 3]) + + for i in range(len(limbSeq)): + index = limbSeq[i] + color = color_list[index[0]] + + x = kps[index][:, 0] + y = kps[index][:, 1] + length = ((x[0] - x[1]) ** 2 + (y[0] - y[1]) ** 2) ** 0.5 + angle = math.degrees(math.atan2(y[0] - y[1], x[0] - x[1])) + polygon = cv2.ellipse2Poly((int(np.mean(x)), int(np.mean(y))), (int(length / 2), stickwidth), int(angle), 0, 360, 1) + out_img = cv2.fillConvexPoly(out_img.copy(), polygon, color) + out_img = (out_img * 0.6).astype(np.uint8) + + for idx_kp, kp in enumerate(kps): + color = color_list[idx_kp] + x, y = kp + out_img = cv2.circle(out_img.copy(), (int(x), int(y)), 10, color, -1) + + out_img_pil = Image.fromarray(out_img.astype(np.uint8)) + return out_img_pil + + +def extract_arcface_bgr_embedding(in_image, landmark, arcface_model=None, in_settings=None): # pylint: disable=unused-argument + kps = landmark + arc_face_image = face_align.norm_crop(in_image, landmark=np.array(kps), image_size=112) + arc_face_image = torch.from_numpy(arc_face_image).unsqueeze(0).permute(0,3,1,2) / 255. + arc_face_image = 2 * arc_face_image - 1 + arc_face_image = arc_face_image.cuda().contiguous() + if arcface_model is None: + arcface_model = init_recognition_model('arcface', device=devices.device) + face_emb = arcface_model(arc_face_image)[0] # [512], normalized + return face_emb + + +def resize_and_pad_image(source_img, target_img_size): + # Get original and target sizes + source_img_size = source_img.size + target_width, target_height = target_img_size + + # Determine the new size based on the shorter side of target_img + if target_width <= target_height: + new_width = target_width + new_height = int(target_width * (source_img_size[1] / source_img_size[0])) + else: + new_height = target_height + new_width = int(target_height * (source_img_size[0] / source_img_size[1])) + + # Resize the source image using LANCZOS interpolation for high quality + resized_source_img = source_img.resize((new_width, new_height), Image.Resampling.LANCZOS) + + # Compute padding to center resized image + pad_left = (target_width - new_width) // 2 + pad_top = (target_height - new_height) // 2 + + # Create a new image with white background + padded_img = Image.new("RGB", target_img_size, (255, 255, 255)) + padded_img.paste(resized_source_img, (pad_left, pad_top)) + + return padded_img + + +class InfUFluxPipeline: + def __init__( + self, + pipe, + image_proj_num_tokens=8, + infu_flux_version='v1.0', + model_version='aes_stage2', + ): + + self.infu_flux_version = infu_flux_version + self.model_version = model_version + + # Load pipeline + local_path = snapshot_download(repo_id='ByteDance/InfiniteYou', cache_dir=shared.opts.hfcache_dir) + infiniteyou_path = os.path.join(local_path, f'infu_flux_{infu_flux_version}', model_version) + infusenet_path = os.path.join(infiniteyou_path, 'InfuseNetModel') + quant_args = model_quant.create_config() + # quant_args = {} + + self.infusenet = FluxControlNetModel.from_pretrained( + infusenet_path, + torch_dtype=devices.dtype, + **quant_args, + ) + + self.pipe = FluxInfuseNetPipeline( + vae=pipe.vae, + text_encoder=pipe.text_encoder, + text_encoder_2=pipe.text_encoder_2, + tokenizer=pipe.tokenizer, + tokenizer_2=pipe.tokenizer_2, + transformer=pipe.transformer, + scheduler=pipe.scheduler, + controlnet=self.infusenet, + ) + + # Load image proj model + num_tokens = image_proj_num_tokens + image_emb_dim = 512 + image_proj_model = Resampler( + dim=1280, + depth=4, + dim_head=64, + heads=20, + num_queries=num_tokens, + embedding_dim=image_emb_dim, + output_dim=4096, + ff_mult=4, + ) + image_proj_model_path = os.path.join(infiniteyou_path, 'image_proj_model.bin') + ipm_state_dict = torch.load(image_proj_model_path, map_location="cpu") + image_proj_model.load_state_dict(ipm_state_dict['image_proj']) + del ipm_state_dict + image_proj_model.to(device=devices.device, dtype=devices.dtype) + image_proj_model.eval() + + self.image_proj_model = image_proj_model + + # Load face encoder + insightface_root_path = os.path.join(local_path, 'supports', 'insightface') + self.app_640 = FaceAnalysis(name='antelopev2', root=insightface_root_path, providers=devices.onnx) + self.app_640.prepare(ctx_id=0, det_size=(640, 640)) + self.app_320 = FaceAnalysis(name='antelopev2', root=insightface_root_path, providers=devices.onnx) + self.app_320.prepare(ctx_id=0, det_size=(320, 320)) + self.app_160 = FaceAnalysis(name='antelopev2', root=insightface_root_path, providers=devices.onnx) + self.app_160.prepare(ctx_id=0, det_size=(160, 160)) + self.arcface_model = init_recognition_model('arcface', device=devices.device) + + def load_loras(self, loras): + names, scales = [],[] + for lora_path, lora_name, lora_scale in loras: + if lora_path != "": + print(f"loading lora {lora_path}") + self.pipe.load_lora_weights(lora_path, adapter_name = lora_name) + names.append(lora_name) + scales.append(lora_scale) + + if len(names) > 0: + self.pipe.set_adapters(names, adapter_weights=scales) + + def _detect_face(self, id_image_cv2): + face_info = self.app_640.get(id_image_cv2) + if len(face_info) > 0: + return face_info + + face_info = self.app_320.get(id_image_cv2) + if len(face_info) > 0: + return face_info + + face_info = self.app_160.get(id_image_cv2) + return face_info + + def __call__( + self, + prompt: str, + id_image: Image.Image, # PIL.Image.Image (RGB) + negative_prompt = None, + control_image: Optional[Image.Image] = None, # PIL.Image.Image (RGB) or None + width = 1024, + height = 1024, + seed = 42, + guidance_scale = 3.5, + controlnet_guidance_scale = 1.0, + num_inference_steps = 30, + infusenet_conditioning_scale = 1.0, + infusenet_guidance_start = 0.0, + infusenet_guidance_end = 1.0, + output_type = 'pil', + generator = None, + *args, **kwargs # pylint: disable=unused-argument + ): + # Extract ID embeddings + id_image_cv2 = cv2.cvtColor(np.array(id_image), cv2.COLOR_RGB2BGR) + face_info = self._detect_face(id_image_cv2) + if len(face_info) == 0: + raise ValueError('No face detected in the input ID image') + + face_info = sorted(face_info, key=lambda x:(x['bbox'][2]-x['bbox'][0])*(x['bbox'][3]-x['bbox'][1]))[-1] # only use the maximum face + landmark = face_info['kps'] + id_embed = extract_arcface_bgr_embedding(id_image_cv2, landmark, self.arcface_model) + id_embed = id_embed.clone().unsqueeze(0).float().cuda() + id_embed = id_embed.reshape([1, -1, 512]) + id_embed = id_embed.to(device=devices.device, dtype=devices.dtype) + with torch.no_grad(): + id_embed = self.image_proj_model(id_embed) + bs_embed, seq_len, _ = id_embed.shape + id_embed = id_embed.repeat(1, 1, 1) + id_embed = id_embed.view(bs_embed * 1, seq_len, -1) + id_embed = id_embed.to(device=devices.device, dtype=devices.dtype) + + # Load control image + if control_image is not None: + control_image = control_image.convert("RGB") + control_image = resize_and_pad_image(control_image, (width, height)) + face_info = self._detect_face(cv2.cvtColor(np.array(control_image), cv2.COLOR_RGB2BGR)) + if len(face_info) == 0: + raise ValueError('No face detected in the control image') + face_info = sorted(face_info, key=lambda x:(x['bbox'][2]-x['bbox'][0])*(x['bbox'][3]-x['bbox'][1]))[-1] # only use the maximum face + control_image = draw_kps(control_image, face_info['kps']) + else: + out_img = np.zeros([height, width, 3]) + control_image = Image.fromarray(out_img.astype(np.uint8)) + + """ + control_image = self.pipe.prepare_image( + image=control_image, + width=width, + height=height, + batch_size=1, + num_images_per_prompt=1, + device=devices.device, + dtype=devices.dtype, + ) + control_image = retrieve_latents(self.pipe.vae.encode(control_image), generator=generator) + control_image = (control_image - self.pipe.vae.config.shift_factor) * self.pipe.vae.config.scaling_factor + # pack + height_control_image, width_control_image = control_image.shape[2:] + num_channels_latents = self.pipe.transformer.config.in_channels // 4 + control_image = self.pipe._pack_latents( + control_image, + 1, + num_channels_latents, + height_control_image, + width_control_image, + ) + """ + + # Perform inference + seed_everything(seed) + latents = self.pipe( + prompt=prompt, + negative_prompt=negative_prompt, + controlnet_prompt_embeds=id_embed, + control_image=control_image, + guidance_scale=guidance_scale, + num_inference_steps=num_inference_steps, + controlnet_guidance_scale=controlnet_guidance_scale, + controlnet_conditioning_scale=infusenet_conditioning_scale, + control_guidance_start=infusenet_guidance_start, + control_guidance_end=infusenet_guidance_end, + height=height, + width=width, + output_type=output_type, + callback_on_step_end=kwargs.get('callback_on_step_end', None), + callback_on_step_end_tensor_inputs=kwargs.get('callback_on_step_end_tensor_inputs', None), + ) + + return latents diff --git a/modules/infiniteyou/resampler.py b/modules/infiniteyou/resampler.py new file mode 100644 index 000000000..6d0011e83 --- /dev/null +++ b/modules/infiniteyou/resampler.py @@ -0,0 +1,121 @@ +# Modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py + +import math + +import torch +import torch.nn as nn + + +# FFN +def FeedForward(dim, mult=4): + inner_dim = int(dim * mult) + return nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, inner_dim, bias=False), + nn.GELU(), + nn.Linear(inner_dim, dim, bias=False), + ) + + +def reshape_tensor(x, heads): + bs, length, width = x.shape + #(bs, length, width) --> (bs, length, n_heads, dim_per_head) + x = x.view(bs, length, heads, -1) + # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) + x = x.transpose(1, 2) + # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) + x = x.reshape(bs, heads, length, -1) + return x + + +class PerceiverAttention(nn.Module): + def __init__(self, *, dim, dim_head=64, heads=8): + super().__init__() + self.scale = dim_head**-0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, l, _ = latents.shape + + q = self.to_q(latents) + kv_input = torch.cat((x, latents), dim=-2) + k, v = self.to_kv(kv_input).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, l, -1) + + return self.to_out(out) + + +class Resampler(nn.Module): + def __init__( + self, + dim=1024, + depth=8, + dim_head=64, + heads=16, + num_queries=8, + embedding_dim=768, + output_dim=1024, + ff_mult=4, + ): + super().__init__() + + self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) + + self.proj_in = nn.Linear(embedding_dim, dim) + + self.proj_out = nn.Linear(dim, output_dim) + self.norm_out = nn.LayerNorm(output_dim) + + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append( + nn.ModuleList( + [ + PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), + FeedForward(dim=dim, mult=ff_mult), + ] + ) + ) + + def forward(self, x): + + latents = self.latents.repeat(x.size(0), 1, 1) + + x = self.proj_in(x) + + for attn, ff in self.layers: + latents = attn(x, latents) + latents + latents = ff(latents) + latents + + latents = self.proj_out(latents) + return self.norm_out(latents) diff --git a/modules/lora/networks.py b/modules/lora/networks.py index b22825639..b1af5d9b6 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -485,7 +485,7 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn weights_backup = getattr(self, "network_weights_backup", None) bias_backup = getattr(self, "network_bias_backup", None) if weights_backup is None and bias_backup is None: - return None, None + return t0 = time.time() if weights_backup is not None: diff --git a/modules/processing_args.py b/modules/processing_args.py index db8794214..a839c9995 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -119,16 +119,16 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t t0 = time.time() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) apply_circular(p.tiling, model) - if hasattr(model, "set_progress_bar_config"): - if disable_pbar: - model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba', disable=disable_pbar) - else: - model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba') args = {} has_vae = hasattr(model, 'vae') or (hasattr(model, 'pipe') and hasattr(model.pipe, 'vae')) if hasattr(model, 'pipe') and not hasattr(model, 'no_recurse'): # recurse model = model.pipe has_vae = has_vae or hasattr(model, 'vae') + if hasattr(model, "set_progress_bar_config"): + if disable_pbar: + model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba', disable=disable_pbar) + else: + model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba') signature = inspect.signature(type(model).__call__, follow_wrapped=True) possible = list(signature.parameters) diff --git a/scripts/infiniteyou_ext.py b/scripts/infiniteyou_ext.py new file mode 100644 index 000000000..59432cd40 --- /dev/null +++ b/scripts/infiniteyou_ext.py @@ -0,0 +1,121 @@ +# https://huggingface.co/ByteDance/InfiniteYou +# https://github.com/bytedance/InfiniteYou +# flux base model + 11.8gb controlnet module + 338mb image module + 428 insightface module + +import gradio as gr +from PIL import Image +from modules import scripts, processing, shared, sd_models, devices + + +prefix = 'InfiniteYou' +model_versions = ['aes_stage2', 'sim_stage1'] +orig_pipeline, orig_prompt_attention = None, None + + +def verify_insightface(): + from installer import installed, install, reload + if not installed('insightface', reload=False, quiet=True): + install('insightface==0.7.3', ignore=False) + install('albumentations==1.4.3', ignore=False, reinstall=True) + install('pydantic==1.10.21', ignore=False, reinstall=True, force=True) + reload('pydantic') + + +def load_infiniteyou(model: str): + from modules.infiniteyou import InfUFluxPipeline + shared.sd_model = InfUFluxPipeline( + pipe=shared.sd_model, + model_version=model, + ) + sd_models.copy_diffuser_options(shared.sd_model, orig_pipeline) + sd_models.set_diffuser_options(shared.sd_model) + + +class Script(scripts.Script): + def title(self): + return f'{prefix}: Flexible Photo Recrafting' + + def show(self, is_img2img): + return not is_img2img if shared.native else False + + # return signature is array of gradio components + def ui(self, _is_img2img): + with gr.Row(): + gr.HTML(f'  {prefix}: Flexible Photo Recrafting
') + with gr.Row(): + model = gr.Dropdown(label='IY model', choices=model_versions, value=model_versions[0]) + restore = gr.Checkbox(label='Restore pipeline on end', value=False) + with gr.Row(): + scale = gr.Slider(label='IY scale', value=1.0, minimum=0.0, maximum=2.0, step=0.05) + with gr.Row(): + start = gr.Slider(label='IY start', value=0.0, minimum=0.0, maximum=1.0, step=0.05) + end = gr.Slider(label='IY end', value=1.0, minimum=0.0, maximum=1.0, step=0.05) + with gr.Row(): + id_guidance = gr.Slider(label='Identity guidance', value=3.5, minimum=0.0, maximum=14.0, step=0.05) + with gr.Row(): + id_image = gr.Image(label='Identity image', type='pil') + with gr.Row(): + control_guidance = gr.Slider(label='Control guidance', value=1.0, minimum=0.0, maximum=14.0, step=0.05) + with gr.Row(): + control_image = gr.Image(label='Control image', type='pil') + return [model, id_image, control_image, scale, start, end, id_guidance, control_guidance, restore] + + def run(self, p: processing.StableDiffusionProcessing, + model: str = None, + id_image: Image.Image = None, + control_image: Image.Image = None, + scale: float = 1.0, + start: float = 0.0, + end: float = 1.0, + id_guidance: float = 3.5, + control_guidance: float = 1.0, + restore: bool = False, + ): # pylint: disable=arguments-differ, unused-argument + + if model is None or model not in model_versions: + return None + if id_image is None: + shared.log.error(f'{prefix}: no init_images') + return None + if shared.sd_model_type != 'f1': + shared.log.error(f'{prefix}: invalid model type: {shared.sd_model_type}') + return None + + global orig_pipeline, orig_prompt_attention # pylint: disable=global-statement + orig_pipeline = shared.sd_model + if shared.sd_model.__class__.__name__ != 'InfUFluxPipeline': + verify_insightface() + load_infiniteyou(model) + devices.torch_gc() + shared.log.info(f'{prefix}: cls={shared.sd_model.__class__.__name__} loaded') + + processing.fix_seed(p) + p.task_args['id_image'] = id_image + p.task_args['control_image'] = control_image + p.task_args['infusenet_conditioning_scale'] = scale + p.task_args['infusenet_guidance_start'] = start + p.task_args['infusenet_guidance_end'] = end + p.task_args['seed'] = p.seed + p.task_args['negative_prompt'] = None + p.task_args['guidance_scale'] = id_guidance + p.task_args['controlnet_guidance_scale'] = control_guidance + p.extra_generation_params['IY model'] = model + p.extra_generation_params['IY guidance'] = f'{scale:.1f}/{start:.1f}/{end:.1f}' + orig_prompt_attention = shared.opts.prompt_attention + shared.opts.data['prompt_attention'] = 'fixed' + shared.log.debug(f'{prefix}: args={p.task_args}') + + processed = processing.process_images(p) + return processed + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, *args, **kwargs): # pylint: disable=unused-argument + # restore pipeline + global orig_pipeline, orig_prompt_attention # pylint: disable=global-statement + restore = args[-1] + if orig_prompt_attention is not None: + shared.opts.data['prompt_attention'] = orig_prompt_attention + orig_prompt_attention = None + if restore and orig_pipeline is not None: + shared.log.info(f'{prefix}: restoring pipeline') + shared.sd_model = orig_pipeline + orig_pipeline = None From 8bcc4527ea4cf1b555057af5975cb6ff0da36432 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 26 Mar 2025 11:10:23 -0400 Subject: [PATCH 065/122] add vlm ByteDance/Sa2VA 1b and 4b Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 14 ++++++++------ modules/interrogate/vqa.py | 39 +++++++++++++++++++++++++++++++++++++- 2 files changed, 46 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 072af4c91..6314aa6af 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -72,10 +72,10 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - new [Caption](https://github.com/vladmandic/sdnext/wiki/Caption) guide - new [VAE](https://github.com/vladmandic/sdnext/wiki/VAE) guide - updated [SD3](https://github.com/vladmandic/sdnext/wiki/SD3) guide - - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide - - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide - - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide - - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide + - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide + - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide + - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide + - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide - **Remote VAE** - add support for remote vae encode in addition to remote vae decode - used by *img2img, inpaint, hires, detailer* @@ -83,9 +83,11 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - add remote vae info to metadata, thanks @iDeNoh - remote vae use `scaling_factor` and `shift_factor` - **Caption/VLM** - - [Google Gemma 3 4B](https://huggingface.co/google/gemma-3-4b-it) + - [Google Gemma 3](https://huggingface.co/google/gemma-3-4b-it) 4B simply select from list of available models in caption tab - - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* + - [ByteDance/Sa2VA](https://huggingface.co/ByteDance/Sa2VA-1B) 1B, 4B + simply select from list of available models in caption tab + - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* - [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) extension updates - add detection of prompt language and alphabet and filter based on those values - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index 2bdec3543..3db5251e8 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -41,6 +41,8 @@ vlm_models = { "AIDC Ovis2 1B": "AIDC-AI/Ovis2-1B", "AIDC Ovis2 2B": "AIDC-AI/Ovis2-2B", "AIDC Ovis2 4B": "AIDC-AI/Ovis2-4B", + "ByteDance Sa2VA 1B": "ByteDance/Sa2VA-1B", + "ByteDance Sa2VA 4B": "ByteDance/Sa2VA-4B", # "OpenGVLab InternVL 2.5 1B": "OpenGVLab/InternVL2_5-1B" # "DeepSeek VL2 Tiny": "deepseek-ai/deepseek-vl2-tiny", # broken # "nVidia Eagle 2 1B": "nvidia/Eagle2-1B", # not compatible with latest transformers @@ -72,7 +74,7 @@ def b64(image): def clean(response, question): - strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:'] + strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:', '<|im_end|>'] if isinstance(response, dict): if 'task' in response: response = response['task'] @@ -451,6 +453,39 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str return response +def sa2(question: str, image: Image.Image, repo: str = None): + global processor, model, loaded # pylint: disable=global-statement + if model is None or loaded != repo: + model = transformers.AutoModel.from_pretrained( + repo, + torch_dtype=devices.dtype, + low_cpu_mem_usage=True, + use_flash_attn=False, + trust_remote_code=True) + model = model.eval() + processor = transformers.AutoTokenizer.from_pretrained( + repo, + trust_remote_code=True, + use_fast=False, + ) + loaded = repo + model = model.to(devices.device, devices.dtype) + if question.startswith('<'): + task = question.split('>', 1)[0] + '>' + else: + task = '' + input_dict = { + 'image': image, + 'text': f'{task}', + 'past_text': '', + 'mask_prompts': None, + 'tokenizer': processor, + } + return_dict = model.predict_forward(**input_dict) + response = return_dict["prediction"] # the text format answer + return response + + def interrogate(question, system_prompt, prompt, image, model_name, quiet:bool=False): if not quiet: shared.state.begin('Interrogate') @@ -516,6 +551,8 @@ def interrogate(question, system_prompt, prompt, image, model_name, quiet:bool=F answer = gemma(question, image, vqa_model, system_prompt) elif 'ovis' in vqa_model.lower(): answer = ovis(question, image, vqa_model) + elif 'sa2' in vqa_model.lower(): + answer = sa2(question, image, vqa_model) else: answer = 'unknown model' except Exception as e: From 068daa1e0915e01141a4b6699f4218f014e5b6c3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 27 Mar 2025 09:13:10 -0400 Subject: [PATCH 066/122] add manual module move for video models Signed-off-by: Vladimir Mandic --- modules/video_models/video_utils.py | 3 ++- modules/video_models/video_vae.py | 1 + 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index bb971d9d5..2b996a7e2 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -1,6 +1,6 @@ import os import time -from modules import shared, sd_models, timer, errors +from modules import shared, sd_models, timer, errors, devices debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -31,6 +31,7 @@ def set_prompt(p): def hijack_encode_prompt(*args, **kwargs): t0 = time.time() try: + sd_models.move_model(shared.sd_model.text_encoder, devices.device) res = shared.sd_model.orig_encode_prompt(*args, **kwargs) except Exception as e: shared.log.error(f'Video encode: {e}') diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py index 8adbc939e..5a107f616 100644 --- a/modules/video_models/video_vae.py +++ b/modules/video_models/video_vae.py @@ -57,6 +57,7 @@ def hijack_vae_decode(*args, **kwargs): if res is None: shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) try: + sd_models.move_model(shared.sd_model.vae, devices.device) if torch.is_tensor(args[0]): latent = args[0] latent = latent.to(device=devices.device, dtype=shared.sd_model.vae.dtype) # upcast to vae dtype From 2c58d3b36c1783439d0ca713d69922dfc4f15ba3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 27 Mar 2025 11:49:30 -0400 Subject: [PATCH 067/122] fastercache and pyramidattentionbroadcast Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 41 ++++++++++--------- TODO.md | 3 +- javascript/sdnext.css | 1 + modules/processing_diffusers.py | 3 +- modules/sd_offload.py | 15 ------- modules/shared.py | 20 ++++++++-- modules/transformer_cache.py | 52 +++++++++++++++++++++++++ modules/ui_video.py | 6 +-- modules/video_models/models_def.py | 1 + modules/video_models/video_cache.py | 45 --------------------- modules/video_models/video_load.py | 9 ++++- modules/video_models/video_overrides.py | 14 +++---- modules/video_models/video_run.py | 6 +-- modules/video_models/video_utils.py | 22 +++++++++-- modules/video_models/video_vae.py | 27 ++++++++++++- 15 files changed, 156 insertions(+), 109 deletions(-) create mode 100644 modules/transformer_cache.py delete mode 100644 modules/video_models/video_cache.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 6314aa6af..b17fbee00 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,18 +2,17 @@ ## Update for 2025-03-26 -### Highlights for 2025-03-26 +### Highlights for 2025-03-27 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to remote VAE, additional docs/guides +Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods -**Flash Attention 2** and Sage Attention is now available on ZLUDA backend! +### Details for 2025-03-27 -### Details for 2025-03-26 - -- **Video tab** +- **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! - new top-level tab, replaces previous *video* script in text/image tabs old scripts are still present, but will be removed in the future @@ -24,15 +23,12 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - [CogVideoX](https://huggingface.co/THUDM/CogVideoX-5b): *2B, 5B* | *T2V, I2V* - [Allegro](https://huggingface.co/rhymes-ai/Allegro): *T2V* - [Mochi1](https://huggingface.co/genmo/mochi-1-preview): *T2V* - - [Latte1](https://huggingface.co/maxin-cn/Latte-1): *T2V + - [Latte1](https://huggingface.co/maxin-cn/Latte-1): *T2V - decoding: - **Default**: use vae from model - **Tiny VAE**: support for *Hunyuan, WAN, Mochi* - **Remote VAE**: support for *Hunyuan* - **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog* - - acceleration: - - [FasterCache](https://huggingface.co/papers/2410.19355): support for *Hunyuan, Mochi, Latte, Allegro, Cog* - - [PyramidAttentionBroadcast](https://huggingface.co/papers/2408.12588): support for *Hunyuan, Mochi, Latte, Allegro, Cog* - additional key points: - all models are auto-downloaded upon first use uses *system paths -> huggingface* folder @@ -43,10 +39,10 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - support for balanced offloading and model offloading uses system settings - on-the-fly quantization: *BnB, Quanto, TorchAO* - uses system settings, granular for *transformer* and *text-encoder* separately + uses system settings, granular for *transformer* and *text-encoder* separately - different video models support different video resolutions, frame counts, etc. and may require specific settings - see model links for details - - see *ToDo/Limitations* section for additional notes + - see *ToDo/Limitations* section for additional notes - **Models & Pipelines** - [THUDM CogView 4](https://huggingface.co/THUDM/CogView4-6B) **6B** variant new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer @@ -76,6 +72,11 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide +- **Acceleration** + - Support for most DiT-based models, for example: *FLUX.1, SD35, Hunyuan, Mochi, Latte, Allegro, Cog* + - Enable and configure in *Settings -> Pipeline modifiers* + - [FasterCache](https://huggingface.co/papers/2410.19355) + - [PyramidAttentionBroadcast](https://huggingface.co/papers/2408.12588) - **Remote VAE** - add support for remote vae encode in addition to remote vae decode - used by *img2img, inpaint, hires, detailer* @@ -93,16 +94,17 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) against top-10 standard harmful content categories - add banned words/expressions check against prompt variations -- **Other** +- **Other** - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method - **upscale**: new experimental support for `libvips` upscaling - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models note: previous method for quanto is still valid and is noted in settings as post-load quantization - - add quantization support to **CogView-3Plus** - - update `diffusers` and other requirements + - add quantization support to **CogView-3Plus** + - update `diffusers` and other requirements - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis -- **IPEX** + - LoRA enable memory cache by default +- **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` - add xpu to profiler - fix untyped_storage, torch.eye and torch.cuda.device ops @@ -112,10 +114,11 @@ Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to r - **ROCm** - add `--upgrade` to torch_command when using `--use-nightly` - disable fp16 for gfx1102 (rx 7600 and rx 7500 series) gpus -- **ZLUDA** - - add `torch.compile` support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) -- **Fixes** +- **ZLUDA** + - add `torch.compile` support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) + - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) + - add Sage Attention support +- **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled - fix cuda errors with *directml* diff --git a/TODO.md b/TODO.md index d9debcb5d..c11b65b7d 100644 --- a/TODO.md +++ b/TODO.md @@ -8,7 +8,8 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: Hunyuan Video I2V: requires `transformers==4.47.1` - Video: Latte 1 T2V: dtype mismatch -- Video: WAN 2.1 14B I2V 480p/720p: broken offload +- Video: WAN 2.1 14B I2V 480p/720p: broken offload +- Video: WAN 2.1 14B I2V 480p/720p: custom number of frames - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output - Video: Allegro T2V: all-gray output diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 872f2b12f..3e841346d 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -114,6 +114,7 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt, #control_prompt, #control_neg_prompt, #video_prompt, #video_neg_prompt { display: contents; } #txt2img_actions_column, #img2img_actions_column, #control_actions, #video_actions { flex-flow: wrap; justify-content: space-between; } #txt2img_seed, #img2img_seed, #control_seed, #video_seed { min-width: 90px !important } +#video_generate_box>button { max-width: unset; } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } .interrogate:hover { background: var(--button-primary-background-fill-hover) !important; } diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 684484217..de21cf8b9 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -5,7 +5,7 @@ import numpy as np import torch import torchvision.transforms.functional as TF from PIL import Image -from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, hidiffusion, timer, modelstats, extra_networks, ras +from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, hidiffusion, timer, modelstats, extra_networks, ras, transformer_cache from modules.processing_helpers import resize_hires, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, save_intermediate, update_sampler, is_txt2img, is_refiner_enabled, get_job_name from modules.processing_args import set_pipeline_args from modules.onnx_impl import preprocess_pipeline as preprocess_onnx_pipeline, check_parameters_changed as olive_check_parameters_changed @@ -87,6 +87,7 @@ def process_base(p: processing.StableDiffusionProcessing): try: t0 = time.time() sd_models_compile.check_deepcache(enable=True) + transformer_cache.set_cache() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) sd_models.move_model(shared.sd_model, devices.device) if hasattr(shared.sd_model, 'unet'): diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 8ab114e93..2e02f98a9 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -3,7 +3,6 @@ import sys import time import inspect import torch -import diffusers import accelerate.hooks from modules import shared, devices, errors, model_quant from modules.timer import process as process_timer @@ -285,22 +284,8 @@ def apply_balanced_offload(sd_model=None, exclude=[]): apply_balanced_offload_to_module(sd_model.prior_pipe) if hasattr(sd_model, "decoder_pipe"): apply_balanced_offload_to_module(sd_model.decoder_pipe) - if shared.opts.layerwise_quantization: model_quant.apply_layerwise(sd_model, quiet=True) # need to reapply since hooks were removed/readded - if shared.opts.pab_enabled and hasattr(sd_model, 'transformer'): - pab_config = diffusers.PyramidAttentionBroadcastConfig( - spatial_attention_block_skip_range=shared.opts.pab_block_skip_range, - spatial_attention_timestep_skip_range=(int(100 * shared.opts.pab_timestep_skip_start), int(100 * shared.opts.pab_timestep_skip_end)), - current_timestep_callback=lambda: sd_model.current_timestep, # pylint: disable=protected-access - ) - try: - diffusers.apply_pyramid_attention_broadcast(sd_model.transformer, pab_config) - except Exception: # hook may already exist - pass - if not cached: - shared.log.info(f'Applying PAB: cls={sd_model.transformer.__class__.__name__} block={shared.opts.pab_block_skip_range} start={shared.opts.pab_timestep_skip_start} end={shared.opts.pab_timestep_skip_end}') - set_accelerate(sd_model) t = time.time() - t0 process_timer.add('offload', t) diff --git a/modules/shared.py b/modules/shared.py index a4a448966..045d87dcf 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -561,9 +561,21 @@ options_templates.update(options_section(('advanced', "Pipeline Modifiers"), { "pab_sep": OptionInfo("

PAB: Pyramid attention broadcast

", "", gr.HTML), "pab_enabled": OptionInfo(False, "Attention cache enabled"), - "pab_block_skip_range": OptionInfo(2, "Block skip range", gr.Slider, {"minimum": 1, "maximum": 4, "step": 1}), - "pab_timestep_skip_start": OptionInfo(0.1, "Timestep skip start", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.05}), - "pab_timestep_skip_end": OptionInfo(0.8, "Timestep skip end", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.05}), + "pab_spacial_skip_range": OptionInfo(2, "FC spacial skip range", gr.Slider, {"minimum": 1, "maximum": 4, "step": 1}), + "pab_spacial_skip_start": OptionInfo(100, "FC spacial skip start", gr.Slider, {"minimum": 0, "maximum": 1000, "step": 1}), + "pab_spacial_skip_end": OptionInfo(800, "FC spacial skip end", gr.Slider, {"minimum": 0, "maximum": 1000, "step": 1}), + + "faster_cache__sep": OptionInfo("

Faster Cache

", "", gr.HTML), + "faster_cache_enabled": OptionInfo(False, "Faster cache enabled"), + "fc_spacial_skip_range": OptionInfo(2, "FC spacial skip range", gr.Slider, {"minimum": 1, "maximum": 4, "step": 1}), + "fc_spacial_skip_start": OptionInfo(0, "FC spacial skip start", gr.Slider, {"minimum": 0, "maximum": 1000, "step": 1}), + "fc_spacial_skip_end": OptionInfo(681, "FC spacial skip end", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01}), + "fc_uncond_skip_range": OptionInfo(5, "FC uncond skip range", gr.Slider, {"minimum": 1, "maximum": 4, "step": 1}), + "fc_uncond_skip_start": OptionInfo(0, "FC uncond skip start", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01}), + "fc_uncond_skip_end": OptionInfo(781, "FC uncond skip end", gr.Slider, {"minimum": 0, "maximum": 1, "step": 1}), + "fc_attention_weight": OptionInfo(0.5, "FC spacial skip range", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.05}), + "fc_tensor_format": OptionInfo("BCFHW", "FC tensor format", gr.Radio, {"choices": ["BCFHW", "BFCHW", "BCHW"]}), + "fc_guidance_distilled": OptionInfo(False, "FC guidance distilled", gr.Checkbox), "para_sep": OptionInfo("

Para-attention

", "", gr.HTML), "para_cache_enabled": OptionInfo(False, "First-block cache enabled"), @@ -915,7 +927,7 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "lora_force_diffusers": OptionInfo(False if not cmd_opts.use_openvino else True, "LoRA load using Diffusers method"), "lora_maybe_diffusers": OptionInfo(False, "LoRA load using Diffusers method for selected models"), "lora_apply_tags": OptionInfo(0, "LoRA auto-apply tags", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}), - "lora_in_memory_limit": OptionInfo(0, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 24, "step": 1}), + "lora_in_memory_limit": OptionInfo(1, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1}), "lora_quant": OptionInfo("NF4","LoRA precision when quantized", gr.Radio, {"choices": ["NF4", "FP4"]}), "extra_networks_styles_sep": OptionInfo("

Styles

", "", gr.HTML), diff --git a/modules/transformer_cache.py b/modules/transformer_cache.py new file mode 100644 index 000000000..3e69126c8 --- /dev/null +++ b/modules/transformer_cache.py @@ -0,0 +1,52 @@ +import os +import diffusers +from modules import shared, errors + + +debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None + + +def set_cache(faster_cache=None, pyramid_attention_broadcast=None): + if not shared.sd_loaded or not hasattr(shared.sd_model, 'transformer'): + return + faster_cache = faster_cache if faster_cache is not None else shared.opts.faster_cache_enabled + pyramid_attention_broadcast = pyramid_attention_broadcast if pyramid_attention_broadcast is not None else shared.opts.pab_enabled + if not faster_cache and not pyramid_attention_broadcast: + return + if not hasattr(shared.sd_model.transformer, 'enable_cache') or not hasattr(shared.sd_model.transformer, 'disable_cache'): + shared.log.debug(f'Transformer cache: cls={shared.sd_model.transformer.__class__.__name__} not supported') + return + try: + if faster_cache: # https://github.com/huggingface/diffusers/pull/10163 + distilled = shared.opts.fc_guidance_distilled or shared.sd_model_type == 'f1' + config = diffusers.FasterCacheConfig( + spatial_attention_block_skip_range=shared.opts.fc_spacial_skip_range, + spatial_attention_timestep_skip_range=(int(shared.opts.fc_spacial_skip_start), int(shared.opts.fc_spacial_skip_end)), + unconditional_batch_skip_range=shared.opts.fc_uncond_skip_range, + unconditional_batch_timestep_skip_range=(int(shared.opts.fc_uncond_skip_start), int(shared.opts.fc_uncond_skip_end)), + attention_weight_callback=lambda _: shared.opts.fc_attention_weight, + tensor_format=shared.opts.fc_tensor_format, # TODO fc: autodetect tensor format based on model + is_guidance_distilled=distilled, # TODO fc: autodetect distilled based on model + current_timestep_callback=lambda: shared.sd_model.current_timestep, + ) + shared.sd_model.transformer.disable_cache() + shared.sd_model.transformer.enable_cache(config) + shared.log.debug(f'Transformer cache: type={config.__class__.__name__}') + shared.log.critical(f'HERE: {vars(config)}') + debug(f'Transformer cache: {vars(config)}') + elif pyramid_attention_broadcast: # https://github.com/huggingface/diffusers/pull/9562 + config = diffusers.PyramidAttentionBroadcastConfig( + spatial_attention_block_skip_range=shared.opts.pab_spacial_skip_range, + spatial_attention_timestep_skip_range=(int(shared.opts.pab_spacial_skip_start), int(shared.opts.pab_spacial_skip_end)), + current_timestep_callback=lambda: shared.sd_model.current_timestep, + ) + shared.sd_model.transformer.disable_cache() + shared.sd_model.transformer.enable_cache(config) + shared.log.debug(f'Transformer cache: type={config.__class__.__name__}') + debug(f'Transformer cache: {vars(config)}') + else: + debug('Transformer cache: not enabled') + shared.sd_model.transformer.disable_cache() + except Exception as e: + shared.log.error(f'Transformer cache: {e}') + errors.display(e, 'Transformer cache') diff --git a/modules/ui_video.py b/modules/ui_video.py index c0495defe..f94e46ff6 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -89,7 +89,7 @@ def create_ui(): url = gr.HTML(label='Model URL', elem_id='video_model_url', value='

') with gr.Accordion(open=True, label="Size", elem_id='video_size_accordion'): with gr.Row(): - width, height = ui_sections.create_resolution_inputs('video', default_width=720, default_height=480) + width, height = ui_sections.create_resolution_inputs('video', default_width=832, default_height=480) with gr.Row(): frames = gr.Slider(label='Frames', minimum=1, maximum=1024, step=1, value=15, elem_id="video_frames") seed = gr.Number(label='Initial seed', value=-1, elem_id="video_seed", container=True) @@ -111,9 +111,6 @@ def create_ui(): gr.HTML("
  Init image") init_image = gr.Image(elem_id="video_image", show_label=False, type="pil", image_mode="RGB", height=512) init_strength = gr.Slider(label='Init strength', minimum=0.0, maximum=1.0, step=0.01, value=0.5, elem_id="video_denoising_strength") - with gr.Accordion(open=False, label="Accelerate", elem_id='video_accelerate_accordion'): - faster_cache = gr.Checkbox(label='FasterCache', value=False, elem_id="video_faster_cache") - pyramid_attention = gr.Checkbox(label='PyramidAttention', value=False, elem_id="video_pyramid_attention") with gr.Accordion(open=True, label="Output", elem_id='video_output_accordion'): with gr.Row(): save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") @@ -168,7 +165,6 @@ def create_ui(): vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, - faster_cache, pyramid_attention, override_settings, ] # generate function diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index f7181b150..69287bb65 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -16,6 +16,7 @@ class Model(): te_cls: classmethod = None te_folder: str = 'text_encoder' te_hijack: bool = True + image_hijack: bool = True vae_hijack: bool = True vae_remote: bool = False diff --git a/modules/video_models/video_cache.py b/modules/video_models/video_cache.py deleted file mode 100644 index be43068d1..000000000 --- a/modules/video_models/video_cache.py +++ /dev/null @@ -1,45 +0,0 @@ -import os -import diffusers -from modules import shared, errors - - -debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None - - -def set_cache(faster_cache=False, pyramid_attention_broadcast=False): - if not shared.sd_loaded or not hasattr(shared.sd_model, 'transformer'): - return - if not hasattr(shared.sd_model.transformer, 'enable_cache'): - debug(f'Video cache: cls={shared.sd_model.transformer.__class__.__name__} not supported') - return - try: - if faster_cache: # https://github.com/huggingface/diffusers/pull/10163 - config = diffusers.FasterCacheConfig( - spatial_attention_block_skip_range=2, - spatial_attention_timestep_skip_range=(-1, 681), - current_timestep_callback=lambda: shared.sd_model.current_timestep, - attention_weight_callback=lambda _: 0.3, - unconditional_batch_skip_range=5, - unconditional_batch_timestep_skip_range=(-1, 781), - tensor_format="BFCHW", - ) - shared.sd_model.transformer.disable_cache() - shared.sd_model.transformer.enable_cache(config) - shared.log.debug(f'Video cache: type={config.__class__.__name__}') - debug(f'Video cache: {vars(config)}') - elif pyramid_attention_broadcast: # https://github.com/huggingface/diffusers/pull/9562 - config = diffusers.PyramidAttentionBroadcastConfig( - spatial_attention_block_skip_range=2, - spatial_attention_timestep_skip_range=(100, 800), - current_timestep_callback=lambda: shared.sd_model.current_timestep, - ) - shared.sd_model.transformer.disable_cache() - shared.sd_model.transformer.enable_cache(config) - shared.log.debug(f'Video cache: type={config.__class__.__name__}') - debug(f'Video cache: {vars(config)}') - else: - debug('Video cache: not enabled') - shared.sd_model.transformer.disable_cache() - except Exception as e: - shared.log.error(f'Video cache: error={e}') - errors.display(e, 'video cache') diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 9254c4ed6..3bfeb2786 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -71,12 +71,17 @@ def load_model(selected: models_def.Model): shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo) shared.sd_model.sd_model_hash = None sd_models.set_diffuser_options(shared.sd_model) - if selected.vae_hijack: + if selected.vae_hijack and hasattr(shared.sd_model.vae, 'decode'): shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode shared.sd_model.vae.decode = video_vae.hijack_vae_decode - if selected.te_hijack: + shared.sd_model.vae.orig_encode = shared.sd_model.vae.encode + shared.sd_model.vae.encode = video_vae.hijack_vae_encode + if selected.te_hijack and hasattr(shared.sd_model, 'encode_prompt'): shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.encode_prompt = video_utils.hijack_encode_prompt + if selected.image_hijack and hasattr(shared.sd_model, 'encode_image'): + shared.sd_model.orig_encode_image = shared.sd_model.encode_image + shared.sd_model.encode_image = video_utils.hijack_encode_image if hasattr(shared.sd_model.vae, 'enable_slicing'): shared.sd_model.vae.enable_slicing() loaded_model = selected.name diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index 7af00f6be..5436b485a 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -11,16 +11,12 @@ debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None e def load_override(selected: Model): kwargs = {} if selected.name == 'Allegro T2V': - kwargs['vae'] = diffusers.AutoencoderKLAllegro.from_pretrained(selected.repo, - subfolder="vae", - torch_dtype=torch.float32, - cache_dir=shared.opts.hfcache_dir) - debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}') + kwargs['vae'] = diffusers.AutoencoderKLAllegro.from_pretrained(selected.repo, subfolder="vae", torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir) if selected.name == 'LTXVideo 0.9.5 I2V': - kwargs['vae'] = diffusers.AutoencoderKLLTXVideo.from_pretrained(selected.repo, - subfolder="vae", - torch_dtype=torch.float32, - cache_dir=shared.opts.hfcache_dir) + kwargs['vae'] = diffusers.AutoencoderKLLTXVideo.from_pretrained(selected.repo, subfolder="vae", torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir) + if selected.name == 'WAN 2.1 14B I2V 480p' or selected.name == 'WAN 2.1 14B I2V 720p': + kwargs['vae'] = diffusers.AutoencoderKLWan.from_pretrained(selected.repo, subfolder="vae", torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir) + debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}') return kwargs diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 5f60efd3c..1c939936a 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -1,14 +1,14 @@ import os import time from modules import shared, errors, sd_models, processing, devices, images, ui_common -from modules.video_models import models_def, video_utils, video_load, video_vae, video_cache, video_overrides +from modules.video_models import models_def, video_utils, video_load, video_vae, video_overrides debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None def generate(*args, **kwargs): - task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, init_strength, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, faster_cache, pyramid_attention, override_settings = args + task_id, ui_state, engine, model, prompt, negative, styles, width, height, frames, steps, sampler_index, sampler_shift, dynamic_shift, seed, guidance_scale, guidance_true, init_image, init_strength, vae_type, vae_tile_frames, save_frames, video_type, video_duration, video_loop, video_pad, video_interpolate, override_settings = args if engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') found = [model.name for model in models_def.models.get(engine, [])] @@ -64,7 +64,6 @@ def generate(*args, **kwargs): # set args processing.fix_seed(p) video_vae.set_vae_params(p) - video_cache.set_cache(faster_cache=faster_cache, pyramid_attention_broadcast=pyramid_attention) video_utils.set_prompt(p) p.task_args['num_inference_steps'] = p.steps p.task_args['width'] = p.width @@ -83,6 +82,7 @@ def generate(*args, **kwargs): shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={p.frames} steps={p.steps}') err = None t0 = time.time() + processed = None try: processed = processing.process_images(p) except Exception as e: diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 2b996a7e2..4e26849b6 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -34,11 +34,27 @@ def hijack_encode_prompt(*args, **kwargs): sd_models.move_model(shared.sd_model.text_encoder, devices.device) res = shared.sd_model.orig_encode_prompt(*args, **kwargs) except Exception as e: - shared.log.error(f'Video encode: {e}') - errors.display(e, 'Video encode') + shared.log.error(f'Video encode prompt: {e}') + errors.display(e, 'Video encode prompt') res = None t1 = time.time() timer.process.add('te', t1-t0) - debug(f'Video encode: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + debug(f'Video encode prompt: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + return res + + +def hijack_encode_image(*args, **kwargs): + t0 = time.time() + try: + sd_models.move_model(shared.sd_model.image_encoder, devices.device) + res = shared.sd_model.orig_encode_image(*args, **kwargs) + except Exception as e: + shared.log.error(f'Video encode image: {e}') + errors.display(e, 'Video encode image') + res = None + t1 = time.time() + timer.process.add('te', t1-t0) + debug(f'Video encode image: te={shared.sd_model.image_encoder.__class__.__name__} time={t1-t0:.2f}') shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) return res diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py index 5a107f616..0dab30c72 100644 --- a/modules/video_models/video_vae.py +++ b/modules/video_models/video_vae.py @@ -65,10 +65,33 @@ def hijack_vae_decode(*args, **kwargs): else: res = shared.sd_model.vae.orig_decode(*args, **kwargs) except Exception as e: - shared.log.error(f'Video VAE: type={vae_type} {e}') + shared.log.error(f'Video VAE decode: type={vae_type} {e}') errors.display(e, 'Video VAE') res = None t1 = time.time() timer.process.add('vae', t1-t0) - debug(f'Video decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}') + debug(f'Video VAE decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}') + return res + + +def hijack_vae_encode(*args, **kwargs): + t0 = time.time() + res = None + if res is None: + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) + try: + sd_models.move_model(shared.sd_model.vae, devices.device) + if torch.is_tensor(args[0]): + latent = args[0] + latent = latent.to(device=devices.device, dtype=shared.sd_model.vae.dtype) # upcast to vae dtype + res = shared.sd_model.vae.orig_encode(latent, *args[1:], **kwargs) + else: + res = shared.sd_model.vae.orig_encode(*args, **kwargs) + except Exception as e: + shared.log.error(f'Video VAE encode: type={vae_type} {e}') + errors.display(e, 'Video VAE') + res = None + t1 = time.time() + timer.process.add('vae', t1-t0) + debug(f'Video VAE encode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}') return res From 2a18890235c3178c6f0dbee16d0ecfb72b7d06d7 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 27 Mar 2025 11:53:03 -0400 Subject: [PATCH 068/122] update changelog and cleanup Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 40 ++++++++++++++++++------------------ modules/transformer_cache.py | 7 +++---- 2 files changed, 23 insertions(+), 24 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index b17fbee00..dd75c5009 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -62,16 +62,6 @@ Pretty big performance updates to a) Any model using DiT based architecture: new download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui -- **Wiki/Docs** - - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide - - new [Caption](https://github.com/vladmandic/sdnext/wiki/Caption) guide - - new [VAE](https://github.com/vladmandic/sdnext/wiki/VAE) guide - - updated [SD3](https://github.com/vladmandic/sdnext/wiki/SD3) guide - - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide - - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide - - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide - - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide - **Acceleration** - Support for most DiT-based models, for example: *FLUX.1, SD35, Hunyuan, Mochi, Latte, Allegro, Cog* - Enable and configure in *Settings -> Pipeline modifiers* @@ -94,16 +84,6 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) against top-10 standard harmful content categories - add banned words/expressions check against prompt variations -- **Other** - - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method - - **upscale**: new experimental support for `libvips` upscaling - - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models - note: previous method for quanto is still valid and is noted in settings as post-load quantization - - add quantization support to **CogView-3Plus** - - update `diffusers` and other requirements - - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis - - LoRA enable memory cache by default - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` - add xpu to profiler @@ -118,6 +98,26 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - add `torch.compile` support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - add Sage Attention support +- **Other** + - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method + - **upscale**: new experimental support for `libvips` upscaling + - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models + note: previous method for quanto is still valid and is noted in settings as post-load quantization + - add quantization support to **CogView-3Plus** + - update `diffusers` and other requirements + - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion + - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis + - LoRA enable memory cache by default +- **Wiki/Docs** + - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info + - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide + - new [Caption](https://github.com/vladmandic/sdnext/wiki/Caption) guide + - new [VAE](https://github.com/vladmandic/sdnext/wiki/VAE) guide + - updated [SD3](https://github.com/vladmandic/sdnext/wiki/SD3) guide + - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide + - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide + - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide + - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/modules/transformer_cache.py b/modules/transformer_cache.py index 3e69126c8..c8615c534 100644 --- a/modules/transformer_cache.py +++ b/modules/transformer_cache.py @@ -11,10 +11,10 @@ def set_cache(faster_cache=None, pyramid_attention_broadcast=None): return faster_cache = faster_cache if faster_cache is not None else shared.opts.faster_cache_enabled pyramid_attention_broadcast = pyramid_attention_broadcast if pyramid_attention_broadcast is not None else shared.opts.pab_enabled - if not faster_cache and not pyramid_attention_broadcast: + if (not faster_cache) and (not pyramid_attention_broadcast): return - if not hasattr(shared.sd_model.transformer, 'enable_cache') or not hasattr(shared.sd_model.transformer, 'disable_cache'): - shared.log.debug(f'Transformer cache: cls={shared.sd_model.transformer.__class__.__name__} not supported') + if (not hasattr(shared.sd_model.transformer, 'enable_cache')) or (not hasattr(shared.sd_model.transformer, 'disable_cache')): + shared.log.debug(f'Transformer cache: cls={shared.sd_model.transformer.__class__.__name__} fc={faster_cache} pab={pyramid_attention_broadcast} not supported') return try: if faster_cache: # https://github.com/huggingface/diffusers/pull/10163 @@ -32,7 +32,6 @@ def set_cache(faster_cache=None, pyramid_attention_broadcast=None): shared.sd_model.transformer.disable_cache() shared.sd_model.transformer.enable_cache(config) shared.log.debug(f'Transformer cache: type={config.__class__.__name__}') - shared.log.critical(f'HERE: {vars(config)}') debug(f'Transformer cache: {vars(config)}') elif pyramid_attention_broadcast: # https://github.com/huggingface/diffusers/pull/9562 config = diffusers.PyramidAttentionBroadcastConfig( From ecb6730838abf2d4b2cc52746433a65cd76f0f8a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 27 Mar 2025 13:22:09 -0400 Subject: [PATCH 069/122] fix sampler metadata when using default sampler Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + TODO.md | 12 +++++------- modules/processing_info.py | 2 +- 3 files changed, 7 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index dd75c5009..c73845f29 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -135,6 +135,7 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - fix sd35 with batch processing - fix extra networks cover and inline views - fix token counter error style with modernui + - fix sampler metadata when using default sampler - improve lora compatibility with balanced offload ## Update for 2025-02-28 diff --git a/TODO.md b/TODO.md index c11b65b7d..1bf6102a6 100644 --- a/TODO.md +++ b/TODO.md @@ -18,13 +18,11 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Flux: NF4 loader: - IPAdapter: negative guidance: - Control: API enhance scripts compatibility -- Video: add generate context menu -- Video: FasterCache and PyramidAttentionBroadcast granular config -- Video: FasterCache and PyramidAttentionBroadcast for LTX and WAN -- Video: API support -- Video: STG: -- Video SmoothCache: https://github.com/huggingface/diffusers/issues/11135 -- FasterCache, PyramidAttentionBroadcast, SmoothCache general support +- Video: add generate context menu +- Video: API support +- Video: STG: +- Video SmoothCache: https://github.com/huggingface/diffusers/issues/11135 +- SoftFill: https://github.com/zacheryvaughn/softfill-pipelines ## Code TODO diff --git a/modules/processing_info.py b/modules/processing_info.py index 4db8ccaa5..d254168cc 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -159,7 +159,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No args["Embeddings"] = ', '.join(sd_hijack.model_hijack.embedding_db.embeddings_used) # samplers - if getattr(p, 'sampler_name', None) is not None: + if getattr(p, 'sampler_name', None) is not None and p.sampler_name.lower() != 'default': args["Sampler eta delta"] = shared.opts.eta_noise_seed_delta if shared.opts.eta_noise_seed_delta != 0 and sd_samplers_common.is_sampler_using_eta_noise_seed_delta(p) else None args["Sampler eta multiplier"] = p.initial_noise_multiplier if getattr(p, 'initial_noise_multiplier', 1.0) != 1.0 else None args['Sampler timesteps'] = shared.opts.schedulers_timesteps if shared.opts.schedulers_timesteps != shared.opts.data_labels.get('schedulers_timesteps').default else None From 0d6301ff25ba33f8adc082c67a59efe5b748dd96 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 27 Mar 2025 16:26:11 -0400 Subject: [PATCH 070/122] samplers add manual sigma adjustment Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 3 +- modules/processing_callbacks.py | 4 +-- modules/schedulers/scheduler_dpm_flowmatch.py | 5 ++- modules/sd_samplers_diffusers.py | 36 +++++++++++-------- modules/shared.py | 3 ++ modules/ui_sections.py | 32 ++++++++++++----- modules/ui_video.py | 2 +- scripts/xyz_grid_classes.py | 1 + 8 files changed, 57 insertions(+), 29 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c73845f29..0dc1b5d8f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -107,7 +107,8 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - update `diffusers` and other requirements - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis - - LoRA enable memory cache by default + - **LoRA** enable memory cache by default + - **Samplers** add ability to set sigma adjustment for each sampler - **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index cb90a5950..0ab91baa6 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -59,8 +59,6 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {} if debug: debug_callback(f'Callback: step={step} timestep={timestep} latents={latents.shape if latents is not None else None} kwargs={list(kwargs)}') shared.state.step() - # order = getattr(pipe.scheduler, "order", 1) if hasattr(pipe, 'scheduler') else 1 - # shared.state.sampling_step = step // order if shared.state.interrupted or shared.state.skipped: raise AssertionError('Interrupted...') if shared.state.paused: @@ -125,6 +123,8 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {} try: shared.state.current_sigma = pipe.scheduler.sigmas[pipe.scheduler.step_index-1] shared.state.current_sigma_next = pipe.scheduler.sigmas[pipe.scheduler.step_index] + if (shared.opts.schedulers_sigma_adjust != 1.0) and (timestep > 1000 * shared.opts.schedulers_sigma_adjust_min) and (timestep < 1000 * shared.opts.schedulers_sigma_adjust_max): + pipe.scheduler.sigmas[pipe.scheduler.step_index+1] = pipe.scheduler.sigmas[pipe.scheduler.step_index+1] * shared.opts.schedulers_sigma_adjust except Exception: pass except Exception as e: diff --git a/modules/schedulers/scheduler_dpm_flowmatch.py b/modules/schedulers/scheduler_dpm_flowmatch.py index ab9aa47a9..c1f045e8a 100644 --- a/modules/schedulers/scheduler_dpm_flowmatch.py +++ b/modules/schedulers/scheduler_dpm_flowmatch.py @@ -509,7 +509,10 @@ class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): def t_fn(_sigma: torch.Tensor) -> torch.Tensor: return _sigma.log().neg() sigma = self.sigmas[self.step_index] - sigma_next = self.sigmas[self.step_index + 1] + try: + sigma_next = self.sigmas[self.step_index + 1] + except Exception: + sigma_next = self.sigmas[-1] sigma_prev = self.sigmas[self.step_index - 1] if self.config.algorithm_type == "dpmsolver2": if self.config.solver_order == 2: diff --git a/modules/sd_samplers_diffusers.py b/modules/sd_samplers_diffusers.py index 523354e66..644379f1b 100644 --- a/modules/sd_samplers_diffusers.py +++ b/modules/sd_samplers_diffusers.py @@ -193,8 +193,7 @@ class DiffusionSampler: self.name = name self.config = {} self.sampler = None - # if not hasattr(model, 'scheduler'): - # return + if getattr(model, "default_scheduler", None) is None and (model is not None): # sanity check model.default_scheduler = copy.deepcopy(model.scheduler) for key, value in config.get('All', {}).items(): # apply global defaults @@ -217,6 +216,7 @@ class DiffusionSampler: for key, value in kwargs.items(): # apply user args, if any if key in self.config: self.config[key] = value + # finally apply user preferences if shared.opts.schedulers_prediction_type != 'default': self.config['prediction_type'] = shared.opts.schedulers_prediction_type @@ -283,6 +283,7 @@ class DiffusionSampler: del self.config['prediction_type'] if 'SGM' in name: self.config['timestep_spacing'] = 'trailing' + # validate all config params signature = inspect.signature(constructor, follow_wrapped=True) possible = signature.parameters.keys() @@ -293,7 +294,8 @@ class DiffusionSampler: debug_log(f'Sampler: name="{name}"') debug_log(f'Sampler: config={self.config}') debug_log(f'Sampler: signature={possible}') - # shared.log.debug_log(f'Sampler: sampler="{name}" config={self.config}') + + # finally create the new sampler try: sampler = constructor(**self.config) except Exception as e: @@ -302,21 +304,25 @@ class DiffusionSampler: errors.display(e, 'Samplers') self.sampler = None return - accept_sigmas = "sigmas" in set(inspect.signature(sampler.set_timesteps).parameters.keys()) - accepts_timesteps = "timesteps" in set(inspect.signature(sampler.set_timesteps).parameters.keys()) - accept_scale_noise = hasattr(sampler, "scale_noise") - debug_log(f'Sampler: sampler="{name}" sigmas={accept_sigmas} timesteps={accepts_timesteps}') - if ('Flux' in model.__class__.__name__) and (not accept_sigmas): - shared.log.warning(f'Sampler: sampler="{name}" does not accept sigmas') - self.sampler = None - return - if ('StableDiffusion3' in model.__class__.__name__) and (not accept_scale_noise): - shared.log.warning(f'Sampler: sampler="{name}" does not implement scale noise') - self.sampler = None - return + + if hasattr(sampler, 'set_timesteps'): + accept_sigmas = "sigmas" in set(inspect.signature(sampler.set_timesteps).parameters.keys()) + accepts_timesteps = "timesteps" in set(inspect.signature(sampler.set_timesteps).parameters.keys()) + accept_scale_noise = hasattr(sampler, "scale_noise") + debug_log(f'Sampler: sampler="{name}" sigmas={accept_sigmas} timesteps={accepts_timesteps}') + if ('Flux' in model.__class__.__name__) and (not accept_sigmas): + shared.log.warning(f'Sampler: sampler="{name}" does not accept sigmas') + self.sampler = None + return + if ('StableDiffusion3' in model.__class__.__name__) and (not accept_scale_noise): + shared.log.warning(f'Sampler: sampler="{name}" does not implement scale noise') + self.sampler = None + return + self.sampler = sampler if name == 'DC Solver': if not hasattr(self.sampler, 'dc_ratios'): pass + # shared.log.debug_log(f'Sampler: class="{self.sampler.__class__.__name__}" config={self.sampler.config}') self.sampler.name = name diff --git a/modules/shared.py b/modules/shared.py index 045d87dcf..95a205e43 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -792,6 +792,9 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"), 'schedulers_timesteps_range': OptionInfo(1000, "Timesteps range", gr.Slider, {"minimum": 250, "maximum": 4000, "step": 1, "visible": native}), 'schedulers_shift': OptionInfo(3, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": False}), 'schedulers_dynamic_shift': OptionInfo(False, "Sampler dynamic shift", gr.Checkbox, {"visible": False}), + 'schedulers_sigma_adjust': OptionInfo(1.0, "Sigma adjust", gr.Slider, {"minimum": 0.5, "maximum": 1.5, "step": 0.01, "visible": False}), + 'schedulers_sigma_adjust_min': OptionInfo(0.2, "Sigma adjust start", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01, "visible": False}), + 'schedulers_sigma_adjust_max': OptionInfo(0.8, "Sigma adjust end", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01, "visible": False}), # managed from ui.py for backend original k-diffusion "always_batch_cond_uncond": OptionInfo(False, "Disable conditional batching", gr.Checkbox, {"visible": not native}), diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 8d038ab13..2785c40b6 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -140,22 +140,22 @@ def create_seed_inputs(tab, reuse_visible=True, accordion=True, subseed_visible= return seed, reuse_seed, subseed, reuse_subseed, subseed_strength, seed_resize_from_h, seed_resize_from_w -def create_video_inputs(tab:str): +def create_video_inputs(tab:str, show_always:bool=False): def video_type_change(video_type): return [ - gr.update(visible=video_type != 'None'), - gr.update(visible=video_type in ['GIF', 'PNG']), - gr.update(visible=video_type not in ['None', 'GIF', 'PNG']), - gr.update(visible=video_type not in ['None', 'GIF', 'PNG']), + gr.update(visible=video_type != 'None' or show_always), + gr.update(visible=video_type in ['GIF', 'PNG'] or show_always), + gr.update(visible=video_type not in ['None', 'GIF', 'PNG'] or show_always), + gr.update(visible=video_type not in ['None', 'GIF', 'PNG'] or show_always), ] with gr.Column(): video_codecs = ['None', 'GIF', 'PNG', 'MP4/MP4V', 'MP4/AVC1', 'MP4/JVT3', 'MKV/H264', 'AVI/DIVX', 'AVI/RGBA', 'MJPEG/MJPG', 'MPG/MPG1', 'AVR/AVR1'] video_type = gr.Dropdown(label='Save video', choices=video_codecs, value='None', elem_id=f"{tab}_video_type") with gr.Column(): - video_duration = gr.Slider(label='Duration', minimum=0.25, maximum=300, step=0.25, value=2, visible=False, elem_id=f"{tab}_video_duration") - video_loop = gr.Checkbox(label='Loop', value=True, visible=False, elem_id=f"{tab}_video_loop") - video_pad = gr.Slider(label='Pad frames', minimum=0, maximum=24, step=1, value=1, visible=False, elem_id=f"{tab}_video_pad") - video_interpolate = gr.Slider(label='Interpolate frames', minimum=0, maximum=24, step=1, value=0, visible=False, elem_id=f"{tab}_video_interpolate") + video_duration = gr.Slider(label='Duration', minimum=0.25, maximum=300, step=0.25, value=2, visible=show_always, elem_id=f"{tab}_video_duration") + video_loop = gr.Checkbox(label='Loop', value=True, visible=show_always, elem_id=f"{tab}_video_loop") + video_pad = gr.Slider(label='Pad frames', minimum=0, maximum=24, step=1, value=1, visible=show_always, elem_id=f"{tab}_video_pad") + video_interpolate = gr.Slider(label='Interpolate frames', minimum=0, maximum=24, step=1, value=0, visible=show_always, elem_id=f"{tab}_video_interpolate") video_type.change(fn=video_type_change, inputs=[video_type], outputs=[video_duration, video_loop, video_pad, video_interpolate]) return video_type, video_duration, video_loop, video_pad, video_interpolate @@ -274,6 +274,13 @@ def create_sampler_options(tabname): shared.opts.schedulers_shift = sampler_shift shared.opts.save(shared.config_filename, silent=True) + def set_sigma_ajust(val, start, end): + shared.log.debug(f'Sampler set options: sigma={val} min={start} max={end}') + shared.opts.schedulers_sigma_adjust = val + shared.opts.schedulers_sigma_adjust_min = start + shared.opts.schedulers_sigma_adjust_max = end + shared.opts.save(shared.config_filename, silent=True) + # 'linear', 'scaled_linear', 'squaredcos_cap_v2' def set_sampler_preset(preset): if preset == 'AYS SD15': @@ -305,6 +312,10 @@ def create_sampler_options(tabname): with gr.Row(elem_classes=['flex-break']): sampler_presets = gr.Dropdown(label='Timesteps presets', elem_id=f"{tabname}_sampler_presets", choices=['None', 'AYS SD15', 'AYS SDXL'], value='None', type='value') sampler_timesteps = gr.Textbox(label='Timesteps override', elem_id=f"{tabname}_sampler_timesteps", value=shared.opts.schedulers_timesteps) + with gr.Row(elem_classes=['flex-break']): + sampler_sigma_adjust_val = gr.Slider(minimum=0.5, maximum=1.5, step=0.01, label='Sigma adjust', value=shared.opts.schedulers_sigma_adjust, elem_id=f"{tabname}_sampler_sigma_adjust") + sampler_sigma_adjust_min = gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Adjust start', value=shared.opts.schedulers_sigma_adjust_min, elem_id=f"{tabname}_sampler_sigma_adjust_min") + sampler_sigma_adjust_max = gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Adjust end', value=shared.opts.schedulers_sigma_adjust_max, elem_id=f"{tabname}_sampler_sigma_adjust_max") with gr.Row(elem_classes=['flex-break']): sampler_order = gr.Slider(minimum=0, maximum=5, step=1, label="Sampler order", value=shared.opts.schedulers_solver_order, elem_id=f"{tabname}_sampler_order") sampler_shift = gr.Slider(minimum=0, maximum=10, step=0.1, label="Flow shift", value=shared.opts.schedulers_shift, elem_id=f"{tabname}_sampler_shift") @@ -326,6 +337,9 @@ def create_sampler_options(tabname): sampler_order.change(fn=set_sampler_order, inputs=[sampler_order], outputs=[]) sampler_shift.change(fn=set_sampler_shift, inputs=[sampler_shift], outputs=[]) sampler_options.change(fn=set_sampler_options, inputs=[sampler_options], outputs=[]) + sampler_sigma_adjust_val.change(fn=set_sigma_ajust, inputs=[sampler_sigma_adjust_val, sampler_sigma_adjust_min, sampler_sigma_adjust_max], outputs=[]) + sampler_sigma_adjust_min.change(fn=set_sigma_ajust, inputs=[sampler_sigma_adjust_val, sampler_sigma_adjust_min, sampler_sigma_adjust_max], outputs=[]) + sampler_sigma_adjust_max.change(fn=set_sigma_ajust, inputs=[sampler_sigma_adjust_val, sampler_sigma_adjust_min, sampler_sigma_adjust_max], outputs=[]) def create_hires_inputs(tab): diff --git a/modules/ui_video.py b/modules/ui_video.py index f94e46ff6..d67bdc5eb 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -115,7 +115,7 @@ def create_ui(): with gr.Row(): save_frames = gr.Checkbox(label='Save image frames', value=False, elem_id="video_save_frames") with gr.Row(): - video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video') + video_type, video_duration, video_loop, video_pad, video_interpolate = ui_sections.create_video_inputs(tab='video', show_always=True) override_settings = ui_common.create_override_inputs('video') # output panel with gallery and video tabs diff --git a/scripts/xyz_grid_classes.py b/scripts/xyz_grid_classes.py index cd4df56e8..abcc1d9dd 100644 --- a/scripts/xyz_grid_classes.py +++ b/scripts/xyz_grid_classes.py @@ -119,6 +119,7 @@ axis_options = [ AxisOptionTxt2Img("[Sampler] Name", str, apply_sampler, fmt=format_value_add_label, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]), AxisOptionImg2Img("[Sampler] Name", str, apply_sampler, fmt=format_value_add_label, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers_for_img2img]), AxisOption("[Sampler] Sigma method", str, apply_setting("schedulers_sigma"), choices=lambda: ['default', 'karras', 'betas', 'exponential', 'lambdas']), + AxisOption("[Sampler] Sigma adjust", float, apply_setting("schedulers_sigma_adjust")), AxisOption("[Sampler] Timestep spacing", str, apply_setting("schedulers_timestep_spacing"), choices=lambda: ['default', 'linspace', 'leading', 'trailing']), AxisOption("[Sampler] Timestep range", int, apply_setting("schedulers_timesteps_range")), AxisOption("[Sampler] Solver order", int, apply_setting("schedulers_solver_order")), From e224bf48c995e48617fc19ade38cbd554a185ad4 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 09:09:20 -0400 Subject: [PATCH 071/122] xyz restore settings on done Signed-off-by: Vladimir Mandic --- TODO.md | 2 +- extensions-builtin/sdnext-modernui | 2 +- javascript/sdnext.css | 1 + scripts/xyz_grid_classes.py | 68 ++++++++++++++++++++++++------ 4 files changed, 58 insertions(+), 15 deletions(-) diff --git a/TODO.md b/TODO.md index 1bf6102a6..dcfeb76a9 100644 --- a/TODO.md +++ b/TODO.md @@ -9,7 +9,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: Hunyuan Video I2V: requires `transformers==4.47.1` - Video: Latte 1 T2V: dtype mismatch - Video: WAN 2.1 14B I2V 480p/720p: broken offload -- Video: WAN 2.1 14B I2V 480p/720p: custom number of frames +- Video: WAN 2.1 14B I2V 480p/720p: broken custom number of frames - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output - Video: Allegro T2V: all-gray output diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index baee2e439..d41350e07 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit baee2e4397537554e6db67e7e2496258f024880f +Subproject commit d41350e07f0c94e733c5ffbda8d64b07e313eea3 diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 3e841346d..e1425271d 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -115,6 +115,7 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_actions_column, #img2img_actions_column, #control_actions, #video_actions { flex-flow: wrap; justify-content: space-between; } #txt2img_seed, #img2img_seed, #control_seed, #video_seed { min-width: 90px !important } #video_generate_box>button { max-width: unset; } +#interrogate_output_prompt>textarea { resize: vertical; } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } .interrogate:hover { background: var(--button-primary-background-fill-hover) !important; } diff --git a/scripts/xyz_grid_classes.py b/scripts/xyz_grid_classes.py index abcc1d9dd..7c7cd1b48 100644 --- a/scripts/xyz_grid_classes.py +++ b/scripts/xyz_grid_classes.py @@ -25,26 +25,56 @@ class AxisOptionTxt2Img(AxisOption): class SharedSettingsStackHelper(object): - vae = None - schedulers_solver_order = None - tome_ratio = None - todo_ratio = None sd_model_checkpoint = None sd_model_refiner = None sd_model_dict = None sd_vae = None sd_unet = None sd_text_encoder = None + prompt_attention = None + freeu_b1 = None + freeu_b2 = None + freeu_s1 = None + freeu_s2 = None + schedulers_sigma_adjust = None + schedulers_beta_schedule = None + schedulers_beta_start = None + schedulers_beta_end = None + schedulers_shift = None + schedulers_sigma = None + schedulers_timestep_spacing = None + schedulers_timesteps_range = None + schedulers_beta_schedule = None + schedulers_beta_start = None + schedulers_beta_end = None + schedulers_shift = None + scheduler_eta = None + schedulers_solver_order = None + eta_noise_seed_delta = None + tome_ratio = None + todo_ratio = None extra_networks_default_multiplier = None disable_weights_auto_swap = None - prompt_attention = None def __enter__(self): - #Save overridden settings so they can be restored later. - self.vae = shared.opts.sd_vae + # Save overridden settings so they can be restored later + self.prompt_attention = shared.opts.prompt_attention + self.schedulers_sigma_adjust = shared.opts.schedulers_sigma_adjust + self.schedulers_timestep_spacing = shared.opts.schedulers_timestep_spacing + self.schedulers_timesteps_range = shared.opts.schedulers_timesteps_range self.schedulers_solver_order = shared.opts.schedulers_solver_order + self.schedulers_beta_schedule = shared.opts.schedulers_beta_schedule + self.schedulers_beta_start = shared.opts.schedulers_beta_start + self.schedulers_beta_end = shared.opts.schedulers_beta_end + self.schedulers_shift = shared.opts.schedulers_shift + self.scheduler_eta = shared.opts.scheduler_eta + self.eta_noise_seed_delta = shared.opts.eta_noise_seed_delta self.tome_ratio = shared.opts.tome_ratio self.todo_ratio = shared.opts.todo_ratio + self.freeu_b1 = shared.opts.freeu_b1 + self.freeu_b2 = shared.opts.freeu_b2 + self.freeu_s1 = shared.opts.freeu_s1 + self.freeu_s2 = shared.opts.freeu_s2 self.sd_model_checkpoint = shared.opts.sd_model_checkpoint self.sd_model_refiner = shared.opts.sd_model_refiner self.sd_model_dict = shared.opts.sd_model_dict @@ -53,18 +83,30 @@ class SharedSettingsStackHelper(object): self.sd_text_encoder = shared.opts.sd_text_encoder self.extra_networks_default_multiplier = shared.opts.extra_networks_default_multiplier self.disable_weights_auto_swap = shared.opts.disable_weights_auto_swap - self.prompt_attention = shared.opts.prompt_attention shared.opts.data["disable_weights_auto_swap"] = False def __exit__(self, exc_type, exc_value, tb): - #Restore overriden settings after plot generation. + # Restore overriden settings after plot generation shared.opts.data["disable_weights_auto_swap"] = self.disable_weights_auto_swap - shared.opts.data["sd_vae"] = self.vae - shared.opts.data["schedulers_solver_order"] = self.schedulers_solver_order - shared.opts.data["tome_ratio"] = self.tome_ratio - shared.opts.data["todo_ratio"] = self.todo_ratio shared.opts.data["extra_networks_default_multiplier"] = self.extra_networks_default_multiplier shared.opts.data["prompt_attention"] = self.prompt_attention + shared.opts.data["schedulers_solver_order"] = self.schedulers_solver_order + shared.opts.data["schedulers_sigma_adjust"] = self.schedulers_sigma_adjust + shared.opts.data["schedulers_timestep_spacing"] = self.schedulers_timestep_spacing + shared.opts.data["schedulers_timesteps_range"] = self.schedulers_timesteps_range + shared.opts.data["schedulers_beta_schedule"] = self.schedulers_beta_schedule + shared.opts.data["schedulers_beta_start"] = self.schedulers_beta_start + shared.opts.data["schedulers_beta_end"] = self.schedulers_beta_end + shared.opts.data["schedulers_shift"] = self.schedulers_shift + shared.opts.data["scheduler_eta"] = self.scheduler_eta + shared.opts.data["eta_noise_seed_delta"] = self.eta_noise_seed_delta + shared.opts.data["freeu_b1"] = self.freeu_b1 + shared.opts.data["freeu_b2"] = self.freeu_b2 + shared.opts.data["freeu_s1"] = self.freeu_s1 + shared.opts.data["freeu_s2"] = self.freeu_s2 + shared.opts.data["tome_ratio"] = self.tome_ratio + shared.opts.data["todo_ratio"] = self.todo_ratio + if self.sd_model_checkpoint != shared.opts.sd_model_checkpoint: shared.opts.data["sd_model_checkpoint"] = self.sd_model_checkpoint sd_models.reload_model_weights(op='model') From d1c3b97c65c9764b64e28be955fa42b332fc6465 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 14:05:28 -0400 Subject: [PATCH 072/122] add prompt enhance Signed-off-by: Vladimir Mandic --- .ruff.toml | 1 + CHANGELOG.md | 7 + extensions-builtin/sdnext-modernui | 2 +- javascript/sdnext.css | 2 + modules/interrogate/deepbooru.py | 6 +- modules/interrogate/deepseek.py | 4 +- modules/interrogate/openclip.py | 16 +- modules/interrogate/vqa.py | 67 ++++++-- modules/model_quant.py | 6 +- modules/processing_class.py | 3 + modules/shared.py | 2 +- scripts/flux_prompt_enhance.py | 99 ++++++++++++ scripts/prompt_enhance.py | 242 ++++++++++++++++++++--------- 13 files changed, 350 insertions(+), 107 deletions(-) create mode 100644 scripts/flux_prompt_enhance.py diff --git a/.ruff.toml b/.ruff.toml index 48f2e9026..6c77aa6f3 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -83,6 +83,7 @@ ignore = [ "F401", # Imported by unused "NPY002", # replace legacy random "RUF005", # Consider iterable unpacking + "RUF008", # Do not use mutable default values for dataclass "RUF010", # Use explicit conversion flag "RUF012", # Mutable class attributes "RUF013", # PEP 484 prohibits implicit `Optional` diff --git a/CHANGELOG.md b/CHANGELOG.md index 0dc1b5d8f..18e2617cb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -79,6 +79,13 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - [ByteDance/Sa2VA](https://huggingface.co/ByteDance/Sa2VA-1B) 1B, 4B simply select from list of available models in caption tab - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* +- **Prompt Enhance** + - new built-in extension available in text/image/control tabs + - can be used to manually or automatically enhance prompts using LLM + - supports **Gemma-3, Qwen-2.5, Phi-4, Llama-3.2, SmolLM2** + models are auto-downloaded on first use + also supports custom models that are compatible with `transformers/AutoModelForCausalLM` + - support quantization and offloading - [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) extension updates - add detection of prompt language and alphabet and filter based on those values - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index d41350e07..1e38e9b56 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit d41350e07f0c94e733c5ffbda8d64b07e313eea3 +Subproject commit 1e38e9b56edf45dd17402aee2a1c281dc26b5286 diff --git a/javascript/sdnext.css b/javascript/sdnext.css index e1425271d..9c858e37b 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -116,6 +116,8 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_seed, #img2img_seed, #control_seed, #video_seed { min-width: 90px !important } #video_generate_box>button { max-width: unset; } #interrogate_output_prompt>textarea { resize: vertical; } +#prompt_enhance_apply, #prompt_enhance_model { max-width: unset; } +#prompt_enhance_system textarea { color: var(--body-text-color-subdued) !important } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } .interrogate:hover { background: var(--button-primary-background-fill-hover) !important; } diff --git a/modules/interrogate/deepbooru.py b/modules/interrogate/deepbooru.py index 30227dc21..1e47e6cc8 100644 --- a/modules/interrogate/deepbooru.py +++ b/modules/interrogate/deepbooru.py @@ -4,7 +4,7 @@ import threading import torch import numpy as np from PIL import Image -from modules import modelloader, paths, devices, shared +from modules import modelloader, paths, devices, shared, sd_models re_special = re.compile(r'([\\()])') load_lock = threading.Lock() @@ -35,11 +35,11 @@ class DeepDanbooru: def start(self): self.load() - self.model.to(devices.device) + sd_models.move_model(self.model, devices.device) def stop(self): if shared.opts.interrogate_offload: - self.model.to(devices.cpu) + sd_models.move_model(self.model, devices.cpu) devices.torch_gc() def tag(self, pil_image): diff --git a/modules/interrogate/deepseek.py b/modules/interrogate/deepseek.py index 5138c5693..b2d340248 100644 --- a/modules/interrogate/deepseek.py +++ b/modules/interrogate/deepseek.py @@ -12,7 +12,7 @@ import os import sys import importlib from transformers import AutoModelForCausalLM -from modules import shared, devices, paths +from modules import shared, devices, paths, sd_models # model_path = "deepseek-ai/deepseek-vl2-small" @@ -73,7 +73,7 @@ def predict(question, image, repo): ).to(device=devices.device, dtype=devices.dtype) inputs_embeds = vl_gpt.prepare_inputs_embeds(**prepare_inputs) inputs_embeds = inputs_embeds.to(device=devices.device, dtype=devices.dtype) - vl_gpt = vl_gpt.to(devices.device) + sd_models.move_model(vl_gpt, devices.device) with devices.inference_context(): outputs = vl_gpt.language.generate( inputs_embeds=inputs_embeds, diff --git a/modules/interrogate/openclip.py b/modules/interrogate/openclip.py index 761c0fa39..792b4df85 100644 --- a/modules/interrogate/openclip.py +++ b/modules/interrogate/openclip.py @@ -10,7 +10,7 @@ import gradio as gr from PIL import Image from torchvision import transforms from torchvision.transforms.functional import InterpolationMode -from modules import devices, paths, shared, lowvram, errors +from modules import devices, paths, shared, lowvram, errors, sd_models caption_models = { @@ -125,7 +125,7 @@ class InterrogateModels: else: model, preprocess = clip.load(clip_model_name, download_root=shared.opts.clip_models_path) model.eval() - model = model.to(devices.device) + sd_models.move_model(model, devices.device) return model, preprocess def load(self): @@ -133,23 +133,23 @@ class InterrogateModels: self.blip_model = self.load_blip_model() if not shared.opts.no_half and not self.running_on_cpu: self.blip_model = self.blip_model.half() - self.blip_model = self.blip_model.to(devices.device) if self.clip_model is None: self.clip_model, self.clip_preprocess = self.load_clip_model() if not shared.opts.no_half and not self.running_on_cpu: self.clip_model = self.clip_model.half() - self.clip_model = self.clip_model.to(devices.device) self.dtype = next(self.clip_model.parameters()).dtype + sd_models.move_model(self.blip_model, devices.device) + sd_models.move_model(self.clip_model, devices.device) def send_clip_to_ram(self): if shared.opts.interrogate_offload: if self.clip_model is not None: - self.clip_model = self.clip_model.to(devices.cpu) + sd_models.move_model(self.blip_model, devices.cpu) def send_blip_to_ram(self): if shared.opts.interrogate_offload: if self.blip_model is not None: - self.blip_model = self.blip_model.to(devices.cpu) + sd_models.move_model(self.blip_model, devices.cpu) def unload(self): self.send_clip_to_ram() @@ -291,8 +291,8 @@ def load_interrogator(clip_model, blip_model): def unload_clip_model(): if ci is not None and shared.opts.interrogate_offload: - ci.caption_model = ci.caption_model.to(devices.cpu) - ci.clip_model = ci.clip_model.to(devices.cpu) + sd_models.move_model(ci.caption_model, devices.cpu) + sd_models.move_model(ci.clip_model, devices.cpu) ci.caption_offloaded = True ci.clip_offloaded = True devices.torch_gc() diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index 3db5251e8..b7f6aabb3 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -7,7 +7,8 @@ import torch import transformers import transformers.dynamic_module_utils from PIL import Image -from modules import shared, devices, errors +from modules import shared, devices, errors, sd_models + processor = None model = None @@ -74,7 +75,7 @@ def b64(image): def clean(response, question): - strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:', '<|im_end|>'] + strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:', '<|im_end|>', ''] if isinstance(response, dict): if 'task' in response: response = response['task'] @@ -113,13 +114,16 @@ def qwen(question: str, image: Image.Image, repo: str = None, system_prompt: str global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.Qwen2VLForConditionalGeneration.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir ) + model = model.to(devices.device, devices.dtype) processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model = model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.vlm_system conversation = [ @@ -157,10 +161,13 @@ def gemma(question: str, image: Image.Image, repo: str = None, system_prompt: st return '' if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.Gemma3ForConditionalGeneration.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + model = model.to(devices.device, devices.dtype) processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model = model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.vlm_system conversation = [ @@ -199,13 +206,16 @@ def paligemma(question: str, image: Image.Image, repo: str = None): if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') processor = transformers.PaliGemmaProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + model = None model = transformers.PaliGemmaForConditionalGeneration.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, ) + model = model.to(devices.device, devices.dtype) loaded = repo - model = model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) question = question.replace('<', '').replace('>', '').replace('_', ' ') model_inputs = processor(text=question, images=image, return_tensors="pt").to(devices.device, devices.dtype) input_len = model_inputs["input_ids"].shape[-1] @@ -228,6 +238,7 @@ def ovis(question: str, image: Image.Image, repo: str = None): global model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.AutoModelForCausalLM.from_pretrained( repo, torch_dtype=devices.dtype, @@ -235,8 +246,10 @@ def ovis(question: str, image: Image.Image, repo: str = None): trust_remote_code=True, cache_dir=shared.opts.hfcache_dir, ) + model = model.to(devices.device, devices.dtype) loaded = repo - model = model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) text_tokenizer = model.get_text_tokenizer() visual_tokenizer = model.get_visual_tokenizer() max_partition = 9 @@ -268,15 +281,18 @@ def smol(question: str, image: Image.Image, repo: str = None, system_prompt: str global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.AutoModelForVision2Seq.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, _attn_implementation="eager", ) + model.to(devices.device, devices.dtype) processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.vlm_system conversation = [ @@ -307,13 +323,16 @@ def git(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.GitForCausalLM.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, ) + model.to(devices.device, devices.dtype) processor = transformers.GitProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) pixel_values = processor(images=image, return_tensors="pt").pixel_values git_dict = {} git_dict['pixel_values'] = pixel_values.to(devices.device, devices.dtype) @@ -332,13 +351,16 @@ def blip(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.BlipForQuestionAnswering.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, ) + model.to(devices.device, devices.dtype) processor = transformers.BlipProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) inputs = processor(image, question, return_tensors="pt") inputs = inputs.to(devices.device, devices.dtype) with devices.inference_context(): @@ -351,13 +373,16 @@ def vilt(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.ViltForQuestionAnswering.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, ) + model.to(devices.device) processor = transformers.ViltProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model.to(devices.device) + devices.torch_gc() + sd_models.move_model(model, devices.device) inputs = processor(image, question, return_tensors="pt") inputs = inputs.to(devices.device) with devices.inference_context(): @@ -372,13 +397,16 @@ def pix(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.Pix2StructForConditionalGeneration.from_pretrained( repo, cache_dir=shared.opts.hfcache_dir, ) + model.to(devices.device) processor = transformers.Pix2StructProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo - model.to(devices.device) + devices.torch_gc() + sd_models.move_model(model, devices.device) if len(question) > 0: inputs = processor(images=image, text=question, return_tensors="pt").to(devices.device) else: @@ -393,6 +421,7 @@ def moondream(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}"') + model = None model = transformers.AutoModelForCausalLM.from_pretrained( repo, revision="2024-08-26", @@ -401,8 +430,10 @@ def moondream(question: str, image: Image.Image, repo: str = None): ) processor = transformers.AutoTokenizer.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo + model.to(devices.device, devices.dtype) model.eval() - model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) question = question.replace('<', '').replace('>', '').replace('_', ' ') encoded = model.encode_image(image) with devices.inference_context(): @@ -424,6 +455,7 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str if model is None or loaded != repo: shared.log.debug(f'Interrogate load: vlm="{repo}" path="{shared.opts.hfcache_dir}"') transformers.dynamic_module_utils.get_imports = get_imports + model = None model = transformers.AutoModelForCausalLM.from_pretrained( repo, trust_remote_code=True, @@ -433,8 +465,10 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir) transformers.dynamic_module_utils.get_imports = _get_imports loaded = repo + model.to(devices.device, devices.dtype) model.eval() - model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) if question.startswith('<'): task = question.split('>', 1)[0] + '>' else: @@ -456,12 +490,14 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str def sa2(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: + model = None model = transformers.AutoModel.from_pretrained( repo, torch_dtype=devices.dtype, low_cpu_mem_usage=True, use_flash_attn=False, trust_remote_code=True) + model = model.to(devices.device, devices.dtype) model = model.eval() processor = transformers.AutoTokenizer.from_pretrained( repo, @@ -469,7 +505,8 @@ def sa2(question: str, image: Image.Image, repo: str = None): use_fast=False, ) loaded = repo - model = model.to(devices.device, devices.dtype) + devices.torch_gc() + sd_models.move_model(model, devices.device) if question.startswith('<'): task = question.split('>', 1)[0] + '>' else: @@ -559,7 +596,7 @@ def interrogate(question, system_prompt, prompt, image, model_name, quiet:bool=F errors.display(e, 'VQA') answer = 'error' if shared.opts.interrogate_offload and model is not None: - model.to(devices.cpu) + sd_models.move_model(model, devices.cpu) devices.torch_gc() answer = clean(answer, question) t1 = time.time() diff --git a/modules/model_quant.py b/modules/model_quant.py index 0430dc9e0..e0fad40f9 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -45,7 +45,7 @@ def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Mode bnb_4bit_quant_type=shared.opts.bnb_quantization_type, bnb_4bit_compute_dtype=devices.dtype ) - log.debug(f'Quantization: module=all type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') + log.debug(f'Quantization: module="{module}" type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') if kwargs is None: return bnb_config else: @@ -62,7 +62,7 @@ def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model' if ao is None: return kwargs ao_config = diffusers.TorchAoConfig(shared.opts.torchao_quantization_type) - log.debug(f'Quantization: module=all type=torchao dtype={shared.opts.torchao_quantization_type}') + log.debug(f'Quantization: module="{module}" type=torchao dtype={shared.opts.torchao_quantization_type}') if kwargs is None: return ao_config else: @@ -82,7 +82,7 @@ def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = weights_dtype=shared.opts.quanto_quantization_type, ) quanto_config.activations = None # patch so it works with transformers - log.debug(f'Quantization: module=all type=quanto dtype={shared.opts.quanto_quantization_type}') + log.debug(f'Quantization: module="{module}" type=quanto dtype={shared.opts.quanto_quantization_type}') if kwargs is None: return quanto_config else: diff --git a/modules/processing_class.py b/modules/processing_class.py index eb4e333ec..e38a44fd9 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -111,6 +111,8 @@ class StableDiffusionProcessing: refiner_prompt: str = '', refiner_negative: str = '', hr_refiner_start: float = 0, + # prompt enhancer + enhance_prompt: bool = False, # save options outpath_samples=None, outpath_grids=None, @@ -145,6 +147,7 @@ class StableDiffusionProcessing: self.is_refiner_pass = False self.is_api = False self.scheduled_prompt = False + self.enhance_prompt = enhance_prompt self.prompt_embeds = [] self.positive_pooleds = [] self.negative_embeds = [] diff --git a/modules/shared.py b/modules/shared.py index 95a205e43..f4967cf93 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -509,7 +509,7 @@ options_templates.update(options_section(('backends', "Backend Settings"), { options_templates.update(options_section(('quantization', "Quantization Settings"), { "bnb_quantization_sep": OptionInfo("

BitsAndBytes

", "", gr.HTML), - "bnb_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder"], "visible": native}), + "bnb_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "LLM"], "visible": native}), "bnb_quantization_type": OptionInfo("nf4", "Quantization type", gr.Dropdown, {"choices": ['nf4', 'fp8', 'fp4'], "visible": native}), "bnb_quantization_storage": OptionInfo("uint8", "Backend storage", gr.Dropdown, {"choices": ["float16", "float32", "int8", "uint8", "float64", "bfloat16"], "visible": native}), diff --git a/scripts/flux_prompt_enhance.py b/scripts/flux_prompt_enhance.py new file mode 100644 index 000000000..17613964a --- /dev/null +++ b/scripts/flux_prompt_enhance.py @@ -0,0 +1,99 @@ +# repo: https://huggingface.co/gokaygokay/Flux-Prompt-Enhance + +import time +import random +import threading +from transformers import AutoTokenizer, AutoModelForSeq2SeqLM +import gradio as gr +from modules import shared, scripts, devices, processing + + +repo_id = "gokaygokay/Flux-Prompt-Enhance" +num_return_sequences = 5 +load_lock = threading.Lock() + + +class Script(scripts.Script): + prompts = [['']] + tokenizer: AutoTokenizer = None + model: AutoModelForSeq2SeqLM = None + prefix: str = "enhance prompt: " + button: gr.Button = None + auto_apply: gr.Checkbox = None + max_length: gr.Slider = None + temperature: gr.Slider = None + repetition_penalty: gr.Slider = None + table: gr.DataFrame = None + prompt: gr.Textbox = None + + def title(self): + return 'Prompt enhance' + + def show(self, is_img2img): + return shared.native + + def load(self): + with load_lock: + if self.tokenizer is None: + self.tokenizer = AutoTokenizer.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir) + if self.model is None: + shared.log.info(f'Prompt enhance: model="{repo_id}"') + self.model = AutoModelForSeq2SeqLM.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir).to(device=devices.cpu, dtype=devices.dtype) + + def enhance(self, prompt, auto_apply: bool = False, temperature: float = 0.7, repetition_penalty: float = 1.2, max_length: int = 128): + self.load() + t0 = time.time() + input_text = self.prefix + prompt + input_ids = self.tokenizer(input_text, return_tensors="pt").input_ids.to(devices.device) + self.model = self.model.to(devices.device) + kwargs = { + 'max_length': int(max_length), + 'num_return_sequences': int(num_return_sequences), + 'do_sample': True, + 'temperature': float(temperature), + 'repetition_penalty': float(repetition_penalty), + } + try: + outputs = self.model.generate(input_ids, **kwargs) + except Exception as e: + shared.log.error(f'Prompt enhance: error="{e}"') + return [['']] + self.model = self.model.to(devices.cpu) + prompts = self.tokenizer.batch_decode(outputs, skip_special_tokens=True) + prompts = [[p] for p in prompts] + t1 = time.time() + shared.log.info(f'Prompt enhance: temperature={temperature} repetition={repetition_penalty} length={max_length} sequences={num_return_sequences} apply={auto_apply} time={t1-t0:.2f}s') + return prompts + + def select(self, cell: gr.SelectData, _table): + prompt = cell.value if hasattr(cell, 'value') else cell + shared.log.info(f'Prompt enhance: prompt="{prompt}"') + return prompt + + def ui(self, _is_img2img): + with gr.Row(): + self.button = gr.Button(value='Enhance prompt') + self.auto_apply = gr.Checkbox(label='Auto apply', default=False) + with gr.Row(): + self.max_length = gr.Slider(label='Length', minimum=64, maximum=512, step=1, value=128) + self.temperature = gr.Slider(label='Temperature', minimum=0.1, maximum=2.0, step=0.05, value=0.7) + self.repetition_penalty = gr.Slider(label='Penalty', minimum=0.1, maximum=2.0, step=0.05, value=1.2) + with gr.Row(): + self.table = gr.DataFrame(self.prompts, label='', show_label=False, interactive=False, wrap=True, datatype="str", col_count=1, max_rows=num_return_sequences, headers=['Prompts']) + + if self.prompt is not None: + self.button.click(fn=self.enhance, inputs=[self.prompt, self.auto_apply, self.temperature, self.repetition_penalty, self.max_length], outputs=[self.table]) + self.table.select(fn=self.select, inputs=[self.table], outputs=[self.prompt]) + return [self.auto_apply, self.temperature, self.repetition_penalty, self.max_length] + + def run(self, p: processing.StableDiffusionProcessing, auto_apply, temperature, repetition_penalty, max_length): # pylint: disable=arguments-differ + if auto_apply: + p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + shared.log.debug(f'Prompt enhance: source="{p.prompt}"') + prompts = self.enhance(p.prompt, auto_apply, temperature, repetition_penalty, max_length) + p.prompt = random.choice(prompts)[0] + shared.log.debug(f'Prompt enhance: prompt="{p.prompt}"') + + def after_component(self, component, **kwargs): # searching for actual ui prompt components + if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt', 'video_prompt']: + self.prompt = component diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 17613964a..8d38f5ef1 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -1,99 +1,193 @@ -# repo: https://huggingface.co/gokaygokay/Flux-Prompt-Enhance - +from dataclasses import dataclass +import re import time -import random -import threading -from transformers import AutoTokenizer, AutoModelForSeq2SeqLM import gradio as gr -from modules import shared, scripts, devices, processing +import transformers +from modules import scripts, shared, devices, processing, sd_models -repo_id = "gokaygokay/Flux-Prompt-Enhance" -num_return_sequences = 5 -load_lock = threading.Lock() +@dataclass +class Options: + models = [ + 'Qwen/Qwen2.5-0.5B-Instruct', + 'Qwen/Qwen2.5-1.5B-Instruct', + 'Qwen/Qwen2.5-3B-Instruct', + 'google/gemma-3-1b-it', + 'google/gemma-3-4b-it', + 'microsoft/Phi-4-mini-instruct', + 'HuggingFaceTB/SmolLM2-135M-Instruct', + 'HuggingFaceTB/SmolLM2-360M-Instruct', + 'HuggingFaceTB/SmolLM2-1.7B-Instruct', + 'meta-llama/Llama-3.2-1B-Instruct', + 'meta-llama/Llama-3.2-3B-Instruct', + ] + default = models[3] + system_prompt: str = 'You are a helpful assistant. You will be given a prompt and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing, and any distinctive details (e.g., colors, textures, styles). Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. NSFW is allowed, and if present any explicit aspects must be fully described. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix.' + max_tokens: int = 50 + do_sample: bool = True + temperature: float = 0.5 + repetition_penalty: float = 1.2 class Script(scripts.Script): - prompts = [['']] - tokenizer: AutoTokenizer = None - model: AutoModelForSeq2SeqLM = None - prefix: str = "enhance prompt: " - button: gr.Button = None - auto_apply: gr.Checkbox = None - max_length: gr.Slider = None - temperature: gr.Slider = None - repetition_penalty: gr.Slider = None - table: gr.DataFrame = None prompt: gr.Textbox = None + model: str = None + llm: transformers.AutoModelForCausalLM = None + tokenizer: transformers.AutoProcessor = None + options = Options() def title(self): return 'Prompt enhance' - def show(self, is_img2img): - return shared.native + def show(self, _is_img2img): + return scripts.AlwaysVisible - def load(self): - with load_lock: - if self.tokenizer is None: - self.tokenizer = AutoTokenizer.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir) - if self.model is None: - shared.log.info(f'Prompt enhance: model="{repo_id}"') - self.model = AutoModelForSeq2SeqLM.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir).to(device=devices.cpu, dtype=devices.dtype) + def load(self, model:str=None): + model = model or self.options.default + if self.model is None or self.model != model: + t0 = time.time() + from modules import modelloader, model_quant + modelloader.hf_login() + quant_args = model_quant.create_config(module='LLM') + self.llm = None + self.llm = transformers.AutoModelForCausalLM.from_pretrained( + model, + trust_remote_code=True, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args, + ) + self.llm.eval() + self.tokenizer = transformers.AutoTokenizer.from_pretrained( + model, + cache_dir=shared.opts.hfcache_dir, + ) + self.model = model + devices.torch_gc() + t1 = time.time() + shared.log.debug(f'Prompt enhance: model="{model}" cls={self.llm.__class__.__name__} time={t1-t0:.2f} loaded') - def enhance(self, prompt, auto_apply: bool = False, temperature: float = 0.7, repetition_penalty: float = 1.2, max_length: int = 128): - self.load() + def clean(self, response): + if isinstance(response, list): + response = response[0] + response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n') + response = re.sub(r'<.*?>', '', response) + if 'prompt:' in response: + response = response.split('prompt:')[1] + if 'Prompt:' in response: + response = response.split('Prompt:')[1] + if '---' in response: + response = response.split('---')[0] + response = response.strip() + return response + + def enhance(self, model: str=None, prompt:str=None, system:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None): + model = model or self.options.default + prompt = prompt or self.prompt.value + system = system or self.options.system_prompt + tokens = tokens or self.options.max_tokens + penalty = penalty or self.options.repetition_penalty + temperature = temperature or self.options.temperature + sample = sample if sample is not None else self.options.do_sample + self.load(model) + if self.llm is None: + shared.log.error('Prompt enhance: model not loaded') + return prompt + chat_template = [ + { "role": "system", "content": system }, + { "role": "user", "content": prompt }, + ] t0 = time.time() - input_text = self.prefix + prompt - input_ids = self.tokenizer(input_text, return_tensors="pt").input_ids.to(devices.device) - self.model = self.model.to(devices.device) - kwargs = { - 'max_length': int(max_length), - 'num_return_sequences': int(num_return_sequences), - 'do_sample': True, - 'temperature': float(temperature), - 'repetition_penalty': float(repetition_penalty), - } try: - outputs = self.model.generate(input_ids, **kwargs) + inputs = self.tokenizer.apply_chat_template( + chat_template, + add_generation_prompt=True, + tokenize=True, + return_dict=True, + return_tensors="pt", + ).to(devices.device).to(devices.dtype) + input_len = inputs['input_ids'].shape[1] except Exception as e: - shared.log.error(f'Prompt enhance: error="{e}"') - return [['']] - self.model = self.model.to(devices.cpu) - prompts = self.tokenizer.batch_decode(outputs, skip_special_tokens=True) - prompts = [[p] for p in prompts] + shared.log.error(f'Prompt enhance tokenize: {e}') + return prompt + try: + with devices.inference_context(): + sd_models.move_model(self.llm, devices.device) + outputs = self.llm.generate( + **inputs, + do_sample=sample, + temperature=float(temperature), + max_new_tokens=int(input_len + tokens), + repetition_penalty=float(penalty), + ) + if shared.opts.diffusers_offload_mode != 'none': + sd_models.move_model(self.llm, devices.cpu) + devices.torch_gc() + raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) + shared.log.trace(f'Prompt enhance: raw="{raw_response}"') + outputs = outputs[:, input_len:] + response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) + except Exception as e: + shared.log.error(f'Prompt enhance generate: {e}') + response = self.clean(response) t1 = time.time() - shared.log.info(f'Prompt enhance: temperature={temperature} repetition={repetition_penalty} length={max_length} sequences={num_return_sequences} apply={auto_apply} time={t1-t0:.2f}s') - return prompts + shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt="{response}"') + return response - def select(self, cell: gr.SelectData, _table): - prompt = cell.value if hasattr(cell, 'value') else cell - shared.log.info(f'Prompt enhance: prompt="{prompt}"') - return prompt + + def apply(self, prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty): + response = self.enhance( + prompt=prompt, + model=llm_model, + system=prompt_system, + sample=do_sample, + tokens=max_tokens, + temperature=temperature, + penalty=repetition_penalty, + ) + if apply_prompt: + return [response, response] + return [response, gr.update()] def ui(self, _is_img2img): - with gr.Row(): - self.button = gr.Button(value='Enhance prompt') - self.auto_apply = gr.Checkbox(label='Auto apply', default=False) - with gr.Row(): - self.max_length = gr.Slider(label='Length', minimum=64, maximum=512, step=1, value=128) - self.temperature = gr.Slider(label='Temperature', minimum=0.1, maximum=2.0, step=0.05, value=0.7) - self.repetition_penalty = gr.Slider(label='Penalty', minimum=0.1, maximum=2.0, step=0.05, value=1.2) - with gr.Row(): - self.table = gr.DataFrame(self.prompts, label='', show_label=False, interactive=False, wrap=True, datatype="str", col_count=1, max_rows=num_return_sequences, headers=['Prompts']) - - if self.prompt is not None: - self.button.click(fn=self.enhance, inputs=[self.prompt, self.auto_apply, self.temperature, self.repetition_penalty, self.max_length], outputs=[self.table]) - self.table.select(fn=self.select, inputs=[self.table], outputs=[self.prompt]) - return [self.auto_apply, self.temperature, self.repetition_penalty, self.max_length] - - def run(self, p: processing.StableDiffusionProcessing, auto_apply, temperature, repetition_penalty, max_length): # pylint: disable=arguments-differ - if auto_apply: - p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) - shared.log.debug(f'Prompt enhance: source="{p.prompt}"') - prompts = self.enhance(p.prompt, auto_apply, temperature, repetition_penalty, max_length) - p.prompt = random.choice(prompts)[0] - shared.log.debug(f'Prompt enhance: prompt="{p.prompt}"') + with gr.Accordion('Prompt enhance', open=False, elem_id='prompt_enhance'): + with gr.Row(): + apply_btn = gr.Button(value='Enhance now', elem_id='prompt_enhance_apply', variant='primary') + with gr.Row(): + apply_prompt = gr.Checkbox(label='Apply to prompt', value=False) + apply_auto = gr.Checkbox(label='Auto enhance', value=False) + with gr.Group(): + with gr.Row(): + llm_model = gr.Dropdown(label='LLM model', choices=self.options.models, value=self.options.default, interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') + with gr.Row(): + prompt_system = gr.Textbox(label='System prompt', value=self.options.system_prompt, interactive=True, lines=4, elem_id='prompt_enhance_system') + with gr.Row(): + max_tokens = gr.Slider(label='Max tokens', value=self.options.max_tokens, minimum=10, maximum=1024, step=1, interactive=True) + do_sample = gr.Checkbox(label='Do sample', value=self.options.do_sample, interactive=True) + with gr.Row(): + temperature = gr.Slider(label='Temperature', value=self.options.temperature, minimum=0.0, maximum=1.0, step=0.01, interactive=True) + repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) + with gr.Row(): + prompt_output = gr.Textbox(label='Output', value='', interactive=True, lines=4) + apply_btn.click(fn=self.apply, inputs=[self.prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty], outputs=[prompt_output, self.prompt]) + return [apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty] def after_component(self, component, **kwargs): # searching for actual ui prompt components if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt', 'video_prompt']: self.prompt = component + + def before_process(self, p: processing.StableDiffusionProcessing, *args, **kwargs): # pylint: disable=unused-argument + apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty = args + if not apply_auto and not p.enhance_prompt: + return + p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + p.styles = [] + p.prompt = self.enhance( + prompt=p.prompt, + model=llm_model, + system=prompt_system, + sample=do_sample, + tokens=max_tokens, + temperature=temperature, + penalty=repetition_penalty, + ) From a6e8e8897432203dd6e99e6eaf9d41166ce2cb8d Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 14:23:54 -0400 Subject: [PATCH 073/122] fix wan2.1-i2v Signed-off-by: Vladimir Mandic --- TODO.md | 2 -- installer.py | 2 +- modules/video_models/video_overrides.py | 6 +++++- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/TODO.md b/TODO.md index dcfeb76a9..595f8938d 100644 --- a/TODO.md +++ b/TODO.md @@ -8,8 +8,6 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: Hunyuan Video I2V: requires `transformers==4.47.1` - Video: Latte 1 T2V: dtype mismatch -- Video: WAN 2.1 14B I2V 480p/720p: broken offload -- Video: WAN 2.1 14B I2V 480p/720p: broken custom number of frames - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output - Video: Allegro T2V: all-gray output diff --git a/installer.py b/installer.py index cb80e3962..d6308f519 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git or args.experimental: return - sha = '1ddf3f3a19095344166ad7207ebc5be7a862d17e' # diffusers commit hash + sha = '617c208bb4cc68fe4518164fee7cbdf5aa44ff78' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index 5436b485a..6a6619ebc 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -34,6 +34,10 @@ def set_overrides(p: processing.StableDiffusionProcessingVideo, selected: Model) p.task_args['generator'] = None if cls == 'LTXConditionPipeline': p.task_args['strength'] = p.denoising_strength - if 'LTX' in shared.sd_model.__class__.__name__: + if 'Wan' in cls: + p.task_args['width'] = 16 * (p.width // 16) + p.task_args['height'] = 16 * (p.height // 16) + p.frames = 4 * (p.frames // 4) + 1 + if 'LTX' in cls: p.task_args['width'] = 32 * (p.width // 32) p.task_args['height'] = 32 * (p.height // 32) From f4fdd496b9ed253587c75d6bf467acdd50fd1c18 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 14:46:52 -0400 Subject: [PATCH 074/122] more granular quantization modules options Signed-off-by: Vladimir Mandic --- TODO.md | 2 +- modules/intel/openvino/__init__.py | 4 ++-- modules/model_cogview.py | 4 ++-- modules/model_flux.py | 8 ++++---- modules/model_lumina.py | 4 ++-- modules/model_quant.py | 2 +- modules/model_sana.py | 3 +-- modules/model_sd3.py | 4 ++-- modules/onnx_impl/pipelines/__init__.py | 2 +- modules/sd_models.py | 2 +- modules/sd_models_utils.py | 2 +- modules/shared.py | 18 +++++++++--------- modules/video_models/video_load.py | 4 ++-- scripts/ltxvideo.py | 2 +- 14 files changed, 30 insertions(+), 31 deletions(-) diff --git a/TODO.md b/TODO.md index 595f8938d..73b11c743 100644 --- a/TODO.md +++ b/TODO.md @@ -19,7 +19,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: add generate context menu - Video: API support - Video: STG: -- Video SmoothCache: https://github.com/huggingface/diffusers/issues/11135 +- Video: SmoothCache: https://github.com/huggingface/diffusers/issues/11135 - SoftFill: https://github.com/zacheryvaughn/softfill-pipelines ## Code TODO diff --git a/modules/intel/openvino/__init__.py b/modules/intel/openvino/__init__.py index 2d2dc0b79..a5584575f 100644 --- a/modules/intel/openvino/__init__.py +++ b/modules/intel/openvino/__init__.py @@ -479,8 +479,8 @@ def openvino_fx(subgraph, example_inputs, options=None): subgraph_type[3] is torch.nn.modules.linear.Linear): dont_use_faketensors = True - dont_use_nncf = bool("Text Encoder" not in shared.opts.nncf_compress_weights) - dont_use_quant = bool("Text Encoder" not in shared.opts.nncf_quantize) + dont_use_nncf = bool("TE" not in shared.opts.nncf_compress_weights) + dont_use_quant = bool("TE" not in shared.opts.nncf_quantize) # Create a hash to be used for caching shared.compiled_model_state.model_hash_str = "" diff --git a/modules/model_cogview.py b/modules/model_cogview.py index a0594cc16..a37dd88ca 100644 --- a/modules/model_cogview.py +++ b/modules/model_cogview.py @@ -38,7 +38,7 @@ def load_cogview3(checkpoint_info, diffusers_load_config={}): **quant_args, ) - diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Text Encoder') + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='TE') text_encoder = transformers.T5EncoderModel.from_pretrained( repo_id, subfolder="text_encoder", @@ -71,7 +71,7 @@ def load_cogview4(checkpoint_info, diffusers_load_config={}): **quant_args, ) - diffusers_load_config, quant_args = load_common(diffusers_load_config, module='Text Encoder') + diffusers_load_config, quant_args = load_common(diffusers_load_config, module='TE') text_encoder = transformers.AutoModelForCausalLM.from_pretrained( repo_id, subfolder="text_encoder", diff --git a/modules/model_flux.py b/modules/model_flux.py index 97b39aedc..74ba5c8dd 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -112,12 +112,12 @@ def load_quants(kwargs, repo_id, cache_dir, allow_quant): quant_args = model_quant.create_config(allow=allow_quant) if not quant_args: return kwargs - if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization): + if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - quant_args = model_quant.create_config(allow=allow_quant, module='Text Encoder') + quant_args = model_quant.create_config(allow=allow_quant, module='TE') if not quant_args: return kwargs - if 'text_encoder_2' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization or 'Text Encoder' in shared.opts.quanto_quantization): + if 'text_encoder_2' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder_2'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) except Exception as e: shared.log.error(f'Quantization: {e}') @@ -192,7 +192,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change transformer, _text_encoder_2 = load_flux_nf4(file_path, prequantized=False) if transformer is not None: return transformer - quant_args = model_quant.create_config() + quant_args = model_quant.create_config(module='Transformer') if quant_args: shared.log.info(f'Load module: type=UNet/Transformer file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant=torchao dtype={devices.dtype}') transformer = diffusers.FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config, **quant_args) diff --git a/modules/model_lumina.py b/modules/model_lumina.py index 9e9ca8eff..e6bef6773 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -35,9 +35,9 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): quant_args = model_quant.create_config() kwargs = {} repo_id = sd_models.path_to_repo(checkpoint_info.name) - if ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): + if (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.Lumina2Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) - if ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization): + if ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization): kwargs['text_encoder'] = transformers.AutoModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) sd_model = diffusers.Lumina2Text2ImgPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config, **quant_args, **kwargs) return sd_model diff --git a/modules/model_quant.py b/modules/model_quant.py index e0fad40f9..4078dd01a 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -238,7 +238,7 @@ def apply_layerwise(sd_model, quiet:bool=False): m.enable_layerwise_casting(compute_dtype=devices.dtype, storage_dtype=storage_dtype, non_blocking=non_blocking) m.quantization_method = 'LayerWise' log.quiet(quiet, f'Quantization: type=layerwise module={module} cls={cls} storage={storage_dtype} compute={devices.dtype} blocking={not non_blocking}') - if module.startswith('text_encoder') and ('Model' in shared.opts.layerwise_quantization or 'Text Encoder' in shared.opts.layerwise_quantization) and ('clip' not in cls.lower()): + if module.startswith('text_encoder') and ('Model' in shared.opts.layerwise_quantization or 'TE' in shared.opts.layerwise_quantization) and ('clip' not in cls.lower()): m = getattr(sd_model, module) if hasattr(m, 'enable_layerwise_casting'): m.enable_layerwise_casting(compute_dtype=devices.dtype, storage_dtype=storage_dtype, non_blocking=non_blocking) diff --git a/modules/model_sana.py b/modules/model_sana.py index 509db4c88..8fb07a6da 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -1,4 +1,3 @@ -import os import time import torch import diffusers @@ -15,7 +14,7 @@ def load_quants(kwargs, repo_id, cache_dir): if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): kwargs['transformer'] = diffusers.models.SanaTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, **load_args, **quant_args) shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'text_encoder' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization): + if 'text_encoder' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder'] = transformers.AutoModelForCausalLM.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=cache_dir, **load_args, **quant_args) shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') return kwargs diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 59789155e..eaae0684f 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -55,10 +55,10 @@ def load_quants(kwargs, repo_id, cache_dir): quant_args = model_quant.create_config() if not quant_args: return kwargs - if 'Model' in shared.opts.bnb_quantization and 'transformer' not in kwargs: + if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'text_encoder_3' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization): + if 'text_encoder_3' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') return kwargs diff --git a/modules/onnx_impl/pipelines/__init__.py b/modules/onnx_impl/pipelines/__init__.py index a11b07fc7..62db80cd5 100644 --- a/modules/onnx_impl/pipelines/__init__.py +++ b/modules/onnx_impl/pipelines/__init__.py @@ -368,7 +368,7 @@ class OnnxRawPipeline(PipelineBase): if shared.opts.cuda_compile_backend == "olive-ai": submodels_for_olive = [] - if "Text Encoder" in shared.opts.cuda_compile: + if "TE" in shared.opts.cuda_compile: if not self.is_refiner: submodels_for_olive.append("text_encoder") if self._is_sdxl: diff --git a/modules/sd_models.py b/modules/sd_models.py index 43e926924..a7307561d 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -130,7 +130,7 @@ def set_diffuser_options(sd_model, vae=None, op:str='model', offload:bool=True, model.requires_grad_(False) model.eval() return model - sd_model = apply_function_to_model(sd_model, eval_model, ["Model", "VAE", "Text Encoder"], op="eval") + sd_model = apply_function_to_model(sd_model, eval_model, ["Model", "VAE", "TE"], op="eval") if len(shared.opts.torchao_quantization) > 0 and shared.opts.torchao_quantization_mode == 'post': sd_model = model_quant.torchao_quantization(sd_model) diff --git a/modules/sd_models_utils.py b/modules/sd_models_utils.py index cfbf38136..0b2794bac 100644 --- a/modules/sd_models_utils.py +++ b/modules/sd_models_utils.py @@ -164,7 +164,7 @@ def apply_function_to_model(sd_model, function, options, op=None): sd_model.prior_pipe.prior = function(sd_model.prior_pipe.prior, op="prior_pipe.prior", sd_model=sd_model) if op == "nncf" and "StableCascade" in sd_model.__class__.__name__: sd_model.prior_pipe.prior.clip_txt_pooled_mapper = backup_clip_txt_pooled_mapper - if "Text Encoder" in options: + if "TE" in options: if hasattr(sd_model, 'text_encoder') and hasattr(sd_model.text_encoder, 'config'): if hasattr(sd_model, 'decoder_pipe') and hasattr(sd_model.decoder_pipe, 'text_encoder') and hasattr(sd_model.decoder_pipe.text_encoder, 'config'): sd_model.decoder_pipe.text_encoder = function(sd_model.decoder_pipe.text_encoder, op="decoder_pipe.text_encoder", sd_model=sd_model) diff --git a/modules/shared.py b/modules/shared.py index f4967cf93..5d785f1a5 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -494,7 +494,7 @@ options_templates.update(options_section(('backends', "Backend Settings"), { "olive_cache_optimized": OptionInfo(True, 'Olive cache optimized models'), "ipex_sep": OptionInfo("

IPEX

", "", gr.HTML, {"visible": devices.backend == "ipex"}), - "ipex_optimize": OptionInfo([], "IPEX Optimize", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "Upscaler"], "visible": devices.backend == "ipex"}), + "ipex_optimize": OptionInfo([], "IPEX Optimize", gr.CheckboxGroup, {"choices": ["Model", "VAE", "TE", "Upscaler"], "visible": devices.backend == "ipex"}), "openvino_sep": OptionInfo("

OpenVINO

", "", gr.HTML, {"visible": cmd_opts.use_openvino}), "openvino_devices": OptionInfo([], "OpenVINO devices to use", gr.CheckboxGroup, {"choices": get_openvino_device_list() if cmd_opts.use_openvino else [], "visible": cmd_opts.use_openvino}), # pylint: disable=E0606 @@ -509,36 +509,36 @@ options_templates.update(options_section(('backends', "Backend Settings"), { options_templates.update(options_section(('quantization', "Quantization Settings"), { "bnb_quantization_sep": OptionInfo("

BitsAndBytes

", "", gr.HTML), - "bnb_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "LLM"], "visible": native}), + "bnb_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "VAE", "TE", "Video", "LLM"], "visible": native}), "bnb_quantization_type": OptionInfo("nf4", "Quantization type", gr.Dropdown, {"choices": ['nf4', 'fp8', 'fp4'], "visible": native}), "bnb_quantization_storage": OptionInfo("uint8", "Backend storage", gr.Dropdown, {"choices": ["float16", "float32", "int8", "uint8", "float64", "bfloat16"], "visible": native}), "quanto_quantization_sep": OptionInfo("

Optimum Quanto

", "", gr.HTML), - "quanto_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder"], "visible": native}), + "quanto_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "VAE", "TE", "Video", "LLM"], "visible": native}), "quanto_quantization_type": OptionInfo("int8", "Quantization weights type", gr.Dropdown, {"choices": ["float8", "int8", "int4", "int2"], "visible": native}), "optimum_quanto_sep": OptionInfo("

Optimum Quanto: post-load

", "", gr.HTML), - "optimum_quanto_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "ControlNet"], "visible": native}), + "optimum_quanto_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "VAE", "TE", "Video", "LLM", "ControlNet"], "visible": native}), "optimum_quanto_weights_type": OptionInfo("qint8", "Quantization weights type", gr.Dropdown, {"choices": ['qint8', 'qfloat8_e4m3fn', 'qfloat8_e5m2', 'qint4', 'qint2'], "visible": native}), "optimum_quanto_activations_type": OptionInfo("none", "Quantization activations type ", gr.Dropdown, {"choices": ['none', 'qint8', 'qfloat8_e4m3fn', 'qfloat8_e5m2'], "visible": native}), "optimum_quanto_shuffle_weights": OptionInfo(False, "Shuffle weights", gr.Checkbox, {"visible": native}), "torchao_sep": OptionInfo("

TorchAO

", "", gr.HTML), - "torchao_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder"], "visible": native}), + "torchao_quantization": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "VAE", "TE", "Video", "LLM"], "visible": native}), "torchao_quantization_mode": OptionInfo("pre", "Quantization mode", gr.Dropdown, {"choices": ['pre', 'post'], "visible": native}), "torchao_quantization_type": OptionInfo("int8_weight_only", "Quantization type", gr.Dropdown, {"choices": ['int4_weight_only', 'int8_dynamic_activation_int4_weight', 'int8_weight_only', 'int8_dynamic_activation_int8_weight', 'float8_weight_only', 'float8_dynamic_activation_float8_weight', 'float8_static_activation_float8_weight'], "visible": native}), "nncf_compress_sep": OptionInfo("

NNCF: Neural Network Compression Framework

", "", gr.HTML), - "nncf_compress_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "ControlNet"], "visible": native}), + "nncf_compress_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "VAE", "TE", "Video", "LLM", "ControlNet"], "visible": native}), "nncf_compress_weights_mode": OptionInfo("INT8", "Quantization type", gr.Dropdown, {"choices": ['INT8', 'INT8_SYM', 'INT4_ASYM', 'INT4_SYM', 'NF4'] if cmd_opts.use_openvino else ['INT8']}), "nncf_compress_weights_raito": OptionInfo(0, "Compress ratio", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": cmd_opts.use_openvino}), "nncf_compress_weights_group_size": OptionInfo(0, "Group size", gr.Slider, {"minimum": -1, "maximum": 512, "step": 1, "visible": cmd_opts.use_openvino}), - "nncf_quantize": OptionInfo([], "OpenVINO enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder"], "visible": cmd_opts.use_openvino}), + "nncf_quantize": OptionInfo([], "OpenVINO enabled", gr.CheckboxGroup, {"choices": ["Model", "VAE", "TE"], "visible": cmd_opts.use_openvino}), "nncf_quantize_mode": OptionInfo("INT8", "OpenVINO mode", gr.Dropdown, {"choices": ['INT8', 'FP8_E4M3', 'FP8_E5M2'], "visible": cmd_opts.use_openvino}), "nncf_quantize_shuffle_weights": OptionInfo(False, "Shuffle weights", gr.Checkbox, {"visible": native}), "layerwise_quantization_sep": OptionInfo("

Layerwise Casting

", "", gr.HTML), - "layerwise_quantization": OptionInfo([], "Layerwise casting enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "Text Encoder"], "visible": native}), + "layerwise_quantization": OptionInfo([], "Layerwise casting enabled", gr.CheckboxGroup, {"choices": ["Model", "Transformer", "TE"], "visible": native}), "layerwise_quantization_storage": OptionInfo("float8_e4m3fn", "Layerwise casting storage", gr.Dropdown, {"choices": ["float8_e4m3fn", "float8_e5m2"], "visible": native}), "layerwise_quantization_nonblocking": OptionInfo(False, "Layerwise non-blocking operations", gr.Checkbox, {"visible": native}), })) @@ -612,7 +612,7 @@ options_templates.update(options_section(('advanced', "Pipeline Modifiers"), { options_templates.update(options_section(('compile', "Model Compile"), { "cuda_compile_sep": OptionInfo("

Model Compile

", "", gr.HTML), - "cuda_compile": OptionInfo([] if not cmd_opts.use_openvino else ["Model", "VAE", "Upscaler"], "Compile Model", gr.CheckboxGroup, {"choices": ["Model", "VAE", "Text Encoder", "Upscaler"]}), + "cuda_compile": OptionInfo([] if not cmd_opts.use_openvino else ["Model", "VAE", "Upscaler"], "Compile Model", gr.CheckboxGroup, {"choices": ["Model", "VAE", "TE", "Upscaler"]}), "cuda_compile_backend": OptionInfo("none" if not cmd_opts.use_openvino else "openvino_fx", "Model compile backend", gr.Radio, {"choices": ['none', 'inductor', 'cudagraphs', 'aot_ts_nvfuser', 'hidet', 'migraphx', 'ipex', 'onediff', 'stable-fast', 'deep-cache', 'olive-ai', 'openvino_fx']}), "cuda_compile_mode": OptionInfo("default", "Model compile mode", gr.Radio, {"choices": ['default', 'reduce-overhead', 'max-autotune', 'max-autotune-no-cudagraphs']}), "cuda_compile_fullgraph": OptionInfo(True if not cmd_opts.use_openvino else False, "Model compile fullgraph"), diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 3bfeb2786..13d0dc935 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -19,7 +19,7 @@ def load_model(selected: models_def.Model): # text encoder try: - quant_args = model_quant.create_config(module='Text Encoder') + quant_args = model_quant.create_config(module='TE') debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={video_utils.get_quant(quant_args)}') text_encoder = selected.te_cls.from_pretrained( pretrained_model_name_or_path=selected.te or selected.repo, @@ -35,7 +35,7 @@ def load_model(selected: models_def.Model): # transformer try: - quant_args = model_quant.create_config(module='Model') + quant_args = model_quant.create_config(module='Video') debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={video_utils.get_quant(quant_args)}') transformer = selected.dit_cls.from_pretrained( pretrained_model_name_or_path=selected.dit or selected.repo, diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py index a1933673e..09f30360e 100644 --- a/scripts/ltxvideo.py +++ b/scripts/ltxvideo.py @@ -24,7 +24,7 @@ def load_quants(kwargs, repo_id): if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): kwargs['transformer'] = diffusers.LTXVideoTransformer3DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, **quant_args) shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'text_encoder' not in kwargs and ('Text Encoder' in shared.opts.bnb_quantization or 'Text Encoder' in shared.opts.torchao_quantization): + if 'text_encoder' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization): kwargs['text_encoder'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, **quant_args) shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') return kwargs From 5ce834c362b738ec377708e58e40958ff36cdf7c Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 14:50:37 -0400 Subject: [PATCH 075/122] fix missing quant options Signed-off-by: Vladimir Mandic --- modules/model_lumina.py | 2 +- modules/model_sana.py | 4 +--- modules/model_sd3.py | 2 -- 3 files changed, 2 insertions(+), 6 deletions(-) diff --git a/modules/model_lumina.py b/modules/model_lumina.py index e6bef6773..f9b09f23b 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -37,7 +37,7 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): repo_id = sd_models.path_to_repo(checkpoint_info.name) if (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.Lumina2Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) - if ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization): + if ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder'] = transformers.AutoModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) sd_model = diffusers.Lumina2Text2ImgPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config, **quant_args, **kwargs) return sd_model diff --git a/modules/model_sana.py b/modules/model_sana.py index 8fb07a6da..a31985ab6 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -11,12 +11,10 @@ def load_quants(kwargs, repo_id, cache_dir): if not quant_args: return kwargs load_args = kwargs.copy() - if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): + if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.models.SanaTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, **load_args, **quant_args) - shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') if 'text_encoder' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder'] = transformers.AutoModelForCausalLM.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=cache_dir, **load_args, **quant_args) - shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') return kwargs diff --git a/modules/model_sd3.py b/modules/model_sd3.py index eaae0684f..5b8006c2a 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -57,10 +57,8 @@ def load_quants(kwargs, repo_id, cache_dir): return kwargs if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') if 'text_encoder_3' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') return kwargs From 720d8a8df6ab6b0701cfe600785b7cc69fcea38a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 20:37:09 -0400 Subject: [PATCH 076/122] improve video interpolation logging Signed-off-by: Vladimir Mandic --- modules/rife/__init__.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/modules/rife/__init__.py b/modules/rife/__init__.py index 2a636eb2f..2e0bb19f9 100644 --- a/modules/rife/__init__.py +++ b/modules/rife/__init__.py @@ -25,7 +25,7 @@ def load(model_path: str = 'rife/flownet-v46.pkl'): from modules import modelloader model_dir = os.path.join(shared.models_path, 'RIFE') model_path = modelloader.load_file_from_url(url=model_url, model_dir=model_dir, file_name='flownet-v46.pkl') - shared.log.debug(f'RIFE load model: file="{model_path}"') + shared.log.debug(f'Video interpolate: model="{model_path}"') model = RifeModel() model.load_model(model_path, -1) model.eval() @@ -46,7 +46,6 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1, item = buffer.get() while item is not None: img = item[:, :, ::-1] - # image = Image.fromarray(cv2.cvtColor(img, cv2.COLOR_BGR2RGB)) image = Image.fromarray(img) item = buffer.get() interpolated.append(image) @@ -76,6 +75,7 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1, pw = ((w - 1) // tmp + 1) * tmp padding = (0, pw - w, 0, ph - h) buffer = Queue(maxsize=8192) + duplicate = 0 _thread.start_new_thread(write, (buffer,)) frame = cv2.cvtColor(np.array(images[0]), cv2.COLOR_RGB2BGR) @@ -93,6 +93,7 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1, I1_small = F.interpolate(I1, (32, 32), mode='bilinear', align_corners=False).to(torch.float32) ssim = ssim_matlab(I0_small[:, :3], I1_small[:, :3]) if ssim > 0.99: # skip duplicate frames + duplicate += 1 continue if ssim < change: output = [] @@ -110,8 +111,8 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1, for _i in range(pad): # fill ending frames buffer.put(frame) - while not buffer.empty(): + while not buffer.qsize() > 0: time.sleep(0.1) t1 = time.time() - shared.log.info(f'RIFE interpolate: input={len(images)} frames={len(interpolated)} width={w} height={h} interpolate={count} scale={scale} pad={pad} change={change} time={round(t1 - t0, 2)}') + shared.log.info(f'Video interpolate: input={len(images)} frames={len(interpolated)} buffer={buffer.qsize()} duplicate={duplicate} width={w} height={h} interpolate={count} scale={scale} pad={pad} change={change} time={round(t1 - t0, 2)}') return interpolated From 8343469be08ebe0c72508a8c14e1f1762e4281b1 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 21:03:27 -0400 Subject: [PATCH 077/122] improve prompt enhancer and update apply styles Signed-off-by: Vladimir Mandic --- TODO.md | 21 ++++++++++++--------- modules/lora/lora_extract.py | 2 +- modules/video_models/video_utils.py | 2 ++ scripts/animatediff.py | 2 ++ scripts/consistory_ext.py | 2 ++ scripts/flux_prompt_enhance.py | 3 +++ scripts/mixture_of_diffusers.py | 2 ++ scripts/prompt_enhance.py | 27 ++++++++++++++++++++++++--- scripts/pulid_ext.py | 2 ++ scripts/x_adapter.py | 2 ++ 10 files changed, 52 insertions(+), 13 deletions(-) diff --git a/TODO.md b/TODO.md index 73b11c743..3674fe5ef 100644 --- a/TODO.md +++ b/TODO.md @@ -24,14 +24,17 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Code TODO -- enable ROCm for windows when available -- resize image: enable full VAE mode for resize-latent -- infotext: handle using regex instead -- processing: remove duplicate mask params -- model loader: implement model in-memory caching -- hypertile: vae breaks when using non-standard sizes -- force-reloading entire model as loading transformers only leads to massive memory usage -- add other quantization types -- lora make support quantized flux - control: support scripts via api +- enable ROCm for windows when available +- fc: autodetect distilled based on model +- fc: autodetect tensor format based on model +- hypertile: vae breaks when using non-standard sizes +- infotext: handle using regex instead +- lora: add other quantization types +- lora: force-reloading entire model as loading transformers only leads to massive memory usage +- lora: required for flux to reapply offload after lora has been applied, but fails with oom +- lora: support pre-quantized flux +- model loader: implement model in-memory caching - modernui: monkey-patch for missing tabs.select event +- processing: remove duplicate mask params +- resize image: enable full VAE mode for resize-latent diff --git a/modules/lora/lora_extract.py b/modules/lora/lora_extract.py index 5050187a2..1217b952d 100644 --- a/modules/lora/lora_extract.py +++ b/modules/lora/lora_extract.py @@ -182,7 +182,7 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): progress.remove_task(task) t3 = time.time() - # TODO: lora make support quantized flux + # TODO: lora support pre-quantized flux # if 'te' in modules and getattr(shared.sd_model, 'transformer', None) is not None: # for name, module in shared.sd_model.transformer.named_modules(): # if "norm" in name and "linear" not in name: diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 4e26849b6..d0a80d42a 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -24,6 +24,8 @@ def get_url(url): def set_prompt(p): p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] p.task_args['prompt'] = p.prompt p.task_args['negative_prompt'] = p.negative_prompt diff --git a/scripts/animatediff.py b/scripts/animatediff.py index 34b7829c1..6fb77a45b 100644 --- a/scripts/animatediff.py +++ b/scripts/animatediff.py @@ -140,6 +140,8 @@ def set_scheduler(p, model, override: bool = False): def set_prompt(p): p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] prompts = p.prompt.split('\n') try: prompt = {} diff --git a/scripts/consistory_ext.py b/scripts/consistory_ext.py index 7a8f21e3b..c02ca1e50 100644 --- a/scripts/consistory_ext.py +++ b/scripts/consistory_ext.py @@ -118,6 +118,8 @@ class Script(scripts.Script): settings = [p.strip() for p in prompts.split('\n') if p.strip() != ''] anchors = [f'{subject} {p}' for p in settings] prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] prompts = [p.strip() for p in prompt.split('\n') if p.strip() != ''] for i, prompt in enumerate(prompts): if subject not in prompt: diff --git a/scripts/flux_prompt_enhance.py b/scripts/flux_prompt_enhance.py index 17613964a..abfbeae6d 100644 --- a/scripts/flux_prompt_enhance.py +++ b/scripts/flux_prompt_enhance.py @@ -89,6 +89,9 @@ class Script(scripts.Script): def run(self, p: processing.StableDiffusionProcessing, auto_apply, temperature, repetition_penalty, max_length): # pylint: disable=arguments-differ if auto_apply: p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] shared.log.debug(f'Prompt enhance: source="{p.prompt}"') prompts = self.enhance(p.prompt, auto_apply, temperature, repetition_penalty, max_length) p.prompt = random.choice(prompts)[0] diff --git a/scripts/mixture_of_diffusers.py b/scripts/mixture_of_diffusers.py index 463948d36..58598ec66 100644 --- a/scripts/mixture_of_diffusers.py +++ b/scripts/mixture_of_diffusers.py @@ -91,6 +91,8 @@ class Script(scripts.Script): p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] p.prompts, guidance = self.get_prompts(x_tiles, y_tiles, prompts, p.prompt, p.cfg_scale) p.all_prompts = p.prompts p.task_args['prompts'] = p.prompts diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 8d38f5ef1..15a9d5d88 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -55,6 +55,7 @@ class Script(scripts.Script): trust_remote_code=True, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir, + _attn_implementation="eager", **quant_args, ) self.llm.eval() @@ -67,6 +68,15 @@ class Script(scripts.Script): t1 = time.time() shared.log.debug(f'Prompt enhance: model="{model}" cls={self.llm.__class__.__name__} time={t1-t0:.2f} loaded') + def unload(self): + if self.llm is not None: + sd_models.move_model(self.llm, devices.cpu) + self.model = None + self.llm = None + self.tokenizer = None + devices.torch_gc() + shared.log.debug('Prompt enhance: model unloaded') + def clean(self, response): if isinstance(response, list): response = response[0] @@ -123,8 +133,8 @@ class Script(scripts.Script): if shared.opts.diffusers_offload_mode != 'none': sd_models.move_model(self.llm, devices.cpu) devices.torch_gc() - raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) - shared.log.trace(f'Prompt enhance: raw="{raw_response}"') + # raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) + # shared.log.trace(f'Prompt enhance: raw="{raw_response}"') outputs = outputs[:, input_len:] response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) except Exception as e: @@ -134,7 +144,6 @@ class Script(scripts.Script): shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt="{response}"') return response - def apply(self, prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty): response = self.enhance( prompt=prompt, @@ -159,6 +168,11 @@ class Script(scripts.Script): with gr.Group(): with gr.Row(): llm_model = gr.Dropdown(label='LLM model', choices=self.options.models, value=self.options.default, interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') + with gr.Row(): + load_btn = gr.Button(value='Load model', elem_id='prompt_enhance_load', variant='secondary') + load_btn.click(fn=self.load, inputs=[llm_model], outputs=[]) + unload_btn = gr.Button(value='Unload model', elem_id='prompt_enhance_unload', variant='secondary') + unload_btn.click(fn=self.unload, inputs=[], outputs=[]) with gr.Row(): prompt_system = gr.Textbox(label='System prompt', value=self.options.system_prompt, interactive=True, lines=4, elem_id='prompt_enhance_system') with gr.Row(): @@ -169,6 +183,11 @@ class Script(scripts.Script): repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) with gr.Row(): prompt_output = gr.Textbox(label='Output', value='', interactive=True, lines=4) + with gr.Row(): + clear_btn = gr.Button(value='Clear', elem_id='prompt_enhance_clear', variant='secondary') + clear_btn.click(fn=lambda: '', inputs=[], outputs=[prompt_output]) + copy_btn = gr.Button(value='Set prompt', elem_id='prompt_enhance_copy', variant='secondary') + copy_btn.click(fn=lambda x: x, inputs=[prompt_output], outputs=[self.prompt]) apply_btn.click(fn=self.apply, inputs=[self.prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty], outputs=[prompt_output, self.prompt]) return [apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty] @@ -181,6 +200,8 @@ class Script(scripts.Script): if not apply_auto and not p.enhance_prompt: return p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) p.styles = [] p.prompt = self.enhance( prompt=p.prompt, diff --git a/scripts/pulid_ext.py b/scripts/pulid_ext.py index 40e726549..dbc62715d 100644 --- a/scripts/pulid_ext.py +++ b/scripts/pulid_ext.py @@ -225,6 +225,8 @@ class Script(scripts.Script): processing.fix_seed(p) p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] with devices.inference_context(): output = shared.sd_model( prompt=p.prompt, diff --git a/scripts/x_adapter.py b/scripts/x_adapter.py index c67eca18b..08874aac9 100644 --- a/scripts/x_adapter.py +++ b/scripts/x_adapter.py @@ -110,6 +110,8 @@ class Script(scripts.Script): shared.opts.data['prompt_attention'] = 'fixed' prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) negative = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + shared.prompt_styles.apply_styles_to_extra(p) + p.styles = [] p.task_args['prompt'] = prompt p.task_args['negative_prompt'] = negative p.task_args['prompt_sd1_5'] = prompt From 1969ac96bbef8397af39cfff862c7c87b4f38a06 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 21:24:05 -0400 Subject: [PATCH 078/122] skip prompt enhance on interrupt Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 15a9d5d88..a218e081d 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -199,6 +199,8 @@ class Script(scripts.Script): apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty = args if not apply_auto and not p.enhance_prompt: return + if shared.state.skipped or shared.state.interrupted: + return p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) shared.prompt_styles.apply_styles_to_extra(p) From 8069c8caa9833a05b587d983c309aa7e3fe8ca00 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 21:29:18 -0400 Subject: [PATCH 079/122] add llm to progress Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index a218e081d..51ebd475e 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -205,6 +205,7 @@ class Script(scripts.Script): p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) shared.prompt_styles.apply_styles_to_extra(p) p.styles = [] + shared.state.begin('LLM') p.prompt = self.enhance( prompt=p.prompt, model=llm_model, @@ -214,3 +215,4 @@ class Script(scripts.Script): temperature=temperature, penalty=repetition_penalty, ) + shared.state.end() From 8014c6381d0fbfbf6349a7c99785f2f9002ca08a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 21:31:52 -0400 Subject: [PATCH 080/122] interpolate dont skip duplicate frames Signed-off-by: Vladimir Mandic --- modules/rife/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modules/rife/__init__.py b/modules/rife/__init__.py index 2e0bb19f9..b8e9fbc6f 100644 --- a/modules/rife/__init__.py +++ b/modules/rife/__init__.py @@ -94,7 +94,7 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1, ssim = ssim_matlab(I0_small[:, :3], I1_small[:, :3]) if ssim > 0.99: # skip duplicate frames duplicate += 1 - continue + # continue if ssim < change: output = [] for _i in range(pad): # fill frames if change rate is above threshold From 0cf30406c5b9570635bd13d9e9b0f07022a60581 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 28 Mar 2025 22:03:00 -0400 Subject: [PATCH 081/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 18e2617cb..cc1607ea6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,16 +1,21 @@ # Change Log for SD.Next -## Update for 2025-03-26 +## Update for 2025-03-28 -### Highlights for 2025-03-27 +### Highlights for 2025-03-28 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! -Plus support for **CogView-4**, **SANA 1.5**, new CLiP models, improvements to remote VAE, additional docs/guides -Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods +Also, support for new models: **CogView-4**, **SANA 1.5**, -### Details for 2025-03-27 +Plus... +- New **Prompt Enhance** using LLM, +- New **CLiP** models, improvements to **remote VAE**, additional wiki/docs/guides +- More quantization options and granular control +- Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods + +### Details for 2025-03-28 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -116,6 +121,7 @@ Pretty big performance updates to a) Any model using DiT based architecture: new - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis - **LoRA** enable memory cache by default - **Samplers** add ability to set sigma adjustment for each sampler + - **ModernUI** updates - **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide From f2f0390e9e397b255f565fb04325e2b82441c965 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 09:18:40 -0400 Subject: [PATCH 082/122] prompt enhance custom model loader Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 16 ++-- javascript/sdnext.css | 2 +- scripts/prompt_enhance.py | 168 ++++++++++++++++++++++++++++---------- 3 files changed, 133 insertions(+), 53 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index cc1607ea6..87ccbf2ee 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -67,6 +67,15 @@ Plus... download text encoders into folder set in settings -> system paths -> text encoders (default is *models/Text-encoder*) load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui +- **Prompt Enhance** + - new built-in extension available in text/image/control tabs + - can be used to manually or automatically enhance prompts using LLM + - built-in presets for **Gemma-3, Qwen-2.5, Phi-4, Llama-3.2, SmolLM2, Dolphin-3** + - support for custom models + load any models hosted on huggingface + load either model in huggingface format or `gguf` format + - models are auto-downloaded on first use + - support quantization and offloading - **Acceleration** - Support for most DiT-based models, for example: *FLUX.1, SD35, Hunyuan, Mochi, Latte, Allegro, Cog* - Enable and configure in *Settings -> Pipeline modifiers* @@ -84,13 +93,6 @@ Plus... - [ByteDance/Sa2VA](https://huggingface.co/ByteDance/Sa2VA-1B) 1B, 4B simply select from list of available models in caption tab - add option to set system prompt for vlm models that support it: *Gemma, Smol, Qwen* -- **Prompt Enhance** - - new built-in extension available in text/image/control tabs - - can be used to manually or automatically enhance prompts using LLM - - supports **Gemma-3, Qwen-2.5, Phi-4, Llama-3.2, SmolLM2** - models are auto-downloaded on first use - also supports custom models that are compatible with `transformers/AutoModelForCausalLM` - - support quantization and offloading - [NudeNet](https://github.com/vladmandic/sd-extension-nudenet/) extension updates - add detection of prompt language and alphabet and filter based on those values - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 9c858e37b..71f8cb43e 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -116,7 +116,7 @@ button.custom-button { border-radius: var(--button-large-radius); padding: var(- #txt2img_seed, #img2img_seed, #control_seed, #video_seed { min-width: 90px !important } #video_generate_box>button { max-width: unset; } #interrogate_output_prompt>textarea { resize: vertical; } -#prompt_enhance_apply, #prompt_enhance_model { max-width: unset; } +#prompt_enhance_apply, #prompt_enhance_model, #prompt_enhance_custom_load { max-width: unset; min-width: 100% !important; } #prompt_enhance_system textarea { color: var(--body-text-color-subdued) !important } .interrogate { position: absolute; right: 2.8em; top: 0.2em; max-width: fit-content; background: none !important; z-index: 50; font-size: 1.5em !important; } diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 51ebd475e..aa719e978 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -3,25 +3,34 @@ import re import time import gradio as gr import transformers -from modules import scripts, shared, devices, processing, sd_models +from modules import scripts, shared, devices, errors, processing, sd_models @dataclass class Options: - models = [ - 'Qwen/Qwen2.5-0.5B-Instruct', - 'Qwen/Qwen2.5-1.5B-Instruct', - 'Qwen/Qwen2.5-3B-Instruct', - 'google/gemma-3-1b-it', - 'google/gemma-3-4b-it', - 'microsoft/Phi-4-mini-instruct', - 'HuggingFaceTB/SmolLM2-135M-Instruct', - 'HuggingFaceTB/SmolLM2-360M-Instruct', - 'HuggingFaceTB/SmolLM2-1.7B-Instruct', - 'meta-llama/Llama-3.2-1B-Instruct', - 'meta-llama/Llama-3.2-3B-Instruct', - ] - default = models[3] + models = { + 'google/gemma-3-1b-it': {}, + 'google/gemma-3-4b-it': {}, + 'Qwen/Qwen2.5-0.5B-Instruct': {}, + 'Qwen/Qwen2.5-1.5B-Instruct': {}, + 'Qwen/Qwen2.5-3B-Instruct': {}, + 'microsoft/Phi-4-mini-instruct': {}, + 'HuggingFaceTB/SmolLM2-135M-Instruct': {}, + 'HuggingFaceTB/SmolLM2-360M-Instruct': {}, + 'HuggingFaceTB/SmolLM2-1.7B-Instruct': {}, + 'meta-llama/Llama-3.2-1B-Instruct': {}, + 'meta-llama/Llama-3.2-3B-Instruct': {}, + 'cognitivecomputations/Dolphin3.0-Llama3.2-1B': {}, + 'cognitivecomputations/Dolphin3.0-Llama3.2-3B': {}, + 'mradermacher/Llama-3.2-1B-Instruct-Uncensored-i1-GGUF': { + 'repo': 'meta-llama/Llama-3.2-1B-Instruct', # original repo so we can load missing components + 'type': 'llama', # required so gguf loader knows what to do + 'gguf': 'mradermacher/Llama-3.2-1B-Instruct-Uncensored-i1-GGUF', # gguf repo + 'file': 'Llama-3.2-1B-Instruct-Uncensored.i1-Q4_0.gguf', # gguf file inside repo + }, + } + default = list(models)[0] + supported = list(transformers.integrations.ggml.GGUF_CONFIG_MAPPING) system_prompt: str = 'You are a helpful assistant. You will be given a prompt and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing, and any distinctive details (e.g., colors, textures, styles). Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. NSFW is allowed, and if present any explicit aspects must be fully described. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix.' max_tokens: int = 50 do_sample: bool = True @@ -34,6 +43,7 @@ class Script(scripts.Script): model: str = None llm: transformers.AutoModelForCausalLM = None tokenizer: transformers.AutoProcessor = None + busy: bool = False options = Options() def title(self): @@ -42,31 +52,61 @@ class Script(scripts.Script): def show(self, _is_img2img): return scripts.AlwaysVisible - def load(self, model:str=None): - model = model or self.options.default - if self.model is None or self.model != model: - t0 = time.time() - from modules import modelloader, model_quant - modelloader.hf_login() - quant_args = model_quant.create_config(module='LLM') + def load(self, name:str=None, model_repo:str=None, model_gguf:str=None, model_type:str=None, model_file:str=None): + name = name or self.options.default + if self.busy: + shared.log.debug('Prompt enhance: busy') + return + self.busy = True + if self.model is not None and self.model == name: + return + + t0 = time.time() + from modules import modelloader, model_quant, ggml + modelloader.hf_login() + model_repo = model_repo or self.options.models.get(name, {}).get('repo', None) or name + model_gguf = model_gguf or self.options.models.get(name, {}).get('gguf', None) or model_repo + model_type = model_type or self.options.models.get(name, {}).get('type', None) + model_file = model_file or self.options.models.get(name, {}).get('file', None) + + gguf_args = {} + if model_type is not None and model_file is not None and len(model_type) > 2 and len(model_file) > 2: + if model_type not in self.options.supported: + shared.log.error(f'Prompt enhance: name="{name}" repo="{model_repo}" fn="{model_file}" type={model_type} gguf not supported') + shared.log.trace(f'Prompt enhance: supported={self.options.supported}') + self.busy = False + return + ggml.install_gguf() + gguf_args['model_type'] = model_type + gguf_args['gguf_file'] = model_file + + quant_args = model_quant.create_config(module='LLM') if not gguf_args else {} + + try: + self.model = None self.llm = None self.llm = transformers.AutoModelForCausalLM.from_pretrained( - model, + pretrained_model_name_or_path=model_repo if not gguf_args else model_gguf, trust_remote_code=True, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir, _attn_implementation="eager", + **gguf_args, **quant_args, ) self.llm.eval() self.tokenizer = transformers.AutoTokenizer.from_pretrained( - model, + pretrained_model_name_or_path=model_repo, cache_dir=shared.opts.hfcache_dir, ) - self.model = model - devices.torch_gc() - t1 = time.time() - shared.log.debug(f'Prompt enhance: model="{model}" cls={self.llm.__class__.__name__} time={t1-t0:.2f} loaded') + self.model = name + except Exception as e: + shared.log.error(f'Prompt enhance: load {e}') + errors.display(e, 'Prompt enhance') + devices.torch_gc() + t1 = time.time() + shared.log.debug(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') + self.busy = False def unload(self): if self.llm is not None: @@ -99,6 +139,8 @@ class Script(scripts.Script): penalty = penalty or self.options.repetition_penalty temperature = temperature or self.options.temperature sample = sample if sample is not None else self.options.do_sample + while self.busy: + time.sleep(0.1) self.load(model) if self.llm is None: shared.log.error('Prompt enhance: model not loaded') @@ -108,6 +150,7 @@ class Script(scripts.Script): { "role": "user", "content": prompt }, ] t0 = time.time() + self.busy = True try: inputs = self.tokenizer.apply_chat_template( chat_template, @@ -119,6 +162,8 @@ class Script(scripts.Script): input_len = inputs['input_ids'].shape[1] except Exception as e: shared.log.error(f'Prompt enhance tokenize: {e}') + errors.display(e, 'Prompt enhance') + self.busy = False return prompt try: with devices.inference_context(): @@ -136,12 +181,19 @@ class Script(scripts.Script): # raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) # shared.log.trace(f'Prompt enhance: raw="{raw_response}"') outputs = outputs[:, input_len:] - response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) + response = self.tokenizer.batch_decode( + outputs, + skip_special_tokens=True, + clean_up_tokenization_spaces=True, + ) except Exception as e: shared.log.error(f'Prompt enhance generate: {e}') + errors.display(e, 'Prompt enhance') + self.busy = False response = self.clean(response) t1 = time.time() shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt="{response}"') + self.busy = False return response def apply(self, prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty): @@ -158,6 +210,13 @@ class Script(scripts.Script): return [response, response] return [response, gr.update()] + def get_custom(self, name): + model_repo = self.options.models.get(name, {}).get('repo', None) or name + model_gguf = self.options.models.get(name, {}).get('gguf', None) + model_type = self.options.models.get(name, {}).get('type', None) + model_file = self.options.models.get(name, {}).get('file', None) + return [model_repo, model_gguf, model_type, model_file] + def ui(self, _is_img2img): with gr.Accordion('Prompt enhance', open=False, elem_id='prompt_enhance'): with gr.Row(): @@ -165,29 +224,48 @@ class Script(scripts.Script): with gr.Row(): apply_prompt = gr.Checkbox(label='Apply to prompt', value=False) apply_auto = gr.Checkbox(label='Auto enhance', value=False) + gr.HTML('
') with gr.Group(): with gr.Row(): - llm_model = gr.Dropdown(label='LLM model', choices=self.options.models, value=self.options.default, interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') + llm_model = gr.Dropdown(label='LLM model', choices=list(self.options.models), value=self.options.default, interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') with gr.Row(): load_btn = gr.Button(value='Load model', elem_id='prompt_enhance_load', variant='secondary') load_btn.click(fn=self.load, inputs=[llm_model], outputs=[]) unload_btn = gr.Button(value='Unload model', elem_id='prompt_enhance_unload', variant='secondary') unload_btn.click(fn=self.unload, inputs=[], outputs=[]) - with gr.Row(): - prompt_system = gr.Textbox(label='System prompt', value=self.options.system_prompt, interactive=True, lines=4, elem_id='prompt_enhance_system') - with gr.Row(): - max_tokens = gr.Slider(label='Max tokens', value=self.options.max_tokens, minimum=10, maximum=1024, step=1, interactive=True) - do_sample = gr.Checkbox(label='Do sample', value=self.options.do_sample, interactive=True) - with gr.Row(): - temperature = gr.Slider(label='Temperature', value=self.options.temperature, minimum=0.0, maximum=1.0, step=0.01, interactive=True) - repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) - with gr.Row(): - prompt_output = gr.Textbox(label='Output', value='', interactive=True, lines=4) - with gr.Row(): - clear_btn = gr.Button(value='Clear', elem_id='prompt_enhance_clear', variant='secondary') - clear_btn.click(fn=lambda: '', inputs=[], outputs=[prompt_output]) - copy_btn = gr.Button(value='Set prompt', elem_id='prompt_enhance_copy', variant='secondary') - copy_btn.click(fn=lambda x: x, inputs=[prompt_output], outputs=[self.prompt]) + with gr.Accordion('Custom model', open=False, elem_id='prompt_enhance_custom'): + with gr.Row(): + model_repo = gr.Textbox(label='Model repo', value=None, interactive=True, elem_id='prompt_enhance_model_repo', placeholder='Original model repo on huggingface') + with gr.Row(): + model_gguf = gr.Textbox(label='Model gguf', value=None, interactive=True, elem_id='prompt_enhance_model_gguf', placeholder='Optional GGUF model repo on huggingface') + with gr.Row(): + model_type = gr.Textbox(label='Model type', value=None, interactive=True, elem_id='prompt_enhance_model_type', placeholder='Optional GGUF model type') + with gr.Row(): + model_file = gr.Textbox(label='Model file', value=None, interactive=True, elem_id='prompt_enhance_model_file', placeholder='Optional GGUF model file inside GGUF model repo') + with gr.Row(): + custom_btn = gr.Button(value='Load custom model', elem_id='prompt_enhance_custom_load', variant='secondary') + custom_btn.click(fn=self.load, inputs=[model_file, model_repo, model_gguf, model_type, model_file], outputs=[]) + llm_model.change(fn=self.get_custom, inputs=[llm_model], outputs=[model_repo, model_gguf, model_type, model_file]) + gr.HTML('
') + with gr.Accordion('Options', open=False, elem_id='prompt_enhance_options'): + with gr.Row(): + max_tokens = gr.Slider(label='Max tokens', value=self.options.max_tokens, minimum=10, maximum=1024, step=1, interactive=True) + do_sample = gr.Checkbox(label='Do sample', value=self.options.do_sample, interactive=True) + with gr.Row(): + temperature = gr.Slider(label='Temperature', value=self.options.temperature, minimum=0.0, maximum=1.0, step=0.01, interactive=True) + repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) + gr.HTML('
') + with gr.Accordion('Input', open=False, elem_id='prompt_enhance_system_prompt'): + with gr.Row(): + prompt_system = gr.Textbox(label='System prompt', value=self.options.system_prompt, interactive=True, lines=4, elem_id='prompt_enhance_system') + with gr.Accordion('Output', open=True, elem_id='prompt_enhance_system_prompt'): + with gr.Row(): + prompt_output = gr.Textbox(label='Enhanced prompt', value='', interactive=True, lines=4) + with gr.Row(): + clear_btn = gr.Button(value='Clear', elem_id='prompt_enhance_clear', variant='secondary') + clear_btn.click(fn=lambda: '', inputs=[], outputs=[prompt_output]) + copy_btn = gr.Button(value='Set prompt', elem_id='prompt_enhance_copy', variant='secondary') + copy_btn.click(fn=lambda x: x, inputs=[prompt_output], outputs=[self.prompt]) apply_btn.click(fn=self.apply, inputs=[self.prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty], outputs=[prompt_output, self.prompt]) return [apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty] From 02f480bf32441f75db819b486f5be549299d82a7 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 09:55:56 -0400 Subject: [PATCH 083/122] add prompt enhance wiki Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 3 +++ modules/lora/extra_networks_lora.py | 3 ++- scripts/prompt_enhance.py | 1 + wiki | 2 +- 4 files changed, 7 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 87ccbf2ee..b029616ea 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -68,12 +68,15 @@ Plus... load using *settings -> text encoder* *tip*: add *sd_text_encoder* to your *settings -> user interface -> quicksettings* list to have it appear at the top of the ui - **Prompt Enhance** + - see [Prompt Enhance Wiki](https://github.com/vladmandic/sdnext/wiki/Prompt-Enhance) for details! - new built-in extension available in text/image/control tabs - can be used to manually or automatically enhance prompts using LLM - built-in presets for **Gemma-3, Qwen-2.5, Phi-4, Llama-3.2, SmolLM2, Dolphin-3** - support for custom models load any models hosted on huggingface load either model in huggingface format or `gguf` format + *note*: any hf model in `transformers.AutoModelForCausalLM` standard should work + *note*: not all model architecture are supported for `gguf` format - models are auto-downloaded on first use - support quantization and offloading - **Acceleration** diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 526926c9a..a775ac91f 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -6,7 +6,7 @@ from modules.lora import networks, network_overrides from modules import extra_networks, shared, sd_models -debug = os.environ.get('SD_SCRIPT_DEBUG', None) is not None +debug = os.environ.get('SD_LORA_DEBUG', None) is not None debug_log = shared.log.trace if debug else lambda *args, **kwargs: None @@ -26,6 +26,7 @@ def get_stepwise(param, step, steps): # from https://github.com/cheald/sd-webui- if m[1][-1] <= 1.0: step = step / (max_steps - step_offset) if max_steps > 0 else 1.0 v = np.interp(step, m[1], m[0]) + debug_log(f"Network load: type=LoRA step={step} steps={max_steps} v={v}") return v else: return m diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index aa719e978..76c53bc6f 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -22,6 +22,7 @@ class Options: 'meta-llama/Llama-3.2-3B-Instruct': {}, 'cognitivecomputations/Dolphin3.0-Llama3.2-1B': {}, 'cognitivecomputations/Dolphin3.0-Llama3.2-3B': {}, + 'nidum/Nidum-Gemma-3-4B-it-Uncensored': {}, 'mradermacher/Llama-3.2-1B-Instruct-Uncensored-i1-GGUF': { 'repo': 'meta-llama/Llama-3.2-1B-Instruct', # original repo so we can load missing components 'type': 'llama', # required so gguf loader knows what to do diff --git a/wiki b/wiki index b654db617..15afa8e1d 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit b654db61752a4d1925397fd08bb9f576943015a2 +Subproject commit 15afa8e1d865450fa123b5dcdb3a2af2317b65af From ac61fce526b87593c2279c4ab67cd25e240a76ed Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 11:43:22 -0400 Subject: [PATCH 084/122] prompt enhance add censor detection and debugging Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/prompt_parser_diffusers.py | 2 +- modules/sd_modules.py | 73 ++++++++++++++++++++++++++++++ scripts/prompt_enhance.py | 60 +++++++++++++++++------- wiki | 2 +- 5 files changed, 120 insertions(+), 18 deletions(-) create mode 100644 modules/sd_modules.py diff --git a/CHANGELOG.md b/CHANGELOG.md index b029616ea..80851083d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -79,6 +79,7 @@ Plus... *note*: not all model architecture are supported for `gguf` format - models are auto-downloaded on first use - support quantization and offloading + - debug using `SD_LLM_DEBUG=true` env variable - **Acceleration** - Support for most DiT-based models, for example: *FLUX.1, SD35, Hunyuan, Mochi, Latte, Allegro, Cog* - Enable and configure in *Settings -> Pipeline modifiers* diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index fc6b2af52..4aaf49a39 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -10,7 +10,7 @@ from modules import shared, prompt_parser, devices, sd_models from modules.prompt_parser_xhinker import get_weighted_text_embeddings_sd15, get_weighted_text_embeddings_sdxl_2p, get_weighted_text_embeddings_sd3, get_weighted_text_embeddings_flux1 debug_enabled = os.environ.get('SD_PROMPT_DEBUG', None) -debug = shared.log.trace if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None +debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None debug('Trace: PROMPT') orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access token_dict = None # used by helper get_tokens diff --git a/modules/sd_modules.py b/modules/sd_modules.py new file mode 100644 index 000000000..9a619a2f6 --- /dev/null +++ b/modules/sd_modules.py @@ -0,0 +1,73 @@ +from dataclasses import dataclass +import inspect +import torch + + +@dataclass +class ModuleStats: + module: str + cls: str + params: float + size: float + quant: str + dtype: str + + def __init__(self, module: str, cls: str, params: float, size: float, quant: str, dtype: str): + self.module = module + self.cls = cls + self.params = params + self.size = size + self.quant = quant + self.dtype = dtype + + def __str__(self): + return f'module="{self.module}" cls={self.cls} params={self.params:.3f} size={self.size:.3f} quant={self.quant} dtype={self.dtype}' + + +def get_signature(cls): + signature = inspect.signature(cls.__init__, follow_wrapped=True) + return signature.parameters + + +def get_module_stats(name, module): + if not isinstance(module, torch.nn.Module): + return + try: + module_size = sum(p.numel() * p.element_size() for p in module.parameters(recurse=True)) / 1024 / 1024 / 1024 + param_num = sum(p.numel() for p in module.parameters(recurse=True)) / 1024 / 1024 / 1024 + except Exception: + module_size = 0 + param_num = 0 + cls = module.__class__.__name__ + quant = getattr(module, "quantization_method", None) + module_stats = ModuleStats(name, cls, param_num, module_size, quant, module.dtype) + return module_stats + + +def get_model_stats(model, exclude=None): + # from transformers import Gemma3ForCausalLM + modules = [] + + if isinstance(model, torch.nn.Module): + module_stats = get_module_stats(model.__class__.__name__, model) + if module_stats is not None: + modules.append(module_stats) + return modules + + if hasattr(model, "_internal_dict"): + modules_names = model._internal_dict.keys() # pylint: disable=protected-access + else: + modules_names = get_signature(model).keys() + + if modules_names is None or not isinstance(modules_names, list) or len(modules_names) == 0: + return modules + + modules_names = [m for m in modules_names if m is not None and m not in exclude and not m.startswith('_')] + for module_name in modules_names: + module = getattr(model, module_name, None) + if module is not None: + module_stats = get_module_stats(module_name, module) + if module_stats is not None: + modules.append(module_stats) + + return modules diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 76c53bc6f..284b9613e 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -1,9 +1,14 @@ from dataclasses import dataclass +import os import re import time import gradio as gr import transformers -from modules import scripts, shared, devices, errors, processing, sd_models +from modules import scripts, shared, devices, errors, processing, sd_models, sd_modules + + +debug_enabled = os.environ.get('SD_LLM_DEBUG', None) +debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None @dataclass @@ -32,7 +37,9 @@ class Options: } default = list(models)[0] supported = list(transformers.integrations.ggml.GGUF_CONFIG_MAPPING) - system_prompt: str = 'You are a helpful assistant. You will be given a prompt and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing, and any distinctive details (e.g., colors, textures, styles). Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. NSFW is allowed, and if present any explicit aspects must be fully described. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix.' + system_prompt: str = 'You are a helpful assistant. You will be given a prompt and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing, and any distinctive details (e.g., colors, textures, styles). Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. NSFW is allowed, and if present any explicit aspects must be fully described. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix and as a simple text without formatting or numbering.' + censored = ["i cannot", "i can't", "i am sorry", "against my programming", "i am not able", "i am unable", 'i am not allowed'] + max_tokens: int = 50 do_sample: bool = True temperature: float = 0.5 @@ -72,9 +79,11 @@ class Script(scripts.Script): gguf_args = {} if model_type is not None and model_file is not None and len(model_type) > 2 and len(model_file) > 2: + if debug: + shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') if model_type not in self.options.supported: shared.log.error(f'Prompt enhance: name="{name}" repo="{model_repo}" fn="{model_file}" type={model_type} gguf not supported') - shared.log.trace(f'Prompt enhance: supported={self.options.supported}') + shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') self.busy = False return ggml.install_gguf() @@ -100,6 +109,10 @@ class Script(scripts.Script): pretrained_model_name_or_path=model_repo, cache_dir=shared.opts.hfcache_dir, ) + if debug: + modules = sd_modules.get_model_stats(self.llm) + sd_modules.get_model_stats(self.tokenizer) + for m in modules: + shared.log.trace(f'Prompt enhance: {m}') self.model = name except Exception as e: shared.log.error(f'Prompt enhance: load {e}') @@ -109,6 +122,10 @@ class Script(scripts.Script): shared.log.debug(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') self.busy = False + def censored(self, response): + text = response.lower().replace("i'm", "i am") + return any(c.lower() in text for c in self.options.censored) + def unload(self): if self.llm is not None: sd_models.move_model(self.llm, devices.cpu) @@ -119,16 +136,14 @@ class Script(scripts.Script): shared.log.debug('Prompt enhance: model unloaded') def clean(self, response): - if isinstance(response, list): - response = response[0] response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n') response = re.sub(r'<.*?>', '', response) - if 'prompt:' in response: - response = response.split('prompt:')[1] - if 'Prompt:' in response: - response = response.split('Prompt:')[1] + if response.startswith('Prompt'): + response = response.split('Prompt', maxsplit=2)[1] + if ':' in response: + response = response.split(':', maxsplit=2)[1] if '---' in response: - response = response.split('---')[0] + response = response.split('---', maxsplit=2)[0] response = response.strip() return response @@ -179,11 +194,12 @@ class Script(scripts.Script): if shared.opts.diffusers_offload_mode != 'none': sd_models.move_model(self.llm, devices.cpu) devices.torch_gc() - # raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) - # shared.log.trace(f'Prompt enhance: raw="{raw_response}"') - outputs = outputs[:, input_len:] + if debug: + raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) + shared.log.trace(f'Prompt enhance: raw="{raw_response}"') + outputs_cropped = outputs[:, input_len:] response = self.tokenizer.batch_decode( - outputs, + outputs_cropped, skip_special_tokens=True, clean_up_tokenization_spaces=True, ) @@ -191,10 +207,22 @@ class Script(scripts.Script): shared.log.error(f'Prompt enhance generate: {e}') errors.display(e, 'Prompt enhance') self.busy = False - response = self.clean(response) + response = f'Error: {str(e)}' t1 = time.time() - shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt="{response}"') + + if isinstance(response, list): + response = response[0] + is_censored = self.censored(response) + if not is_censored: + response = self.clean(response) + shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') + if debug: + shared.log.trace(f'Prompt enhance: prompt="{prompt}"') + shared.log.trace(f'Prompt enhance: response="{response}"') self.busy = False + if is_censored: + shared.log.warning(f'Prompt enhance: censored response="{response}"') + return prompt return response def apply(self, prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty): diff --git a/wiki b/wiki index 15afa8e1d..9aff8cd69 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 15afa8e1d865450fa123b5dcdb3a2af2317b65af +Subproject commit 9aff8cd69b01570bd7fd2d52b0f9da6baec9b3be From 9beae12112d4e0e65ee0fa1d20db23417c2553f3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 11:44:12 -0400 Subject: [PATCH 085/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 80851083d..40840a2a2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -79,6 +79,7 @@ Plus... *note*: not all model architecture are supported for `gguf` format - models are auto-downloaded on first use - support quantization and offloading + - auto-detect censored output - debug using `SD_LLM_DEBUG=true` env variable - **Acceleration** - Support for most DiT-based models, for example: *FLUX.1, SD35, Hunyuan, Mochi, Latte, Allegro, Cog* From 2cd89378bd870cd0a5f691d230cf8adfc73cc6bb Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 12:54:48 -0400 Subject: [PATCH 086/122] styles change settings on-the-fly Signed-off-by: Vladimir Mandic --- modules/styles.py | 1 + 1 file changed, 1 insertion(+) diff --git a/modules/styles.py b/modules/styles.py index 41812222c..ecc3204f7 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -260,6 +260,7 @@ class StyleDatabase: for future in concurrent.futures.as_completed(future_items): future.result() + self.built_in = shared.opts.extra_networks_styles list_folder(self.path) t1 = time.time() shared.log.info(f'Available Styles: path="{self.path}" items={len(self.styles.keys())} time={t1-t0:.2f}') From 68d440af5e09e5e8995c4bb31c01a83b9d6b170e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 13:28:14 -0400 Subject: [PATCH 087/122] optimize default system prompt Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 284b9613e..f45123060 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -37,12 +37,12 @@ class Options: } default = list(models)[0] supported = list(transformers.integrations.ggml.GGUF_CONFIG_MAPPING) - system_prompt: str = 'You are a helpful assistant. You will be given a prompt and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing, and any distinctive details (e.g., colors, textures, styles). Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. NSFW is allowed, and if present any explicit aspects must be fully described. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix and as a simple text without formatting or numbering.' + system_prompt: str = 'You are a helpful assistant. You will be given a prompt used to create an image and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness based on original prompt. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing (if any), and any distinctive details (e.g., colors, textures, styles). NSFW and nudity is allowed, and if present it must be fully described. Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix and as a simple text without formatting or numbering.' censored = ["i cannot", "i can't", "i am sorry", "against my programming", "i am not able", "i am unable", 'i am not allowed'] max_tokens: int = 50 do_sample: bool = True - temperature: float = 0.5 + temperature: float = 0.15 repetition_penalty: float = 1.2 @@ -119,7 +119,7 @@ class Script(scripts.Script): errors.display(e, 'Prompt enhance') devices.torch_gc() t1 = time.time() - shared.log.debug(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') + shared.log.info(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') self.busy = False def censored(self, response): @@ -215,8 +215,9 @@ class Script(scripts.Script): is_censored = self.censored(response) if not is_censored: response = self.clean(response) - shared.log.debug(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') + shared.log.info(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') if debug: + shared.log.trace(f'Prompt enhance: sample={sample} tokens={tokens} temperature={temperature} penalty={penalty}') shared.log.trace(f'Prompt enhance: prompt="{prompt}"') shared.log.trace(f'Prompt enhance: response="{response}"') self.busy = False From 1137297e3ab3061408da78db8784d331d5ab93a4 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 29 Mar 2025 13:43:53 -0400 Subject: [PATCH 088/122] fix latte1-t2v Signed-off-by: Vladimir Mandic --- TODO.md | 2 +- installer.py | 2 +- modules/ui_video.py | 14 ++++++++++---- modules/video_models/video_overrides.py | 8 +++++--- 4 files changed, 17 insertions(+), 9 deletions(-) diff --git a/TODO.md b/TODO.md index 3674fe5ef..58dae8a02 100644 --- a/TODO.md +++ b/TODO.md @@ -7,9 +7,9 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ### Issues/Limitations - Video: Hunyuan Video I2V: requires `transformers==4.47.1` -- Video: Latte 1 T2V: dtype mismatch - Video: CogVideoX 1.5 5B T2V/I2V: all-gray output - Video: Allegro T2V: all-gray output +- Video: Latte1 T2V: garbage output ## Future Candidates diff --git a/installer.py b/installer.py index d6308f519..4f384f2b8 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git or args.experimental: return - sha = '617c208bb4cc68fe4518164fee7cbdf5aa44ff78' # diffusers commit hash + sha = '75d7e5cc459f66a53652445d5b281054b297680d' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/ui_video.py b/modules/ui_video.py index d67bdc5eb..b53810d4b 100644 --- a/modules/ui_video.py +++ b/modules/ui_video.py @@ -14,6 +14,14 @@ def engine_change(engine): return gr.update(choices=found, value=found[0] if len(found) > 0 else None) +def get_selected(engine, model): + found = [model.name for model in models_def.models.get(engine, [])] + if len(models_def.models[engine]) > 0 and len(found) > 0: + selected = [m for m in models_def.models[engine] if m.name == model][0] + return selected + return None + + def model_change(engine, model): debug(f'Video change: engine="{engine}" model="{model}"') found = [model.name for model in models_def.models.get(engine, [])] @@ -23,8 +31,7 @@ def model_change(engine, model): def model_load(engine, model): debug(f'Video load: engine="{engine}" model="{model}"') - found = [model.name for model in models_def.models.get(engine, [])] - selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + selected = get_selected(engine, model) yield f'Video model loading: {selected.name}' if selected: if 'None' in selected.name: @@ -43,8 +50,7 @@ def model_load(engine, model): def run_video(*args): engine, model = args[2], args[3] debug(f'Video run: engine="{engine}" model="{model}"') - found = [model.name for model in models_def.models.get(engine, [])] - selected = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None + selected = get_selected(engine, model) if not selected or engine is None or model is None or engine == 'None' or model == 'None': return video_utils.queue_err('model not selected') debug(f'Video run: {str(selected)}') diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index 6a6619ebc..4168a1441 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -27,17 +27,19 @@ def set_overrides(p: processing.StableDiffusionProcessingVideo, selected: Model) shared.sd_model.vae.enable_tiling() # Latte if selected.name == 'Latte 1 T2V': - p.task_args['enable_temporal_attentions'] = False - p.task_args['video_length'] = p.frames + p.task_args['enable_temporal_attentions'] = True + p.task_args['video_length'] = 16 * (max(p.frames // 16, 1)) # LTX if cls == 'LTXImageToVideoPipeline' or cls == 'LTXConditionPipeline': p.task_args['generator'] = None if cls == 'LTXConditionPipeline': p.task_args['strength'] = p.denoising_strength + # WAN if 'Wan' in cls: p.task_args['width'] = 16 * (p.width // 16) p.task_args['height'] = 16 * (p.height // 16) - p.frames = 4 * (p.frames // 4) + 1 + p.frames = 4 * (max(p.frames // 4, 1)) + 1 + # LTX if 'LTX' in cls: p.task_args['width'] = 32 * (p.width // 32) p.task_args['height'] = 32 * (p.height // 32) From a467e23d729f071409af6034339177163727e3cd Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 30 Mar 2025 15:04:17 -0400 Subject: [PATCH 089/122] full ui-settings refactor Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 10 +- extensions-builtin/sd-extension-system-info | 2 +- javascript/black-teal.css | 19 +- javascript/sdnext.css | 33 +- javascript/settings.js | 11 +- modules/img2img.py | 3 +- modules/sd_modules.py | 2 +- modules/shared.py | 6 +- modules/txt2img.py | 2 +- modules/ui.py | 417 ++------------------ modules/ui_extensions.py | 167 ++++---- modules/ui_settings.py | 371 +++++++++++++++++ wiki | 2 +- 13 files changed, 526 insertions(+), 519 deletions(-) create mode 100644 modules/ui_settings.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 40840a2a2..9b836a13e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2025-03-28 +## Update for 2025-03-30 -### Highlights for 2025-03-28 +### Highlights for 2025-03-30 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! @@ -15,7 +15,7 @@ Plus... - More quantization options and granular control - Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods -### Details for 2025-03-28 +### Details for 2025-03-30 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -129,6 +129,10 @@ Plus... - **LoRA** enable memory cache by default - **Samplers** add ability to set sigma adjustment for each sampler - **ModernUI** updates + - **CSS** updates + - settings vertiocal/dirty indicator restores to default setting instead to previous value + - video interpolate do not skip duplicate frames + - **settings UI** full refactor - **Wiki/Docs** - updated [Models](https://github.com/vladmandic/sdnext/wiki/Models) info - new [Video](https://github.com/vladmandic/sdnext/wiki/Video) guide diff --git a/extensions-builtin/sd-extension-system-info b/extensions-builtin/sd-extension-system-info index 8c7edb3be..ce373b9c2 160000 --- a/extensions-builtin/sd-extension-system-info +++ b/extensions-builtin/sd-extension-system-info @@ -1 +1 @@ -Subproject commit 8c7edb3be11b8b8c2d2dcd0421e93345bd20fcae +Subproject commit ce373b9c27544f56ad73a1f7fe2c5530a89c1c32 diff --git a/javascript/black-teal.css b/javascript/black-teal.css index 851c03953..8a99cb91c 100644 --- a/javascript/black-teal.css +++ b/javascript/black-teal.css @@ -54,15 +54,21 @@ --line-md: 1.4em; --line-lg: 1.5em; --range-shadow: - -20em 0 0 0 hsl(180, 54%, 2%), -19em 0 0 0 hsl(180, 54%, 5%), -18em 0 0 0 hsl(180, 54%, 0%), -17em 0 0 0 hsl(180, 54%, 11%), - -16em 0 0 0 hsl(180, 54%, 14%), -15em 0 0 0 hsl(180, 54%, 17%), -14em 0 0 0 hsl(180, 54%, 20%), -13em 0 0 0 hsl(180, 54%, 23%), + -32em 0 0 0 hsl(180, 54%, 6%), -31em 0 0 0 hsl(180, 54%, 7%), -30em 0 0 0 hsl(180, 54%, 8%), -29em 0 0 0 hsl(180, 54%, 9%), + -28em 0 0 0 hsl(180, 54%, 10%), -27em 0 0 0 hsl(180, 54%, 11%), -26em 0 0 0 hsl(180, 54%, 12%), -25em 0 0 0 hsl(180, 54%, 13%), + -24em 0 0 0 hsl(180, 54%, 14%), -23em 0 0 0 hsl(180, 54%, 15%), -22em 0 0 0 hsl(180, 54%, 16%), -21em 0 0 0 hsl(180, 54%, 17%), + -20em 0 0 0 hsl(180, 54%, 18%), -19em 0 0 0 hsl(180, 54%, 19%), -18em 0 0 0 hsl(180, 54%, 20%), -17em 0 0 0 hsl(180, 54%, 21%), + -16em 0 0 0 hsl(180, 54%, 22%), -15em 0 0 0 hsl(180, 54%, 23%), -14em 0 0 0 hsl(180, 54%, 24%), -13em 0 0 0 hsl(180, 54%, 25%), -12em 0 0 0 hsl(180, 54%, 26%), -11em 0 0 0 hsl(180, 54%, 29%), -10em 0 0 0 hsl(180, 54%, 32%), -9em 0 0 0 hsl(180, 54%, 35%), -8em 0 0 0 hsl(180, 54%, 38%), -7em 0 0 0 hsl(180, 54%, 41%), -6em 0 0 0 hsl(180, 54%, 44%), -5em 0 0 0 hsl(180, 54%, 47%), -4em 0 0 0 hsl(180, 54%, 50%), -3em 0 0 0 hsl(180, 54%, 53%), -2em 0 0 0 hsl(180, 54%, 56%), -1em 0 0 0 hsl(180, 54%, 59%), 1em 0 0 0 var(--neutral-800), 2em 0 0 0 var(--neutral-800), 3em 0 0 0 var(--neutral-800), 4em 0 0 0 var(--neutral-800), 5em 0 0 0 var(--neutral-800), 6em 0 0 0 var(--neutral-800), 7em 0 0 0 var(--neutral-800), 8em 0 0 0 var(--neutral-800), 9em 0 0 0 var(--neutral-800), 10em 0 0 0 var(--neutral-800), 11em 0 0 0 var(--neutral-800), 12em 0 0 0 var(--neutral-800), - 13em 0 0 0 var(--neutral-800), 14em 0 0 0 var(--neutral-800), 15em 0 0 0 var(--neutral-800), 16em 0 0 0 var(--neutral-800); + 13em 0 0 0 var(--neutral-800), 14em 0 0 0 var(--neutral-800), 15em 0 0 0 var(--neutral-800), 16em 0 0 0 var(--neutral-800), + 17em 0 0 0 var(--neutral-800), 18em 0 0 0 var(--neutral-800), 19em 0 0 0 var(--neutral-800), 20em 0 0 0 var(--neutral-800), + 21em 0 0 0 var(--neutral-800), 22em 0 0 0 var(--neutral-800), 23em 0 0 0 var(--neutral-800), 24em 0 0 0 var(--neutral-800), + 25em 0 0 0 var(--neutral-800), 26em 0 0 0 var(--neutral-800), 27em 0 0 0 var(--neutral-800), 28em 0 0 0 var(--neutral-800); } html { font-size: var(--font-size); font-family: var(--font); } @@ -70,13 +76,6 @@ body, button, input, select, textarea { font-family: var(--font); } button { max-width: 400px; white-space: nowrap; } img { background-color: var(--background-color); } -/* -input[type=range] { height: var(--line-xs) !important; appearance: none !important; margin-top: 0 !important; min-width: max(4em, 100%) !important; background-color: var(--background-color) !important; width: 100% !important; background: transparent !important; } -input[type=range]::-webkit-slider-runnable-track { width: 100% !important; height: 6px !important; cursor: pointer !important; background: var(--input-background-fill) !important; border-radius: var(--radius-lg) !important; border: 0px solid var(--neutral-900) !important; } -input[type=range]::-moz-range-track { width: 100% !important; height: 6px !important; cursor: pointer !important; background: var(--input-background-fill) !important; border-radius: var(--radius-lg) !important; border: 0px solid var(--neutral-900) !important; } -input[type=range]::-webkit-slider-thumb { border: 0px solid #000000 !important; height: var(--line-xs) !important; width: var(--line-md) !important; border-radius: var(--radius-lg) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px !important; } -input[type=range]::-moz-range-thumb { border: 0px solid #000000 !important; height: var(--line-xs) !important; width: var(--line-md) !important; border-radius: var(--radius-lg) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px !important; } -*/ input[type='range'] { display: block; margin: 0; padding: 0; height: 1em; background-color: transparent; overflow: hidden; cursor: pointer; box-shadow: 0 0 0 0 transparent; -webkit-appearance: none; appearance: none; } input[type='range']::-webkit-slider-thumb { height: .9em; width: .9em; background-color: hsl(180, 54%, 61%); box-shadow: var(--range-shadow); border-radius: var(--radius-xs); } input[type='range']::-webkit-slider-runnable-track, input[type='range']::-webkit-slider-thumb { -webkit-appearance: none; } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 71f8cb43e..f5e1a949d 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -154,25 +154,24 @@ div#extras_scale_to_tab div.form { flex-direction: row; } #si-sparkline-memo, #si-sparkline-load { background-color: #111; } #quicksettings { width: fit-content; } #quicksettings>button { padding: 0 1em 0 0; align-self: end; margin-bottom: 6px; } -#settings { display: flex; gap: var(--layout-gap); } -#settings div { border: none; gap: 0; margin: 0 0 var(--layout-gap) 0px; padding: 0; } -#settings>div.tab-content { flex: 10 0 75%; display: grid; } -#settings>div.tab-content>div { border: none; padding: 0; } + +#settings { display: flex; margin-left: 0.5em; } +#settings>div.tab-content { margin-top: 1em; } +#settings>div.tab-content>div>div { gap: 0; } #settings>div.tab-content>div>div>div>div>div { flex-direction: unset; } -#settings>div.tab-nav { display: grid; grid-template-columns: repeat(auto-fill, .5em minmax(10em, 1fr)); flex: 1 0 auto; width: 12em; align-self: flex-start; gap: 8px; } -#settings>div.tab-nav button { display: block; border: none; text-align: left; white-space: initial; padding: 0; } -#settings>div.tab-nav>#settings_show_all_pages { padding: var(--size-2) var(--size-4); } +#settings>div.tab-nav { width: 14em; display: block; background: var(--neutral-900); border-radius: var(--block-radius); margin-right: 1em;} +#settings>div.tab-nav button { width: 100%; height: 2em; text-align: left; border: none; border-radius: var(--block-radius); } +#settings .dirtyable.hidden { visibility: hidden; } +#settings .modification-indicator { background: none; border-radius: var(--radius-lg); padding: 0; width: 4px !important; height: 2em !important; position: absolute; float: left; left: -6px; } +#settings .modification-indicator:disabled { background: none; } +#settings .modification-indicator.saved { background: var(--color-accent-soft); } +#settings .modification-indicator.changed { background: var(--color-accent); } +#settings .modification-indicator.changed.unsaved { background: var(--color-warning); } #settings .block.gradio-checkbox { margin: 0; width: auto; } -#settings .dirtyable { gap: .5em; } -#settings .dirtyable.hidden { display: none; } -#settings .modification-indicator { height: 1.2em; border-radius: 1em !important; padding: 0; width: 0; margin-right: 0.5em; border-left: inset; } -#settings .modification-indicator:disabled { visibility: hidden; } -#settings .modification-indicator.saved { background: var(--color-accent-soft); width: var(--spacing-sm); } -#settings .modification-indicator.changed { background: var(--color-accent); width: var(--spacing-sm); } -#settings .modification-indicator.changed.unsaved { background-image: linear-gradient(var(--color-accent) 25%, var(--color-accent-soft) 75%); width: var(--spacing-sm); } -#settings_result { margin: 0 1.2em; } -#tab_settings .gradio-slider, #tab_settings .gradio-dropdown { width: 300px !important; max-width: 300px; } -#tab_settings textarea { max-width: 500px; } +#settings .block.gradio-number { min-width: 500px; } +#settings .gradio-slider, #tab_settings .gradio-dropdown { width: 500px !important; max-width: 500px; } +#settings textarea { width: 500px !important; max-width: 500px; } + .licenses { display: block !important; } /* live preview */ diff --git a/javascript/settings.js b/javascript/settings.js index 1a891c844..7a43969ac 100644 --- a/javascript/settings.js +++ b/javascript/settings.js @@ -49,8 +49,8 @@ async function updateOpts(json_string) { function showAllSettings() { // Try to ensure that the show all settings tab is opened by clicking on its tab button - const tab_dirty_indicator = gradioApp().getElementById('modification_indicator_show_all_pages'); - if (tab_dirty_indicator && tab_dirty_indicator.nextSibling) tab_dirty_indicator.nextSibling.click(); + // const tab_dirty_indicator = gradioApp().getElementById('modification_indicator_show_all_pages'); + // if (tab_dirty_indicator && tab_dirty_indicator.nextSibling) tab_dirty_indicator.nextSibling.click(); getSettingsTabs().forEach((elem) => { if (elem.id === 'settings_tab_licenses' || elem.id === 'settings_show_all_pages') return; elem.style.display = 'block'; @@ -192,9 +192,12 @@ async function initSettings() { tabContentWrapper.className = 'tab-content'; tabNavElements.parentElement.insertBefore(tabContentWrapper, tabNavElements.nextSibling); tabElements.forEach((elem, index) => { - const tabName = elem.id.replace('settings_', ''); + const tabName = elem.id.replace('settings_section_tab_', ''); const indicator = gradioApp().getElementById(`modification_indicator_${tabName}`); - tabNavElements.insertBefore(indicator, tabNavButtons[index]); + if (indicator) { + tabNavElements.insertBefore(document.createElement('br'), tabNavButtons[index]); + tabNavElements.insertBefore(indicator, tabNavButtons[index]); + } tabContentWrapper.appendChild(elem); observer.observe(elem, { attributes: true, attributeFilter: ['style'] }); }); diff --git a/modules/img2img.py b/modules/img2img.py index 75971c608..ca71ff0e7 100644 --- a/modules/img2img.py +++ b/modules/img2img.py @@ -6,9 +6,10 @@ from PIL import Image, ImageOps, ImageFilter, ImageEnhance, ImageChops, Unidenti import modules.scripts from modules import shared, processing, images from modules.generation_parameters_copypaste import create_override_settings_dict -from modules.ui import plaintext_to_html +from modules.ui_common import plaintext_to_html from modules.memstats import memory_stats + debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: PROCESS') diff --git a/modules/sd_modules.py b/modules/sd_modules.py index 9a619a2f6..781287f9a 100644 --- a/modules/sd_modules.py +++ b/modules/sd_modules.py @@ -31,7 +31,7 @@ def get_signature(cls): def get_module_stats(name, module): if not isinstance(module, torch.nn.Module): - return + return None try: module_size = sum(p.numel() * p.element_size() for p in module.parameters(recurse=True)) / 1024 / 1024 / 1024 param_num = sum(p.numel() for p in module.parameters(recurse=True)) / 1024 / 1024 / 1024 diff --git a/modules/shared.py b/modules/shared.py index 5d785f1a5..89eba4655 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -46,7 +46,7 @@ tab_names = [] extra_networks = [] options_templates = {} hypernetworks = {} -settings_components = None +settings_components = {} restricted_opts = { "samples_filename_pattern", "directories_filename_pattern", @@ -635,7 +635,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), { "unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True), "te_dir": OptionInfo(os.path.join(paths.models_path, 'Text-encoder'), "Folder with Text encoder files", folder=True), "lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True), - "styles_dir": OptionInfo(os.path.join(paths.data_path, 'styles.csv'), "File or Folder with user-defined styles", folder=True), + "styles_dir": OptionInfo(os.path.join(paths.models_path, 'styles'), "File or Folder with user-defined styles", folder=True), "wildcards_dir": OptionInfo(os.path.join(paths.models_path, 'wildcards'), "Folder with user-defined wildcards", folder=True), "embeddings_dir": OptionInfo(os.path.join(paths.models_path, 'embeddings'), "Folder with textual inversion embeddings", folder=True), "hypernetwork_dir": OptionInfo(os.path.join(paths.models_path, 'hypernetworks'), "Folder with Hypernetwork models", folder=True), @@ -946,7 +946,7 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "wildcards_enabled": OptionInfo(True, "Enable file wildcards support"), })) -options_templates.update(options_section((None, "Internal options"), { +options_templates.update(options_section((None, "Hidden options"), { "diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}), "disabled_extensions": OptionInfo([], "Disable these extensions"), "sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint"), diff --git a/modules/txt2img.py b/modules/txt2img.py index 08184fe43..151110d07 100644 --- a/modules/txt2img.py +++ b/modules/txt2img.py @@ -1,7 +1,7 @@ import os from modules import shared, processing, scripts from modules.generation_parameters_copypaste import create_override_settings_dict -from modules.ui import plaintext_to_html +from modules.ui_common import plaintext_to_html debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None diff --git a/modules/ui.py b/modules/ui.py index 91fcf3cf3..1c4b7ec0f 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -1,28 +1,17 @@ -import os import mimetypes import gradio as gr import gradio.routes import gradio.utils -from modules.call_queue import wrap_gradio_call, wrap_gradio_gpu_call # pylint: disable=unused-import -from modules import timer, gr_hijack, shared, theme, sd_models, script_callbacks, modelloader, ui_common, ui_loadsave, ui_symbols, ui_javascript, ui_sections, generation_parameters_copypaste, call_queue +from modules import errors, timer, gr_hijack, shared, script_callbacks, ui_common, ui_symbols, ui_javascript, ui_sections, generation_parameters_copypaste, call_queue, scripts from modules.paths import script_path, data_path # pylint: disable=unused-import -from modules.dml import directml_override_opts -from modules.onnx_impl import install_olive -import modules.scripts -import modules.errors -modules.errors.install() +errors.install() mimetypes.init() mimetypes.add_type('application/javascript', '.js') mimetypes.add_type('image/webp', '.webp') mimetypes.add_type('image/jxl', '.jxl') -log = shared.log -opts = shared.opts -cmd_opts = shared.cmd_opts -ui_system_tabs = None -paste_function = None -wrap_queued_call = call_queue.wrap_queued_call +gr_hijack.init() switch_values_symbol = ui_symbols.switch detect_image_size_symbol = ui_symbols.detect paste_symbol = ui_symbols.paste @@ -32,11 +21,16 @@ folder_symbol = ui_symbols.folder extra_networks_symbol = ui_symbols.networks apply_style_symbol = ui_symbols.apply save_style_symbol = ui_symbols.save -gr_hijack.init() +wrap_queued_call = call_queue.wrap_queued_call # compatibility item +wrap_gradio_call = call_queue.wrap_gradio_call # compatibility item +wrap_gradio_gpu_call = call_queue.wrap_gradio_gpu_call # compatibility item +plaintext_to_html = ui_common.plaintext_to_html # compatibility item +infotext_to_html = ui_common.infotext_to_html # compatibility item create_sampler_and_steps_selection = ui_sections.create_sampler_and_steps_selection # compatibility item +ui_system_tabs = None # required for system-info -if not cmd_opts.share and not cmd_opts.listen: +if not shared.cmd_opts.share and not shared.cmd_opts.listen: # fix gradio phoning home gradio.utils.version_check = lambda: None gradio.utils.get_local_ip_address = lambda: '127.0.0.1' @@ -55,14 +49,6 @@ def create_output_panel(tabname, outdir): # pylint: disable=unused-argument # ou return a, b, c, e -def plaintext_to_html(text): # may be referenced by extensions - return ui_common.plaintext_to_html(text) - - -def infotext_to_html(text): # may be referenced by extensions - return ui_common.infotext_to_html(text) - - def send_gradio_gallery_to_image(x): if len(x) == 0: return None @@ -81,43 +67,6 @@ def setup_progressbar(*args, **kwargs): # pylint: disable=unused-argument pass -def apply_setting(key, value): - if value is None: - return gr.update() - if shared.cmd_opts.freeze: - return gr.update() - if key == 'sd_backend': - return gr.update() - if shared.opts.disable_weights_auto_swap and key in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_model_dict', 'sd_vae', 'sd_unet', 'sd_text_encoder']: - return gr.update() - if key == "sd_model_checkpoint": - ckpt_info = sd_models.get_closet_checkpoint_match(value) - if ckpt_info is not None: - value = ckpt_info.title - else: - return gr.update() - comp_args = opts.data_labels[key].component_args - if comp_args and isinstance(comp_args, dict) and comp_args.get('visible') is False: - return gr.update() - valtype = type(opts.data_labels[key].default) - oldval = opts.data.get(key, None) - opts.data[key] = valtype(value) if valtype != type(None) else value - if oldval != value and opts.data_labels[key].onchange is not None: - opts.data_labels[key].onchange() - opts.save(shared.config_filename) - return getattr(opts, key) - - -def get_value_for_setting(key): - value = getattr(opts, key) - info = opts.data_labels[key] - args = info.component_args() if callable(info.component_args) else info.component_args or {} - args = {k: v for k, v in args.items() if k not in {'precision', 'multiselect', 'visible'}} - # if not args: - # return gr.update() - return gr.update(value=value, **args) - - def ordered_ui_categories(): return ['dimensions', 'sampler', 'seed', 'denoising', 'cfg', 'checkboxes', 'accordions', 'override_settings', 'scripts'] # a1111 compatibility item, not implemented @@ -127,6 +76,7 @@ def create_ui(startup_timer = None): timer.startup = timer.Timer() ui_javascript.reload_javascript() generation_parameters_copypaste.reset() + scripts.scripts_current = None with gr.Blocks(analytics_enabled=False) as txt2img_interface: from modules import ui_txt2img @@ -138,8 +88,6 @@ def create_ui(startup_timer = None): ui_img2img.create_ui() timer.startup.record("ui-img2img") - modules.scripts.scripts_current = None - with gr.Blocks(analytics_enabled=False) as control_interface: if shared.native: from modules import ui_control @@ -172,224 +120,13 @@ def create_ui(startup_timer = None): ui_gallery.create_ui() timer.startup.record("ui-gallery") - def create_setting_component(key, is_quicksettings=False): - def fun(): - return opts.data[key] if key in opts.data else opts.data_labels[key].default - - info = opts.data_labels[key] - t = type(info.default) - args = (info.component_args() if callable(info.component_args) else info.component_args) or {} - if info.component is not None: - comp = info.component - elif t == str: - comp = gr.Textbox - elif t == int: - comp = gr.Number - elif t == bool: - comp = gr.Checkbox - else: - raise ValueError(f'bad options item type: {t} for key {key}') - elem_id = f"setting_{key}" - dirty_indicator = None - - if not is_quicksettings: - dirtyable_setting = gr.Group(elem_classes="dirtyable", visible=args.get("visible", True)) - dirtyable_setting.__enter__() - dirty_indicator = gr.Button("", elem_classes="modification-indicator", elem_id="modification_indicator_" + key) - - if info.refresh is not None: - if is_quicksettings: - res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) - ui_common.create_refresh_button(res, info.refresh, info.component_args, f"refresh_{key}") - else: - with gr.Row(): - res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) - ui_common.create_refresh_button(res, info.refresh, info.component_args, f"refresh_{key}") - elif info.folder is not None: - with gr.Row(): - res = comp(label=info.label, value=fun(), elem_id=elem_id, elem_classes="folder-selector", **args) - # ui_common.create_browse_button(res, f"folder_{key}") - else: - try: - res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) - except Exception as e: - log.error(f'Error creating setting: {key} {e}') - res = None - - if res is not None and not is_quicksettings: - res.change(fn=None, inputs=res, _js=f'(val) => markIfModified("{key}", val)') - if dirty_indicator is not None: - dirty_indicator.click(fn=lambda: getattr(opts, key), outputs=res, show_progress=False) - dirtyable_setting.__exit__() - - return res - - def create_dirty_indicator(key, keys_to_reset, **kwargs): - def get_opt_values(): - return [getattr(opts, _key) for _key in keys_to_reset] - - elements_to_reset = [component_dict[_key] for _key in keys_to_reset if component_dict[_key] is not None] - indicator = gr.Button("", elem_classes="modification-indicator", elem_id=f"modification_indicator_{key}", **kwargs) - indicator.click(fn=get_opt_values, outputs=elements_to_reset, show_progress=False) - return indicator - - loadsave = ui_loadsave.UiLoadsave(cmd_opts.ui_config) - components = [] - component_dict = {} - shared.settings_components = component_dict - - script_callbacks.ui_settings_callback() - opts.reorder() - - def run_settings(*args): - changed = [] - for key, value, comp in zip(opts.data_labels.keys(), args, components): - if comp == dummy_component or value=='dummy': - continue - if getattr(comp, 'visible', True) is False: - continue - if not opts.same_type(value, opts.data_labels[key].default): - log.error(f'Setting bad value: {key}={value} expecting={type(opts.data_labels[key].default).__name__}') - continue - if opts.set(key, value): - changed.append(key) - if shared.opts.cuda_compile_backend == "olive-ai": - install_olive() - if cmd_opts.use_directml: - directml_override_opts() - if cmd_opts.use_openvino: - if "Model" not in shared.opts.cuda_compile: - shared.log.warning("OpenVINO: Enabling Torch Compile Model") - shared.opts.cuda_compile.append("Model") - if shared.opts.cuda_compile_backend != "openvino_fx": - shared.log.warning("OpenVINO: Setting Torch Compiler backend to OpenVINO FX") - shared.opts.cuda_compile_backend = "openvino_fx" - if shared.opts.sd_backend != "diffusers": - shared.log.warning("OpenVINO: Setting backend to Diffusers") - shared.opts.sd_backend = "diffusers" - try: - if len(changed) > 0: - opts.save(shared.config_filename) - log.info(f'Settings: changed={len(changed)} {changed}') - except RuntimeError: - log.error(f'Settings failed: change={len(changed)} {changed}') - return opts.dumpjson(), f'{len(changed)} Settings changed without save: {", ".join(changed)}' - return opts.dumpjson(), f'{len(changed)} Settings changed{": " if len(changed) > 0 else ""}{", ".join(changed)}' - - def run_settings_single(value, key, progress=False): - if not opts.same_type(value, opts.data_labels[key].default): - return gr.update(visible=True), opts.dumpjson() - if not opts.set(key, value): - return gr.update(value=getattr(opts, key)), opts.dumpjson() - if key == "cuda_compile_backend" and value == "olive-ai": - install_olive() - if cmd_opts.use_directml: - directml_override_opts() - opts.save(shared.config_filename) - log.debug(f'Setting changed: {key}={value} progress={progress}') - return get_value_for_setting(key), opts.dumpjson() - with gr.Blocks(analytics_enabled=False) as settings_interface: - with gr.Row(elem_id="system_row"): - restart_submit = gr.Button(value="Restart server", variant='primary', elem_id="restart_submit") - shutdown_submit = gr.Button(value="Shutdown server", variant='primary', elem_id="shutdown_submit") - unload_sd_model = gr.Button(value='Unload model', variant='primary', elem_id="sett_unload_sd_model") - reload_sd_model = gr.Button(value='Reload model', variant='primary', elem_id="sett_reload_sd_model") - enable_profiling = gr.Button(value='Start profiling', variant='primary', elem_id="enable_profiling") - - with gr.Tabs(elem_id="system") as system_tabs: - global ui_system_tabs # pylint: disable=global-statement - ui_system_tabs = system_tabs - with gr.TabItem("Settings", id="system_settings", elem_id="tab_settings"): - with gr.Row(elem_id="settings_row"): - settings_submit = gr.Button(value="Apply settings", variant='primary', elem_id="settings_submit") - preview_theme = gr.Button(value="Preview theme", variant='primary', elem_id="settings_preview_theme") - defaults_submit = gr.Button(value="Restore defaults", variant='primary', elem_id="defaults_submit") - with gr.Row(): - _settings_search = gr.Text(label="Search", elem_id="settings_search") - - result = gr.HTML(elem_id="settings_result") - quicksettings_names = opts.quicksettings_list - quicksettings_names = {x: i for i, x in enumerate(quicksettings_names) if x != 'quicksettings'} - quicksettings_list = [] - - previous_section = [] - tab_item_keys = [] - current_tab = None - current_row = None - dummy_component = gr.Textbox(visible=False, value='dummy') - with gr.Tabs(elem_id="settings"): - for i, (k, item) in enumerate(opts.data_labels.items()): - section_must_be_skipped = item.section[0] is None - if previous_section != item.section and not section_must_be_skipped: - if len(item.section) == 2: - elem_id, text = item.section - elif len(item.section) == 3: - _category, elem_id, text = item.section - else: - shared.log.error(f'Settings: section={item.section} invalid') - continue - if current_tab is not None and len(previous_section) > 0: - create_dirty_indicator(previous_section[0], tab_item_keys) - tab_item_keys = [] - current_row.__exit__() - current_tab.__exit__() - current_tab = gr.TabItem(elem_id=f"settings_{elem_id}", label=text) - current_tab.__enter__() - current_row = gr.Column(variant='compact') - current_row.__enter__() - previous_section = item.section - if k in quicksettings_names and not shared.cmd_opts.freeze: - quicksettings_list.append((i, k, item)) - components.append(dummy_component) - elif section_must_be_skipped: - components.append(dummy_component) - else: - component = create_setting_component(k) - component_dict[k] = component - tab_item_keys.append(k) - components.append(component) - if current_tab is not None and len(previous_section) > 0: - create_dirty_indicator(previous_section[0], tab_item_keys) - tab_item_keys = [] - current_row.__exit__() - current_tab.__exit__() - - request_notifications = gr.Button(value='Request browser notifications', elem_id="request_notifications", visible=False) - with gr.TabItem("Show all pages", elem_id="settings_show_all_pages"): - create_dirty_indicator("show_all_pages", [], interactive=False) - - with gr.TabItem("Update", id="system_update", elem_id="tab_update"): - from modules import update - update.create_ui() - - with gr.TabItem("User interface", id="system_config", elem_id="tab_config"): - loadsave.create_ui() - create_dirty_indicator("tab_defaults", [], interactive=False) - - with gr.TabItem("ONNX", id="onnx_config", elem_id="tab_onnx"): - from modules.onnx_impl import ui as ui_onnx - ui_onnx.create_ui() - - def unload_sd_weights(): - modules.sd_models.unload_model_weights(op='model') - modules.sd_models.unload_model_weights(op='refiner') - - def reload_sd_weights(): - modules.sd_models.reload_model_weights(force=True) - - def switch_profiling(): - shared.cmd_opts.profile = not shared.cmd_opts.profile - shared.log.warning(f'Profiling: {shared.cmd_opts.profile}') - return 'Stop profiling' if shared.cmd_opts.profile else 'Start profiling' - - unload_sd_model.click(fn=unload_sd_weights, inputs=[], outputs=[]) - reload_sd_model.click(fn=reload_sd_weights, inputs=[], outputs=[]) - enable_profiling.click(fn=switch_profiling, inputs=[], outputs=[enable_profiling]) - request_notifications.click(fn=lambda: None, inputs=[], outputs=[], _js='function(){}') - preview_theme.click(fn=None, _js='previewTheme', inputs=[], outputs=[]) - - timer.startup.record("ui-settings") + from modules import ui_settings + ui_settings.create_ui() + global ui_system_tabs # pylint: disable=global-statement + ui_system_tabs = ui_settings.ui_system_tabs + shared.opts.reorder() + timer.startup.record("ui-extensions") with gr.Blocks(analytics_enabled=False) as info_interface: with gr.Tabs(elem_id="tabs_info"): @@ -401,6 +138,11 @@ def create_ui(startup_timer = None): from modules import ui_docs ui_docs.create_ui_wiki() + with gr.Blocks(analytics_enabled=False) as extensions_interface: + from modules import ui_extensions + ui_extensions.create_ui() + timer.startup.record("ui-extensions") + interfaces = [] interfaces += [(txt2img_interface, "Text", "txt2img")] interfaces += [(img2img_interface, "Image", "img2img")] @@ -415,123 +157,12 @@ def create_ui(startup_timer = None): interfaces += script_callbacks.ui_tabs_callback() interfaces += [(settings_interface, "System", "system")] interfaces += [(info_interface, "Info", "info")] - - from modules import ui_extensions - extensions_interface = ui_extensions.create_ui() interfaces += [(extensions_interface, "Extensions", "extensions")] - timer.startup.record("ui-extensions") + + ui_app = ui_settings.create_quicksettings(interfaces) shared.tab_names = [] for _interface, label, _ifid in interfaces: shared.tab_names.append(label) - with gr.Blocks(theme=theme.gradio_theme, analytics_enabled=False, title="SD.Next") as ui_app: - with gr.Row(elem_id="quicksettings", variant="compact"): - for _i, k, _item in sorted(quicksettings_list, key=lambda x: quicksettings_names.get(x[1], x[0])): - component = create_setting_component(k, is_quicksettings=True) - component_dict[k] = component - - generation_parameters_copypaste.connect_paste_params_buttons() - - with gr.Tabs(elem_id="tabs") as tabs: - for interface, label, ifid in interfaces: - if interface is None: - continue - # if label in shared.opts.hidden_tabs or label == '': - # continue - with gr.TabItem(label, id=ifid, elem_id=f"tab_{ifid}"): - # log.debug(f'UI render: id={ifid}') - interface.render() - for interface, _label, ifid in interfaces: - if interface is None: - continue - if ifid in ["extensions", "system"]: - continue - loadsave.add_block(interface, ifid) - loadsave.add_component(f"webui/Tabs@{tabs.elem_id}", tabs) - loadsave.setup_ui() - - if opts.notification_audio_enable and os.path.exists(os.path.join(script_path, opts.notification_audio_path)): - gr.Audio(interactive=False, value=os.path.join(script_path, opts.notification_audio_path), elem_id="audio_notification", visible=False) - - text_settings = gr.Textbox(elem_id="settings_json", value=lambda: opts.dumpjson(), visible=False) - components = [c for c in components if c is not None] - settings_submit.click( - fn=wrap_gradio_call(run_settings, extra_outputs=[gr.update()]), - inputs=components, - outputs=[text_settings, result], - ) - defaults_submit.click(fn=lambda: shared.restore_defaults(restart=True), _js="restartReload") - restart_submit.click(fn=lambda: shared.restart_server(restart=True), _js="restartReload") - shutdown_submit.click(fn=lambda: shared.restart_server(restart=False), _js="restartReload") - - for _i, k, _item in quicksettings_list: - component = component_dict[k] - info = opts.data_labels[k] - if isinstance(component, gr.components.Textbox): - change_handlers = [component.blur, component.submit] - else: - change_handlers = [component.release if hasattr(component, 'release') else component.change] - for change_handler in change_handlers: - change_handler( - fn=lambda value, k=k, progress=info.refresh is not None: run_settings_single(value, key=k, progress=progress), - inputs=[component], - outputs=[component, text_settings], - show_progress=info.refresh is not None, - ) - - dummy_component = gr.Textbox(visible=False, value='dummy') - button_set_checkpoint = gr.Button('Change model', elem_id='change_checkpoint', visible=False) - button_set_checkpoint.click( - fn=lambda value, _: run_settings_single(value, key='sd_model_checkpoint'), - _js="function(v){ var res = desiredCheckpointName; desiredCheckpointName = ''; return [res || v, null]; }", - inputs=[component_dict['sd_model_checkpoint'], dummy_component], - outputs=[component_dict['sd_model_checkpoint'], text_settings], - ) - button_set_refiner = gr.Button('Change refiner', elem_id='change_refiner', visible=False) - button_set_refiner.click( - fn=lambda value, _: run_settings_single(value, key='sd_model_checkpoint'), - _js="function(v){ var res = desiredCheckpointName; desiredCheckpointName = ''; return [res || v, null]; }", - inputs=[component_dict['sd_model_refiner'], dummy_component], - outputs=[component_dict['sd_model_refiner'], text_settings], - ) - button_set_vae = gr.Button('Change VAE', elem_id='change_vae', visible=False) - button_set_vae.click( - fn=lambda value, _: run_settings_single(value, key='sd_vae'), - _js="function(v){ var res = desiredVAEName; desiredVAEName = ''; return [res || v, null]; }", - inputs=[component_dict['sd_vae'], dummy_component], - outputs=[component_dict['sd_vae'], text_settings], - ) - - def reference_submit(model): - if '@' not in model: # diffusers - loaded = modelloader.load_reference(model) - return model if loaded else opts.sd_model_checkpoint - else: # civitai - model, url = model.split('@') - loaded = modelloader.load_civitai(model, url) - return loaded if loaded is not None else opts.sd_model_checkpoint - - button_set_reference = gr.Button('Change reference', elem_id='change_reference', visible=False) - button_set_reference.click( - fn=reference_submit, - _js="function(v){ return desiredCheckpointName; }", - inputs=[component_dict['sd_model_checkpoint']], - outputs=[component_dict['sd_model_checkpoint']], - ) - component_keys = [k for k in opts.data_labels.keys() if k in component_dict] - - def get_settings_values(): - return [get_value_for_setting(key) for key in component_keys] - - ui_app.load( - fn=get_settings_values, - inputs=[], - outputs=[component_dict[k] for k in component_keys if component_dict[k] is not None], - queue=False, - ) - - timer.startup.record("ui-defaults") - loadsave.dump_defaults() - ui_app.ui_loadsave = loadsave return ui_app diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index 2b75df0e5..ef40092ec 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -5,7 +5,8 @@ import errno import html from datetime import datetime, timedelta import gradio as gr -from modules import extensions, shared, paths, errors, ui_symbols +from modules import extensions, shared, paths, errors, ui_symbols, call_queue + debug = shared.log.debug if os.environ.get('SD_EXT_DEBUG', None) is not None else lambda *args, **kwargs: None extensions_index = "https://vladmandic.github.io/sd-data/pages/extensions.json" @@ -437,86 +438,84 @@ def create_html(search_text, sort_column): def create_ui(): import modules.ui - with gr.Blocks(analytics_enabled=False) as ui: - extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "user", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all", visible=False) - extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False, container=False) - extensions_update_list = gr.Text(elem_id="extensions_update_list", visible=False, container=False) - with gr.Tabs(elem_id="tabs_extensions"): - with gr.TabItem("Manage extensions", id="manage"): - with gr.Row(elem_id="extensions_installed_top"): - extension_to_install = gr.Text(elem_id="extension_to_install", visible=False) - install_extension_button = gr.Button(elem_id="install_extension_button", visible=False) - uninstall_extension_button = gr.Button(elem_id="uninstall_extension_button", visible=False) - update_extension_button = gr.Button(elem_id="update_extension_button", visible=False) - with gr.Column(scale=4): - search_text = gr.Text(label="Search") - with gr.Column(scale=1): - sort_column = gr.Dropdown(value="default", label="Sort by", choices=list(sort_ordering.keys()), multiselect=False) - with gr.Column(scale=1): - refresh_extensions_button = gr.Button(value="Refresh extension list", variant="primary") - check = gr.Button(value="Update all installed", variant="primary") - apply = gr.Button(value="Apply changes", variant="primary") - list_extensions() - gr.HTML(''' -

Extension list

- - Refesh extension list to download latest list with status
- - Check status of an extension by looking at status icon before installing it
- - After any operation such as install/uninstall or enable/disable, please restart the server
-
''') - gr.HTML('') - info = gr.HTML('') - extensions_table = gr.HTML(create_html(search_text.value, sort_column.value)) - check.click( - fn=modules.ui.wrap_gradio_call(check_updates, extra_outputs=[gr.update()]), - _js="extensions_check", - inputs=[info, extensions_disabled_list, search_text, sort_column], - outputs=[extensions_table, info], - ) - apply.click( - fn=apply_changes, - _js="extensions_apply", - inputs=[extensions_disabled_list, extensions_update_list, extensions_disable_all], - outputs=[], - ) - refresh_extensions_button.click( - fn=modules.ui.wrap_gradio_call(refresh_extensions_list, extra_outputs=[gr.update(), gr.update()]), - inputs=[search_text, sort_column], - outputs=[extensions_table, info], - ) - install_extension_button.click( - fn=modules.ui.wrap_gradio_call(install_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), - inputs=[extension_to_install, search_text, sort_column], - outputs=[extensions_table, info], - ) - uninstall_extension_button.click( - fn=modules.ui.wrap_gradio_call(uninstall_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), - inputs=[extension_to_install, search_text, sort_column], - outputs=[extensions_table, info], - ) - update_extension_button.click( - fn=modules.ui.wrap_gradio_call(update_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), - inputs=[extension_to_install, search_text, sort_column], - outputs=[extensions_table, info], - ) - search_text.change( - fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]), - inputs=[search_text, sort_column], - outputs=[extensions_table, info], - ) - sort_column.change( - fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]), - inputs=[search_text, sort_column], - outputs=[extensions_table, info], - ) - with gr.TabItem("Manual install", id="install_from_url"): - install_url = gr.Text(label="Extension GIT repository URL") - install_branch = gr.Text(label="Specific branch name", placeholder="Leave empty for default main branch") - install_dirname = gr.Text(label="Local directory name", placeholder="Leave empty for auto") - install_button = gr.Button(value="Install", variant="primary") - info = gr.HTML(elem_id="extension_info") - install_button.click( - fn=modules.ui.wrap_gradio_call(install_extension_from_url, extra_outputs=[gr.update()]), - inputs=[install_dirname, install_url, install_branch, search_text, sort_column], - outputs=[extensions_table, info], - ) - return ui + extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "user", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all", visible=False) + extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False, container=False) + extensions_update_list = gr.Text(elem_id="extensions_update_list", visible=False, container=False) + with gr.Tabs(elem_id="tabs_extensions"): + with gr.TabItem("Manage extensions", id="manage"): + with gr.Row(elem_id="extensions_installed_top"): + extension_to_install = gr.Text(elem_id="extension_to_install", visible=False) + install_extension_button = gr.Button(elem_id="install_extension_button", visible=False) + uninstall_extension_button = gr.Button(elem_id="uninstall_extension_button", visible=False) + update_extension_button = gr.Button(elem_id="update_extension_button", visible=False) + with gr.Column(scale=4): + search_text = gr.Text(label="Search") + with gr.Column(scale=1): + sort_column = gr.Dropdown(value="default", label="Sort by", choices=list(sort_ordering.keys()), multiselect=False) + with gr.Column(scale=1): + refresh_extensions_button = gr.Button(value="Refresh extension list", variant="primary") + check = gr.Button(value="Update all installed", variant="primary") + apply = gr.Button(value="Apply changes", variant="primary") + list_extensions() + gr.HTML(''' +

Extension list

+ - Refesh extension list to download latest list with status
+ - Check status of an extension by looking at status icon before installing it
+ - After any operation such as install/uninstall or enable/disable, please restart the server
+
''') + gr.HTML('') + info = gr.HTML('') + extensions_table = gr.HTML(create_html(search_text.value, sort_column.value)) + check.click( + fn=call_queue.wrap_gradio_call(check_updates, extra_outputs=[gr.update()]), + _js="extensions_check", + inputs=[info, extensions_disabled_list, search_text, sort_column], + outputs=[extensions_table, info], + ) + apply.click( + fn=apply_changes, + _js="extensions_apply", + inputs=[extensions_disabled_list, extensions_update_list, extensions_disable_all], + outputs=[], + ) + refresh_extensions_button.click( + fn=call_queue.wrap_gradio_call(refresh_extensions_list, extra_outputs=[gr.update(), gr.update()]), + inputs=[search_text, sort_column], + outputs=[extensions_table, info], + ) + install_extension_button.click( + fn=call_queue.wrap_gradio_call(install_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), + inputs=[extension_to_install, search_text, sort_column], + outputs=[extensions_table, info], + ) + uninstall_extension_button.click( + fn=call_queue.wrap_gradio_call(uninstall_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), + inputs=[extension_to_install, search_text, sort_column], + outputs=[extensions_table, info], + ) + update_extension_button.click( + fn=call_queue.wrap_gradio_call(update_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]), + inputs=[extension_to_install, search_text, sort_column], + outputs=[extensions_table, info], + ) + search_text.change( + fn=call_queue.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]), + inputs=[search_text, sort_column], + outputs=[extensions_table, info], + ) + sort_column.change( + fn=call_queue.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]), + inputs=[search_text, sort_column], + outputs=[extensions_table, info], + ) + with gr.TabItem("Manual install", id="install_from_url"): + install_url = gr.Text(label="Extension GIT repository URL") + install_branch = gr.Text(label="Specific branch name", placeholder="Leave empty for default main branch") + install_dirname = gr.Text(label="Local directory name", placeholder="Leave empty for auto") + install_button = gr.Button(value="Install", variant="primary") + info = gr.HTML(elem_id="extension_info") + install_button.click( + fn=call_queue.wrap_gradio_call(install_extension_from_url, extra_outputs=[gr.update()]), + inputs=[install_dirname, install_url, install_branch, search_text, sort_column], + outputs=[extensions_table, info], + ) diff --git a/modules/ui_settings.py b/modules/ui_settings.py new file mode 100644 index 000000000..0a4f3b1eb --- /dev/null +++ b/modules/ui_settings.py @@ -0,0 +1,371 @@ +import os +import gradio as gr +from modules import timer, shared, paths, theme, sd_models, modelloader, ui_common, ui_loadsave, generation_parameters_copypaste, call_queue, script_callbacks + + +ui_system_tabs = None # required for system-info +dummy_component = gr.Textbox(visible=False, value='dummy') +text_settings = gr.Textbox(elem_id="settings_json", value=lambda: shared.opts.dumpjson(), visible=False) +loadsave = ui_loadsave.UiLoadsave(shared.cmd_opts.ui_config) +quicksettings_names = {x: i for i, x in enumerate(shared.opts.quicksettings_list) if x != 'quicksettings'} +quicksettings_list = [] +components = [] + + +def apply_setting(key, value): + if value is None: + return gr.update() + if shared.cmd_opts.freeze: + return gr.update() + if key == 'sd_backend': + return gr.update() + if shared.opts.disable_weights_auto_swap and key in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_model_dict', 'sd_vae', 'sd_unet', 'sd_text_encoder']: + return gr.update() + if key == "sd_model_checkpoint": + ckpt_info = sd_models.get_closet_checkpoint_match(value) + if ckpt_info is not None: + value = ckpt_info.title + else: + return gr.update() + comp_args = shared.opts.data_labels[key].component_args + if comp_args and isinstance(comp_args, dict) and comp_args.get('visible') is False: + return gr.update() + valtype = type(shared.opts.data_labels[key].default) + oldval = shared.opts.data.get(key, None) + shared.opts.data[key] = valtype(value) if valtype != type(None) else value + if oldval != value and shared.opts.data_labels[key].onchange is not None: + shared.opts.data_labels[key].onchange() + shared.opts.save(shared.config_filename) + return getattr(shared.opts, key) + + +def get_value_for_setting(key): + value = getattr(shared.opts, key) + info = shared.opts.data_labels[key] + args = info.component_args() if callable(info.component_args) else info.component_args or {} + args = {k: v for k, v in args.items() if k not in {'precision', 'multiselect', 'visible'}} + return gr.update(value=value, **args) + + +def ordered_ui_categories(): + return ['dimensions', 'sampler', 'seed', 'denoising', 'cfg', 'checkboxes', 'accordions', 'override_settings', 'scripts'] # a1111 compatibility item, not implemented + + +def create_setting_component(key, is_quicksettings=False): + def fun(): + return shared.opts.data[key] if key in shared.opts.data else shared.opts.data_labels[key].default + + info = shared.opts.data_labels[key] + t = type(info.default) + args = (info.component_args() if callable(info.component_args) else info.component_args) or {} + if info.component is not None: + comp = info.component + elif t == str: + comp = gr.Textbox + elif t == int: + comp = gr.Number + elif t == bool: + comp = gr.Checkbox + else: + raise ValueError(f'bad options item type: {t} for key {key}') + elem_id = f"setting_{key}" + dirty_indicator = None + + if not is_quicksettings: + dirtyable_setting = gr.Group(elem_classes="dirtyable", visible=args.get("visible", True)) + dirtyable_setting.__enter__() + dirty_indicator = gr.Button("", elem_classes="modification-indicator", elem_id=f"modification_indicator_{key}") + + if info.refresh is not None: + if is_quicksettings: + res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) + ui_common.create_refresh_button(res, info.refresh, info.component_args, f"refresh_{key}") + else: + with gr.Row(): + res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) + ui_common.create_refresh_button(res, info.refresh, info.component_args, f"refresh_{key}") + elif info.folder is not None: + with gr.Row(): + res = comp(label=info.label, value=fun(), elem_id=elem_id, elem_classes="folder-selector", **args) + # ui_common.create_browse_button(res, f"folder_{key}") + else: + try: + res = comp(label=info.label, value=fun(), elem_id=elem_id, **args) + except Exception as e: + shared.log.error(f'Error creating setting: {key} {e}') + res = None + + if res is not None and not is_quicksettings: + res.change(fn=None, inputs=res, _js=f'(val) => markIfModified("{key}", val)') + if dirty_indicator is not None: + dirty_indicator.click(fn=lambda: shared.opts.get_default(key), outputs=[res], show_progress=False) + dirtyable_setting.__exit__() + + return res + +def create_dirty_indicator(key, keys_to_reset, **kwargs): + def get_default_values(): + values = [shared.opts.get_default(key) for key in keys_to_reset] + shared.log.debug(f'Settings restore: section={key} keys={keys_to_reset} values={values}') + return values + + elements_to_reset = [shared.settings_components[_key] for _key in keys_to_reset if shared.settings_components[_key] is not None] + indicator = gr.Button('', elem_classes="modification-indicator", elem_id=f"modification_indicator_{key}", **kwargs) + indicator.click(fn=get_default_values, outputs=elements_to_reset, show_progress=True) + return indicator + + +def run_settings(*args): + changed = [] + for key, value, comp in zip(shared.opts.data_labels.keys(), args, components): + if comp == dummy_component or value=='dummy': + continue + if getattr(comp, 'visible', True) is False: + continue + if not shared.opts.same_type(value, shared.opts.data_labels[key].default): + shared.log.error(f'Setting bad value: {key}={value} expecting={type(shared.opts.data_labels[key].default).__name__}') + continue + if shared.opts.set(key, value): + changed.append(key) + if shared.opts.cuda_compile_backend == "olive-ai": + from modules.onnx_impl import install_olive + install_olive() + if shared.cmd_opts.use_directml: + from modules.dml import directml_override_opts + directml_override_opts() + if shared.cmd_opts.use_openvino: + if "Model" not in shared.opts.cuda_compile: + shared.log.warning("OpenVINO: Enabling Torch Compile Model") + shared.opts.cuda_compile.append("Model") + if shared.opts.cuda_compile_backend != "openvino_fx": + shared.log.warning("OpenVINO: Setting Torch Compiler backend to OpenVINO FX") + shared.opts.cuda_compile_backend = "openvino_fx" + if shared.opts.sd_backend != "diffusers": + shared.log.warning("OpenVINO: Setting backend to Diffusers") + shared.opts.sd_backend = "diffusers" + try: + if len(changed) > 0: + shared.opts.save(shared.config_filename) + shared.log.info(f'Settings: changed={len(changed)} {changed}') + except RuntimeError: + shared.log.error(f'Settings failed: change={len(changed)} {changed}') + return shared.opts.dumpjson(), f'{len(changed)} Settings changed without save: {", ".join(changed)}' + return shared.opts.dumpjson(), f'{len(changed)} Settings changed{": " if len(changed) > 0 else ""}{", ".join(changed)}' + +def run_settings_single(value, key, progress=False): + if not shared.opts.same_type(value, shared.opts.data_labels[key].default): + return gr.update(visible=True), shared.opts.dumpjson() + if not shared.opts.set(key, value): + return gr.update(value=getattr(shared.opts, key)), shared.opts.dumpjson() + if key == "cuda_compile_backend" and value == "olive-ai": + from modules.onnx_impl import install_olive + install_olive() + if shared.cmd_opts.use_directml: + from modules.dml import directml_override_opts + directml_override_opts() + shared.opts.save(shared.config_filename) + shared.log.debug(f'Setting changed: {key}={value} progress={progress}') + return get_value_for_setting(key), shared.opts.dumpjson() + + +def create_ui(): + with gr.Row(elem_id="system_row"): + restart_submit = gr.Button(value="Restart server", variant='primary', elem_id="restart_submit") + shutdown_submit = gr.Button(value="Shutdown server", variant='primary', elem_id="shutdown_submit") + unload_sd_model = gr.Button(value='Unload model', variant='primary', elem_id="sett_unload_sd_model") + reload_sd_model = gr.Button(value='Reload model', variant='primary', elem_id="sett_reload_sd_model") + enable_profiling = gr.Button(value='Start profiling', variant='primary', elem_id="enable_profiling") + + with gr.Tabs(elem_id="system") as system_tabs: + global ui_system_tabs # pylint: disable=global-statement + ui_system_tabs = system_tabs + with gr.TabItem("Settings", id="system_settings", elem_id="tab_settings"): + with gr.Row(elem_id="settings_row"): + settings_submit = gr.Button(value="Apply settings", variant='primary', elem_id="settings_submit") + preview_theme = gr.Button(value="Preview theme", variant='primary', elem_id="settings_preview_theme") + defaults_submit = gr.Button(value="Restore defaults", variant='primary', elem_id="defaults_submit") + with gr.Row(): + _settings_search = gr.Text(label="Search", elem_id="settings_search") + + result = gr.HTML(elem_id="settings_result") + script_callbacks.ui_settings_callback() # let extensions create settings + sections = [] + for item in shared.opts.data_labels.values(): # get unique sections from all items + if len(item.section) == 2: + section_id, section_text = item.section + elif len(item.section) == 3: # compatibility item with a1111 extensions + _category, section_id, section_text = item.section + item.section = section_id, section_text + else: + section_id = None + item.section = None, 'Hidden' + if (section_id, section_text) not in sections: + sections.append((section_id, section_text)) + + with gr.Tabs(elem_id="settings"): + for (section_id, section_text) in sections: + items = [item for item in shared.opts.data_labels.items() if item[1].section[0] == section_id] # find all items in this section + hidden = section_id is None or 'hidden' in section_id.lower() or 'hidden' in section_text.lower() + shared.log.trace(f'Settings: section="{section_id}" title="{section_text}" items={len(items)} hidden={hidden}') + if hidden: + components.append(dummy_component) + else: + with gr.TabItem(elem_id=f"settings_section_tab_{section_id}", label=section_text): + current_items = [] + for (key, item) in items: + if key in quicksettings_names: + quicksettings_list.append((key, item)) + components.append(dummy_component) # TODO: quicksettings should clone insetad of move + else: + with gr.Row(elem_id=f"settings_section_row_{section_id}"): # only so we can add dirty indicator at the start of the row + component = create_setting_component(key) + shared.settings_components[key] = component + current_items.append(key) + components.append(component) + create_dirty_indicator(section_id, current_items) + + with gr.TabItem("Show all pages", elem_id="settings_show_all_pages"): + create_dirty_indicator("show_all_pages", []) + request_notifications = gr.Button(value='Request browser notifications', elem_id="request_notifications", visible=False) + + + with gr.TabItem("Update", id="system_update", elem_id="tab_update"): + from modules import update + update.create_ui() + + with gr.TabItem("User interface", id="system_config", elem_id="tab_config"): + loadsave.create_ui() + create_dirty_indicator("tab_defaults", [], interactive=False) + + with gr.TabItem("ONNX", id="onnx_config", elem_id="tab_onnx"): + from modules.onnx_impl import ui as ui_onnx + ui_onnx.create_ui() + + def unload_sd_weights(): + sd_models.unload_model_weights(op='model') + sd_models.unload_model_weights(op='refiner') + + def reload_sd_weights(): + sd_models.reload_model_weights(force=True) + + def switch_profiling(): + shared.cmd_opts.profile = not shared.cmd_opts.profile + shared.log.warning(f'Profiling: {shared.cmd_opts.profile}') + return 'Stop profiling' if shared.cmd_opts.profile else 'Start profiling' + + unload_sd_model.click(fn=unload_sd_weights, inputs=[], outputs=[]) + reload_sd_model.click(fn=reload_sd_weights, inputs=[], outputs=[]) + enable_profiling.click(fn=switch_profiling, inputs=[], outputs=[enable_profiling]) + request_notifications.click(fn=lambda: None, inputs=[], outputs=[], _js='function(){}') + preview_theme.click(fn=None, _js='previewTheme', inputs=[], outputs=[]) + settings_submit.click( + fn=call_queue.wrap_gradio_call(run_settings, extra_outputs=[gr.update()]), + inputs=components, + outputs=[text_settings, result], + ) + defaults_submit.click(fn=lambda: shared.restore_defaults(restart=True), _js="restartReload") + restart_submit.click(fn=lambda: shared.restart_server(restart=True), _js="restartReload") + shutdown_submit.click(fn=lambda: shared.restart_server(restart=False), _js="restartReload") + + +def create_quicksettings(interfaces): + shared.tab_names = [] + for _interface, label, _ifid in interfaces: + shared.tab_names.append(label) + + with gr.Blocks(theme=theme.gradio_theme, analytics_enabled=False, title="SD.Next") as ui_app: + with gr.Row(elem_id="quicksettings", variant="compact"): + for k, _item in sorted(quicksettings_list, key=lambda x: quicksettings_names.get(x[1], x[0])): + component = create_setting_component(k, is_quicksettings=True) + shared.settings_components[k] = component + + generation_parameters_copypaste.connect_paste_params_buttons() + + with gr.Tabs(elem_id="tabs") as tabs: + for interface, label, ifid in interfaces: + if interface is None: + continue + with gr.TabItem(label, id=ifid, elem_id=f"tab_{ifid}"): + interface.render() + for interface, _label, ifid in interfaces: + if interface is None: + continue + if ifid in ["extensions", "system"]: + continue + loadsave.add_block(interface, ifid) + loadsave.add_component(f"webui/Tabs@{tabs.elem_id}", tabs) + loadsave.setup_ui() + + if shared.opts.notification_audio_enable and os.path.exists(os.path.join(paths.script_path, shared.opts.notification_audio_path)): + gr.Audio(interactive=False, value=os.path.join(paths.script_path, shared.opts.notification_audio_path), elem_id="audio_notification", visible=False) + + for k, _item in quicksettings_list: + component = shared.settings_components[k] + info = shared.opts.data_labels[k] + if isinstance(component, gr.components.Textbox): + change_handlers = [component.blur, component.submit] + else: + change_handlers = [component.release if hasattr(component, 'release') else component.change] + for change_handler in change_handlers: + change_handler( + fn=lambda value, k=k, progress=info.refresh is not None: run_settings_single(value, key=k, progress=progress), + inputs=[component], + outputs=[component, text_settings], + show_progress=info.refresh is not None, + ) + + dummy_component = gr.Textbox(visible=False, value='dummy') + button_set_checkpoint = gr.Button('Change model', elem_id='change_checkpoint', visible=False) + button_set_checkpoint.click( + fn=lambda value, _: run_settings_single(value, key='sd_model_checkpoint'), + _js="function(v){ var res = desiredCheckpointName; desiredCheckpointName = ''; return [res || v, null]; }", + inputs=[shared.settings_components['sd_model_checkpoint'], dummy_component], + outputs=[shared.settings_components['sd_model_checkpoint'], text_settings], + ) + button_set_refiner = gr.Button('Change refiner', elem_id='change_refiner', visible=False) + button_set_refiner.click( + fn=lambda value, _: run_settings_single(value, key='sd_model_checkpoint'), + _js="function(v){ var res = desiredCheckpointName; desiredCheckpointName = ''; return [res || v, null]; }", + inputs=[shared.settings_components['sd_model_refiner'], dummy_component], + outputs=[shared.settings_components['sd_model_refiner'], text_settings], + ) + button_set_vae = gr.Button('Change VAE', elem_id='change_vae', visible=False) + button_set_vae.click( + fn=lambda value, _: run_settings_single(value, key='sd_vae'), + _js="function(v){ var res = desiredVAEName; desiredVAEName = ''; return [res || v, null]; }", + inputs=[shared.settings_components['sd_vae'], dummy_component], + outputs=[shared.settings_components['sd_vae'], text_settings], + ) + + def reference_submit(model): + if '@' not in model: # diffusers + loaded = modelloader.load_reference(model) + return model if loaded else shared.opts.sd_model_checkpoint + else: # civitai + model, url = model.split('@') + loaded = modelloader.load_civitai(model, url) + return loaded if loaded is not None else shared.opts.sd_model_checkpoint + + button_set_reference = gr.Button('Change reference', elem_id='change_reference', visible=False) + button_set_reference.click( + fn=reference_submit, + _js="function(v){ return desiredCheckpointName; }", + inputs=[shared.settings_components['sd_model_checkpoint']], + outputs=[shared.settings_components['sd_model_checkpoint']], + ) + component_keys = [k for k in shared.opts.data_labels.keys() if k in shared.settings_components] + + def get_settings_values(): + return [get_value_for_setting(key) for key in component_keys] + + ui_app.load( + fn=get_settings_values, + inputs=[], + outputs=[shared.settings_components[k] for k in component_keys if shared.settings_components[k] is not None], + queue=False, + ) + + timer.startup.record("ui-defaults") + loadsave.dump_defaults() + ui_app.ui_loadsave = loadsave + return ui_app diff --git a/wiki b/wiki index 9aff8cd69..90e18e0d1 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 9aff8cd69b01570bd7fd2d52b0f9da6baec9b3be +Subproject commit 90e18e0d17ab43cddbe3d8ff9169707d1d289a41 From def571dae96932b48a40cfbb4ff8016d52bdd916 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 30 Mar 2025 15:29:18 -0400 Subject: [PATCH 090/122] fix settings save Signed-off-by: Vladimir Mandic --- modules/ui_settings.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/modules/ui_settings.py b/modules/ui_settings.py index 0a4f3b1eb..06d37c251 100644 --- a/modules/ui_settings.py +++ b/modules/ui_settings.py @@ -9,6 +9,7 @@ text_settings = gr.Textbox(elem_id="settings_json", value=lambda: shared.opts.du loadsave = ui_loadsave.UiLoadsave(shared.cmd_opts.ui_config) quicksettings_names = {x: i for i, x in enumerate(shared.opts.quicksettings_list) if x != 'quicksettings'} quicksettings_list = [] +hidden_list = [] components = [] @@ -118,9 +119,7 @@ def create_dirty_indicator(key, keys_to_reset, **kwargs): def run_settings(*args): changed = [] for key, value, comp in zip(shared.opts.data_labels.keys(), args, components): - if comp == dummy_component or value=='dummy': - continue - if getattr(comp, 'visible', True) is False: + if comp == dummy_component or value=='dummy': # or getattr(comp, 'visible', True) is False or key in hidden_list: continue if not shared.opts.same_type(value, shared.opts.data_labels[key].default): shared.log.error(f'Setting bad value: {key}={value} expecting={type(shared.opts.data_labels[key].default).__name__}') @@ -202,20 +201,23 @@ def create_ui(): if (section_id, section_text) not in sections: sections.append((section_id, section_text)) + shared.log.debug(f'UI settings: sections={len(sections)} settings={len(list(shared.opts.data_labels))}') with gr.Tabs(elem_id="settings"): for (section_id, section_text) in sections: items = [item for item in shared.opts.data_labels.items() if item[1].section[0] == section_id] # find all items in this section hidden = section_id is None or 'hidden' in section_id.lower() or 'hidden' in section_text.lower() - shared.log.trace(f'Settings: section="{section_id}" title="{section_text}" items={len(items)} hidden={hidden}') + # shared.log.trace(f'Settings: section="{section_id}" title="{section_text}" items={len(items)} hidden={hidden}') if hidden: - components.append(dummy_component) + for (key, _item) in items: + hidden_list.append(key) + components.append(dummy_component) else: with gr.TabItem(elem_id=f"settings_section_tab_{section_id}", label=section_text): current_items = [] for (key, item) in items: if key in quicksettings_names: quicksettings_list.append((key, item)) - components.append(dummy_component) # TODO: quicksettings should clone insetad of move + components.append(dummy_component) else: with gr.Row(elem_id=f"settings_section_row_{section_id}"): # only so we can add dirty indicator at the start of the row component = create_setting_component(key) @@ -314,7 +316,6 @@ def create_quicksettings(interfaces): show_progress=info.refresh is not None, ) - dummy_component = gr.Textbox(visible=False, value='dummy') button_set_checkpoint = gr.Button('Change model', elem_id='change_checkpoint', visible=False) button_set_checkpoint.click( fn=lambda value, _: run_settings_single(value, key='sd_model_checkpoint'), From daec94a9e9a082778eddb7a25c464e1c9c9725b6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 30 Mar 2025 15:39:44 -0400 Subject: [PATCH 091/122] settings css improvements Signed-off-by: Vladimir Mandic --- javascript/sdnext.css | 7 ++++--- modules/shared.py | 12 ++++++------ 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/javascript/sdnext.css b/javascript/sdnext.css index f5e1a949d..83b08a64d 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -168,9 +168,10 @@ div#extras_scale_to_tab div.form { flex-direction: row; } #settings .modification-indicator.changed { background: var(--color-accent); } #settings .modification-indicator.changed.unsaved { background: var(--color-warning); } #settings .block.gradio-checkbox { margin: 0; width: auto; } -#settings .block.gradio-number { min-width: 500px; } -#settings .gradio-slider, #tab_settings .gradio-dropdown { width: 500px !important; max-width: 500px; } -#settings textarea { width: 500px !important; max-width: 500px; } +#settings .block.gradio-number { min-width: 500px !important; } +#settings .gradio-slider, #tab_settings .gradio-dropdown { width: 500px !important; max-width: 500px !important; } +#settings .gradio-radio { padding: var(--block-padding) !important; } +#settings textarea { width: 500px !important; max-width: 500px !important; } .licenses { display: block !important; } diff --git a/modules/shared.py b/modules/shared.py index 89eba4655..3829408b9 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -409,7 +409,7 @@ options_templates.update(options_section(('sd', "Models & Loading"), { "diffusers_eval": OptionInfo(True, "Force model eval", gr.Checkbox, {"visible": False }), "diffusers_to_gpu": OptionInfo(False, "Load model directly to GPU"), "disable_accelerate": OptionInfo(False, "Disable accelerate", gr.Checkbox, {"visible": False }), - "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints), + "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles(), "visible": False}, refresh=refresh_checkpoints), "sd_checkpoint_cache": OptionInfo(0, "Cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": not native }), })) @@ -422,7 +422,7 @@ options_templates.update(options_section(('vae_encoder', "Variable Auto Encoder" "diffusers_vae_tile_size": OptionInfo(0, "VAE tile size", gr.Slider, {"minimum": 0, "maximum": 4096, "step": 8 }), "diffusers_vae_tile_overlap": OptionInfo(0.25, "VAE tile overlap", gr.Slider, {"minimum": 0, "maximum": 0.95, "step": 0.05 }), "sd_vae_sliced_encode": OptionInfo(False, "VAE sliced encode", gr.Checkbox, {"visible": not native}), - "nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox), + "nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox, {"visible": False}), "remote_vae_type": OptionInfo('raw', "Remote VAE image type", gr.Dropdown, {"choices": ['raw', 'jpg', 'png']}), "remote_vae_encode": OptionInfo(False, "Remote VAE for encode"), "rollback_vae": OptionInfo(False, "Attempt VAE roll back for NaN values", gr.Checkbox, {"visible": not native}), @@ -623,7 +623,6 @@ options_templates.update(options_section(('compile', "Model Compile"), { })) options_templates.update(options_section(('system-paths', "System Paths"), { - "clean_temp_dir_at_start": OptionInfo(True, "Cleanup temporary folder on startup"), "models_paths_sep_options": OptionInfo("

Models Paths

", "", gr.HTML), "models_dir": OptionInfo('models', "Root model folder", folder=True), "model_paths_sep_options": OptionInfo("

Paths for specific models

", "", gr.HTML), @@ -651,6 +650,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), { "ldsr_models_path": OptionInfo(os.path.join(paths.models_path, 'LDSR'), "Folder with LDSR models", folder=True), "clip_models_path": OptionInfo(os.path.join(paths.models_path, 'CLIP'), "Folder with CLIP models", folder=True), "other_paths_sep_options": OptionInfo("

Cache folders

", "", gr.HTML), + "clean_temp_dir_at_start": OptionInfo(True, "Cleanup temporary folder on startup"), "temp_dir": OptionInfo("", "Directory for temporary images; leave empty for default", folder=True), "accelerate_offload_path": OptionInfo('cache/accelerate', "Folder for disk offload", folder=True), "openvino_cache_path": OptionInfo('cache', "Folder for OpenVINO cache", folder=True), @@ -659,15 +659,15 @@ options_templates.update(options_section(('system-paths', "System Paths"), { })) options_templates.update(options_section(('saving-images', "Image Options"), { - "keep_incomplete": OptionInfo(True, "Keep incomplete images"), "samples_save": OptionInfo(True, "Save all generated images"), + "keep_incomplete": OptionInfo(False, "Keep incomplete images"), "samples_format": OptionInfo('jpg', 'File format', gr.Dropdown, {"choices": ["jpg", "png", "webp", "tiff", "jp2", "jxl"]}), "jpeg_quality": OptionInfo(90, "Image quality", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1}), "img_max_size_mp": OptionInfo(1000, "Maximum image size (MP)", gr.Slider, {"minimum": 100, "maximum": 2000, "step": 1}), "webp_lossless": OptionInfo(False, "WebP lossless compression"), - "save_selected_only": OptionInfo(True, "Save only saves selected image"), + "save_selected_only": OptionInfo(True, "UI save only saves selected image"), "include_mask": OptionInfo(False, "Include mask in outputs"), - "samples_save_zip": OptionInfo(True, "Create ZIP archive"), + "samples_save_zip": OptionInfo(False, "Create ZIP archive for multiple images"), "image_background": OptionInfo("#000000", "Resize background color", gr.ColorPicker, {}), "image_sep_metadata": OptionInfo("

Metadata/Logging

", "", gr.HTML), From 68b019c3ed50e03c02f2f9d5b85b9535642cfb88 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 09:25:13 -0400 Subject: [PATCH 092/122] update settings for modernui Signed-off-by: Vladimir Mandic --- extensions-builtin/sdnext-modernui | 2 +- javascript/sdnext.css | 4 +--- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 1e38e9b56..ed8291c74 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 1e38e9b56edf45dd17402aee2a1c281dc26b5286 +Subproject commit ed8291c74faea33c51b25df379f48704e36bef37 diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 83b08a64d..5ffc03d36 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -151,10 +151,8 @@ div#extras_scale_to_tab div.form { flex-direction: row; } #txt2img_styles_refresh, #img2img_styles_refresh, #control_styles_refresh, #video_styles_refresh { padding: 0; margin-top: 1em; } /* settings */ -#si-sparkline-memo, #si-sparkline-load { background-color: #111; } #quicksettings { width: fit-content; } #quicksettings>button { padding: 0 1em 0 0; align-self: end; margin-bottom: 6px; } - #settings { display: flex; margin-left: 0.5em; } #settings>div.tab-content { margin-top: 1em; } #settings>div.tab-content>div>div { gap: 0; } @@ -172,10 +170,10 @@ div#extras_scale_to_tab div.form { flex-direction: row; } #settings .gradio-slider, #tab_settings .gradio-dropdown { width: 500px !important; max-width: 500px !important; } #settings .gradio-radio { padding: var(--block-padding) !important; } #settings textarea { width: 500px !important; max-width: 500px !important; } - .licenses { display: block !important; } /* live preview */ +#si-sparkline-memo, #si-sparkline-load { background-color: #111; } .progressDiv { position: relative; height: 20px; background: #b4c0cc; margin-bottom: -3px; } .dark .progressDiv { background: #424c5b; } .progressDiv .progress { width: 0%; height: 20px; background: #0060df; color: white; font-weight: bold; line-height: 20px; padding: 0 8px 0 0; text-align: right; overflow: visible; white-space: nowrap; padding: 0 0.5em; } From f965e7251570c7add96e6135102ba292a9894e07 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 09:40:02 -0400 Subject: [PATCH 093/122] fix nncf check Signed-off-by: Vladimir Mandic --- modules/sd_models.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modules/sd_models.py b/modules/sd_models.py index a7307561d..b500c51b7 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -575,7 +575,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No prompt_parser_diffusers.cache.clear() set_diffuser_options(sd_model, vae, op, offload=False) - if shared.opts.nncf_compress_weights and not ('Model' in shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): + if 'Model' in shared.opts.nncf_compress_weights and not ('Model' in shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): sd_model = model_quant.nncf_compress_weights(sd_model) # run this before move model so it can be compressed in CPU if shared.opts.optimum_quanto_weights: sd_model = model_quant.optimum_quanto_weights(sd_model) # run this before move model so it can be compressed in CPU From 15752f9299d4cf0e4526c3312a188fdabd52e5ae Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 09:49:44 -0400 Subject: [PATCH 094/122] prompt enhance loader exception logging Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index f45123060..ba1d41f15 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -69,7 +69,6 @@ class Script(scripts.Script): if self.model is not None and self.model == name: return - t0 = time.time() from modules import modelloader, model_quant, ggml modelloader.hf_login() model_repo = model_repo or self.options.models.get(name, {}).get('repo', None) or name @@ -93,6 +92,7 @@ class Script(scripts.Script): quant_args = model_quant.create_config(module='LLM') if not gguf_args else {} try: + t0 = time.time() self.model = None self.llm = None self.llm = transformers.AutoModelForCausalLM.from_pretrained( @@ -114,12 +114,12 @@ class Script(scripts.Script): for m in modules: shared.log.trace(f'Prompt enhance: {m}') self.model = name + t1 = time.time() + shared.log.info(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') except Exception as e: shared.log.error(f'Prompt enhance: load {e}') errors.display(e, 'Prompt enhance') devices.torch_gc() - t1 = time.time() - shared.log.info(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') self.busy = False def censored(self, response): From f46ee37f3ac2409a471ee339af2352f989ec5147 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Mon, 31 Mar 2025 22:49:38 +0900 Subject: [PATCH 095/122] zluda log & install improvements --- installer.py | 1 - modules/zluda_hijacks.py | 2 -- modules/zluda_installer.py | 13 +++++++++---- 3 files changed, 9 insertions(+), 7 deletions(-) diff --git a/installer.py b/installer.py index 4f384f2b8..289b2d5b3 100644 --- a/installer.py +++ b/installer.py @@ -652,7 +652,6 @@ def install_rocm_zluda(): zluda_installer.make_copy() zluda_installer.load() torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.6.0 torchvision --index-url https://download.pytorch.org/whl/cu118') - log.info(f'Using ZLUDA in {zluda_installer.path}') except Exception as e: error = e log.warning(f'Failed to load ZLUDA: {e}') diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index 4f1224923..ea60c846a 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -42,8 +42,6 @@ def torch__C__cuda_getCurrentRawStream(device): def do_hijack(): - torch.version.hip = rocm.version - if zluda.default_agent is not None: DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = zluda.default_agent.name torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 236ca85f5..54a0f9234 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -6,6 +6,7 @@ import shutil import zipfile import urllib.request from typing import Union +from installer import args, log from modules import rocm @@ -25,8 +26,6 @@ path = os.path.abspath(os.environ.get('ZLUDA', '.zluda')) default_agent: Union[rocm.Agent, None] = None hipBLASLt_enabled = False -nightly = os.environ.get("ZLUDA_NIGHTLY", "0") == "1" - class ZLUDAResult(ctypes.Structure): _fields_ = [ @@ -101,7 +100,10 @@ def install() -> None: platform = "windows" commit = os.environ.get("ZLUDA_HASH", "dba64c0966df2c71e82255e942c96e2e1cea3a2d") - if nightly: + if os.environ.get("ZLUDA_NIGHTLY", "0") == "1": + log.warning("Environment variable 'ZLUDA_NIGHTLY' will be removed. Please use command-line argument '--use-nightly' instead.") + args.use_nightly = True + if args.use_nightly: platform = "nightly-" + platform urllib.request.urlretrieve(f'https://github.com/lshqqytiger/ZLUDA/releases/download/rel.{commit}/ZLUDA-{platform}-rocm{rocm.version[0]}-amd64.zip', '_zluda') with zipfile.ZipFile('_zluda', 'r') as archive: @@ -150,6 +152,8 @@ def make_copy() -> None: def load() -> None: + log.info(f"ZLUDA load: path='{path}' nightly={bool(core.get_nightly_flag())}") + os.environ["ZLUDA_COMGR_LOG_LEVEL"] = "1" os.environ["ZLUDA_NVRTC_LIB"] = os.path.join([v for v in site.getsitepackages() if v.endswith("site-packages")][0], "torch", "lib", "nvrtc64_112_0.dll") @@ -171,7 +175,8 @@ def load() -> None: ctypes.windll.LoadLibrary(os.path.join(path, 'cudnn64_9.dll')) def conceal(): - import torch # pylint: disable=unused-import + import torch + torch.version.hip = rocm.version platform = sys.platform sys.platform = "" from torch.utils import cpp_extension From ff6d5e6e885c4d4ea49309c7fa8685f555b666c9 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 10:29:32 -0400 Subject: [PATCH 096/122] prompt enhance clean prompt improvements Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index ba1d41f15..ad63878ba 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -37,9 +37,10 @@ class Options: } default = list(models)[0] supported = list(transformers.integrations.ggml.GGUF_CONFIG_MAPPING) - system_prompt: str = 'You are a helpful assistant. You will be given a prompt used to create an image and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness based on original prompt. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing (if any), and any distinctive details (e.g., colors, textures, styles). NSFW and nudity is allowed, and if present it must be fully described. Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without prefix or suffix and as a simple text without formatting or numbering.' + system_prompt: str = 'You are a helpful assistant. You will be given a prompt used to create an image and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness based on original prompt. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing (if any), and any distinctive details (e.g., colors, textures, styles). NSFW and nudity is allowed, and if present it must be fully described. Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without explanation, prefix or suffix. Output as a simple text without formatting or numbering.' censored = ["i cannot", "i can't", "i am sorry", "against my programming", "i am not able", "i am unable", 'i am not allowed'] + max_delim_index: int = 60 max_tokens: int = 50 do_sample: bool = True temperature: float = 0.15 @@ -93,8 +94,10 @@ class Script(scripts.Script): try: t0 = time.time() + if self.llm is not None: + self.llm = None + shared.log.debug(f'Prompt enhance: name="{self.model}" unload') self.model = None - self.llm = None self.llm = transformers.AutoModelForCausalLM.from_pretrained( pretrained_model_name_or_path=model_repo if not gguf_args else model_gguf, trust_remote_code=True, @@ -136,14 +139,17 @@ class Script(scripts.Script): shared.log.debug('Prompt enhance: model unloaded') def clean(self, response): - response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n') + response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n').replace(' ', ' ') response = re.sub(r'<.*?>', '', response) + removed = '' if response.startswith('Prompt'): - response = response.split('Prompt', maxsplit=2)[1] - if ':' in response: - response = response.split(':', maxsplit=2)[1] - if '---' in response: - response = response.split('---', maxsplit=2)[0] + removed, response = response.split('Prompt', maxsplit=2) + if 0 <= response.find(':') < self.options.max_delim_index: + removed, response = response.split(':', maxsplit=2) + if 0 <= response.find('---') < self.options.max_delim_index: + response, removed = response.split('---', maxsplit=2) + if len(removed) > 0: + debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') response = response.strip() return response @@ -323,4 +329,5 @@ class Script(scripts.Script): temperature=temperature, penalty=repetition_penalty, ) + p.extra_generation_params['LLM'] = llm_model shared.state.end() From 5e79cff7a3d6c2f546143c4ccbd1832d2fb7f99f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 10:53:35 -0400 Subject: [PATCH 097/122] prompt enhance rewrite lists Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index ad63878ba..b1ce33636 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -139,8 +139,15 @@ class Script(scripts.Script): shared.log.debug('Prompt enhance: model unloaded') def clean(self, response): - response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n').replace(' ', ' ') + # remove special characters + response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n').replace(' ', ' ').replace('...', '.') + + # remove comments between brackets response = re.sub(r'<.*?>', '', response) + response = re.sub(r'\[.*?\]', '', response) + response = re.sub(r'\/.*?\/', '', response) + + # remove llm commentary removed = '' if response.startswith('Prompt'): removed, response = response.split('Prompt', maxsplit=2) @@ -150,6 +157,12 @@ class Script(scripts.Script): response, removed = response.split('---', maxsplit=2) if len(removed) > 0: debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') + + # remove bullets and lists + lines = response.splitlines() + filtered = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line) for line in lines] + response = '\n'.join(filtered) + response = response.strip() return response From b4c9eb193d2bdcdf24f6115620712f76caf081c0 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 11:00:21 -0400 Subject: [PATCH 098/122] cleanup Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index b1ce33636..abcfdbd76 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -140,7 +140,9 @@ class Script(scripts.Script): def clean(self, response): # remove special characters - response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '').replace('\n\n', '\n').replace(' ', ' ').replace('...', '.') + response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '') + # remove repeating characters + response = response.replace('\n\n', '\n').replace(' ', ' ').replace('...', '.') # remove comments between brackets response = re.sub(r'<.*?>', '', response) @@ -159,9 +161,8 @@ class Script(scripts.Script): debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') # remove bullets and lists - lines = response.splitlines() - filtered = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line) for line in lines] - response = '\n'.join(filtered) + lines = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line).strip() for line in response.splitlines()] + response = '\n'.join(lines) response = response.strip() return response From 85e37c0f03772c5449697a751fcf32f113eefbe3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 11:07:25 -0400 Subject: [PATCH 099/122] update changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 6 +++--- wiki | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 9b836a13e..905e1d6b7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2025-03-30 +## Update for 2025-03-31 -### Highlights for 2025-03-30 +### Highlights for 2025-03-31 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! @@ -15,7 +15,7 @@ Plus... - More quantization options and granular control - Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods -### Details for 2025-03-30 +### Details for 2025-03-31 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! diff --git a/wiki b/wiki index 90e18e0d1..bb4e16f8d 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 90e18e0d17ab43cddbe3d8ff9169707d1d289a41 +Subproject commit bb4e16f8d88d353f83f00f854852d1dd649fcd28 From e8ae79cc49e4ae28644c70f507fb81a2756463c5 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 13:59:46 -0400 Subject: [PATCH 100/122] hdr detect packed latents Signed-off-by: Vladimir Mandic --- modules/processing_correction.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/modules/processing_correction.py b/modules/processing_correction.py index 73d83d3cd..e06adae15 100644 --- a/modules/processing_correction.py +++ b/modules/processing_correction.py @@ -12,6 +12,14 @@ debug_enabled = os.environ.get('SD_HDR_DEBUG', None) is not None debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None debug('Trace: HDR') skip_correction = False +warned = False + + +def warn_once(message): + global warned # pylint: disable=global-statement + if not warned: + shared.log.warning(f'VAE: {message}') + warned = True def sharpen_tensor(tensor, ratio=0): @@ -121,6 +129,9 @@ def correction_callback(p, timestep, kwargs, initial: bool = False): return kwargs latents = kwargs["latents"] # debug(f'HDR correction: latents={latents.shape}') + if len(latents.shape) <= 3: # packed latent + warn_once(f'HDR correction: shape={latents.shape} packed latent') + return kwargs if len(latents.shape) == 4: # standard batched latent for i in range(latents.shape[0]): latents[i] = correction(p, timestep, latents[i]) @@ -135,6 +146,6 @@ def correction_callback(p, timestep, kwargs, initial: bool = False): latents[i] = correction(p, timestep, latents[i]) latents = latents.permute(1, 0, 2, 3).unsqueeze(0) else: - shared.log.debug(f'HDR correction: unknown latent shape {latents.shape}') + warn_once(f'HDR correction: shape={latents.shape} unknown latent') kwargs["latents"] = latents return kwargs From 0eaa2c037892bc9899a36426de0d411822fb24d4 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Tue, 1 Apr 2025 21:14:02 +0900 Subject: [PATCH 101/122] rocm wsl better arch detection --- modules/rocm.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/modules/rocm.py b/modules/rocm.py index b06e92ef6..a742556a2 100644 --- a/modules/rocm.py +++ b/modules/rocm.py @@ -172,12 +172,12 @@ else: return f'{arr[0]}.{arr[1]}' if len(arr) >= 2 else None def get_agents() -> List[Agent]: - if is_wsl: # WSL does not have 'rocm_agent_enumerator' - agents = spawn("rocminfo").split("\n") - agents = [x.strip().split(" ")[-1] for x in agents if x.startswith(' Name:') and "CPU" not in x] - else: + try: agents = spawn("rocm_agent_enumerator").split("\n") agents = [x for x in agents if x and x != 'gfx000'] + except Exception: # old version of ROCm WSL doesn't have rocm_agent_enumerator + agents = spawn("rocminfo").split("\n") + agents = [x.strip().split(" ")[-1] for x in agents if x.startswith(' Name:') and "CPU" not in x] return [Agent(x) for x in agents] def load_hsa_runtime() -> None: From 5906eb6792ff98df76f5bf0a15dae4be8e06c579 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 15:53:21 -0400 Subject: [PATCH 102/122] lora apply on gpu vs cpu settings option Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 4 +++- modules/call_queue.py | 8 ++++++-- modules/lora/networks.py | 24 +++++++++++++++--------- modules/modelloader.py | 3 ++- modules/shared.py | 1 + 5 files changed, 27 insertions(+), 13 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 905e1d6b7..e7230fbf7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,7 +33,9 @@ Plus... - **Default**: use vae from model - **Tiny VAE**: support for *Hunyuan, WAN, Mochi* - **Remote VAE**: support for *Hunyuan* - - **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog* + - **LoRA** + - support for *Hunyuan, LTX, WAN, Mochi, Cog* + - add option to apply LoRA directly on GPU or use CPU first in low-memory scenarios - additional key points: - all models are auto-downloaded upon first use uses *system paths -> huggingface* folder diff --git a/modules/call_queue.py b/modules/call_queue.py index 8368e5759..b33a9b6b2 100644 --- a/modules/call_queue.py +++ b/modules/call_queue.py @@ -81,8 +81,12 @@ def wrap_gradio_call(func, extra_outputs=None, add_stats=False, name=None): ooms = mem_mon_read.pop("oom") retries = mem_mon_read.pop("retries") vram = {k: v//1048576 for k, v in mem_mon_read.items()} - peak = max(vram['active_peak'], vram['reserved_peak'], vram['used']) - used = round(100.0 * peak / vram['total']) if vram['total'] > 0 else 0 + if 'active_peak' in vram: + peak = max(vram['active_peak'], vram['reserved_peak'], vram['used']) + used = round(100.0 * peak / vram['total']) if vram['total'] > 0 else 0 + else: + peak = 0 + used = 0 if peak > 0: gpu += f"| GPU {peak} MB" gpu += f" {used}%" if used > 0 else '' diff --git a/modules/lora/networks.py b/modules/lora/networks.py index b1af5d9b6..ecbaebce8 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -420,7 +420,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. return batch_updown, batch_ex_bias -def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False): +def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, device: torch.device = None): if lora_weights is None: return None if deactivate: @@ -429,16 +429,17 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G model_weights = self.weight # TODO lora: add other quantization types weight = None + device = device or devices.device if self.__class__.__name__ == 'Linear4bit' and bnb is not None: try: - dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(devices.device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) - new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) + dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) + new_weight = dequant_weight.to(device) + lora_weights.to(device) weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) except Exception as e: shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') else: try: - new_weight = model_weights.to(devices.device) + lora_weights.to(devices.device) + new_weight = model_weights.to(device) + lora_weights.to(device) except Exception: new_weight = model_weights + lora_weights # try without device cast weight = torch.nn.Parameter(new_weight, requires_grad=False) @@ -450,7 +451,7 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G return weight -def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False): +def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False, device: torch.device = None): weights_backup = getattr(self, "network_weights_backup", False) bias_backup = getattr(self, "network_bias_backup", False) if not isinstance(weights_backup, bool): # remove previous backup if we switched settings @@ -459,19 +460,20 @@ def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. bias_backup = True if not weights_backup and not bias_backup: return + device = device or devices.device t0 = time.time() if weights_backup: if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate) + weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate, device=device) if weight is not None: self.weight = weight if bias_backup: if ex_bias is not None: - bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate) + bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, device=device) if bias is not None: self.bias = bias @@ -593,6 +595,10 @@ def network_activate(include=[], exclude=[]): pbar = nullcontext() applied_weight = 0 applied_bias = 0 + if shared.opts.lora_apply_gpu: + device = devices.device + else: + device = devices.cpu with devices.inference_context(), pbar: wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) if len(loaded_networks) > 0 else () applied_layers.clear() @@ -609,7 +615,7 @@ def network_activate(include=[], exclude=[]): backup_size += network_backup_weights(module, network_layer_name, wanted_names) batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name) if shared.opts.lora_fuse_diffusers: - network_apply_direct(module, batch_updown, batch_ex_bias) + network_apply_direct(module, batch_updown, batch_ex_bias, device) else: network_apply_weights(module, batch_updown, batch_ex_bias, orig_device) if batch_updown is not None or batch_ex_bias is not None: @@ -627,7 +633,7 @@ def network_activate(include=[], exclude=[]): pbar.remove_task(task) # hide progress bar for no action timer.activate += time.time() - t0 if debug and len(loaded_networks) > 0: - shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} device={device} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() if len(loaded_networks) > 0 and (applied_weight > 0 or applied_bias > 0): if shared.opts.diffusers_offload_mode == "sequential": diff --git a/modules/modelloader.py b/modules/modelloader.py index 2a1d6745a..43325b7c3 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -64,7 +64,8 @@ def download_civit_meta(model_path: str, model_id): def download_civit_preview(model_path: str, preview_url: str): ext = os.path.splitext(preview_url)[1] preview_file = os.path.splitext(model_path)[0] + ext - if preview_file.endswith('.mp4'): + is_video = preview_file.lower().endswith('.mp4') + if is_video: shared.log.warning(f'CivitAI download: url="{preview_url}" skip video') return '' if os.path.exists(preview_file): diff --git a/modules/shared.py b/modules/shared.py index 3829408b9..06b72790b 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -926,6 +926,7 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "lora_preferred_name": OptionInfo("filename", "LoRA preferred name", gr.Radio, {"choices": ["filename", "alias"], "visible": False}), "lora_add_hashes_to_infotext": OptionInfo(False, "LoRA add hash info to metadata"), "lora_fuse_diffusers": OptionInfo(True, "LoRA fuse directly to model"), + "lora_apply_gpu": OptionInfo(True, "LoRA load directly on GPU"), "lora_legacy": OptionInfo(not native, "LoRA load using legacy method"), "lora_force_diffusers": OptionInfo(False if not cmd_opts.use_openvino else True, "LoRA load using Diffusers method"), "lora_maybe_diffusers": OptionInfo(False, "LoRA load using Diffusers method for selected models"), From 6b8299eac86fce357d7ea8cef2614461d47e7ece Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 16:21:20 -0400 Subject: [PATCH 103/122] lora mpt4 preview support Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 2 ++ modules/modelloader.py | 42 +++++++++++++++++++++++++++++------ modules/ui_control_helpers.py | 26 ++-------------------- modules/ui_models.py | 3 +++ modules/video.py | 23 +++++++++++++++++++ 5 files changed, 65 insertions(+), 31 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e7230fbf7..9af36d7bd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -36,6 +36,8 @@ Plus... - **LoRA** - support for *Hunyuan, LTX, WAN, Mochi, Cog* - add option to apply LoRA directly on GPU or use CPU first in low-memory scenarios + - improve metadata and preview parallel fetch + - support for mp4 so first frame is extracted as used as lora preview - additional key points: - all models are auto-downloaded upon first use uses *system paths -> huggingface* folder diff --git a/modules/modelloader.py b/modules/modelloader.py index 43325b7c3..412ef0aa1 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -18,6 +18,7 @@ from modules.paths import script_path, models_path loggedin = None diffuser_repos = [] debug = shared.log.trace if os.environ.get('SD_DOWNLOAD_DEBUG', None) is not None else lambda *args, **kwargs: None +pbar = None def hf_login(token=None): @@ -61,12 +62,34 @@ def download_civit_meta(model_path: str, model_id): return f'CivitAI download error: id={model_id} url={url} code={r.status_code}' +def save_video_frame(filepath: str): + from modules import video + try: + frames, fps, duration, w, h, codec, frame = video.get_video_params(filepath, capture=True) + except Exception as e: + shared.log.error(f'Video: file={filepath} {e}') + return None + if frame is not None: + basename = os.path.splitext(filepath) + thumb = f'{basename[0]}.thumb.jpg' + shared.log.debug(f'Video: file={filepath} frames={frames} fps={fps} size={w}x{h} codec={codec} duration={duration} thumb={thumb}') + frame.save(thumb) + else: + shared.log.error(f'Video: file={filepath} no frames found') + return frame + + def download_civit_preview(model_path: str, preview_url: str): + global pbar # pylint: disable=global-statement + if model_path is None: + pbar = None + return '' ext = os.path.splitext(preview_url)[1] preview_file = os.path.splitext(model_path)[0] + ext is_video = preview_file.lower().endswith('.mp4') - if is_video: - shared.log.warning(f'CivitAI download: url="{preview_url}" skip video') + is_json = preview_file.lower().endswith('.json') + if is_json: + shared.log.warning(f'CivitAI download: url="{preview_url}" skip json') return '' if os.path.exists(preview_file): return '' @@ -77,20 +100,25 @@ def download_civit_preview(model_path: str, preview_url: str): written = 0 img = None shared.state.begin('CivitAI') + if pbar is None: + pbar = p.Progress(p.TextColumn('[cyan]{task.description}'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), console=shared.console) try: with open(preview_file, 'wb') as f: - with p.Progress(p.TextColumn('[cyan]{task.description}'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), console=shared.console) as progress: - task = progress.add_task(description="Download starting", total=total_size) + with pbar: + task = pbar.add_task(description="Download starting", total=total_size) for data in r.iter_content(block_size): written = written + len(data) f.write(data) - progress.update(task, advance=block_size, description="Downloading") + pbar.update(task, advance=block_size, description="Downloading") if written < 1024: # min threshold os.remove(preview_file) raise ValueError(f'removed invalid download: bytes={written}') - img = Image.open(preview_file) + if is_video: + img = save_video_frame(preview_file) + else: + img = Image.open(preview_file) except Exception as e: - os.remove(preview_file) + # os.remove(preview_file) res += f' error={e}' shared.log.error(f'CivitAI download error: url={preview_url} file="{preview_file}" written={written} {e}') shared.state.end() diff --git a/modules/ui_control_helpers.py b/modules/ui_control_helpers.py index 2b10e21d3..de8983be2 100644 --- a/modules/ui_control_helpers.py +++ b/modules/ui_control_helpers.py @@ -1,7 +1,7 @@ import os import gradio as gr from PIL import Image -from modules import shared, scripts, masking # pylint: disable=ungrouped-imports +from modules import shared, scripts, masking, video # pylint: disable=ungrouped-imports gr_height = None @@ -82,31 +82,9 @@ def display_units(num_units): return (num_units * [gr.update(visible=True)]) + ((max_units - num_units) * [gr.update(visible=False)]) -def get_video_params(filepath: str, capture: bool = False): - import cv2 - from modules.control.util import decode_fourcc - video = cv2.VideoCapture(filepath) - if not video.isOpened(): - msg = f'Control: video open failed: path="{filepath}"' - shared.log.error(msg) - raise RuntimeError(msg) - frames = int(video.get(cv2.CAP_PROP_FRAME_COUNT)) - fps = video.get(cv2.CAP_PROP_FPS) - duration = float(frames) / fps - w, h = int(video.get(cv2.CAP_PROP_FRAME_WIDTH)), int(video.get(cv2.CAP_PROP_FRAME_HEIGHT)) - codec = decode_fourcc(video.get(cv2.CAP_PROP_FOURCC)) - frame = None - if capture: - _status, frame = video.read() - frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) - frame = Image.fromarray(frame) - video.release() - return frames, fps, duration, w, h, codec, frame - - def get_video(filepath: str): try: - frames, fps, duration, w, h, codec, _cap = get_video_params(filepath) + frames, fps, duration, w, h, codec, _cap = video.get_video_params(filepath) shared.log.debug(f'Control: input video: path={filepath} frames={frames} fps={fps} size={w}x{h} codec={codec}') msg = f'Control input | Video | Size {w}x{h} | Frames {frames} | FPS {fps:.2f} | Duration {duration:.2f} | Codec {codec}' return msg diff --git a/modules/ui_models.py b/modules/ui_models.py index 4f6355a9a..59d8a9196 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -573,6 +573,8 @@ def create_ui(): def atomic_civit_search_metadata(item, res, rehash): from modules.modelloader import download_civit_preview, download_civit_meta + if item is None: + return meta = os.path.splitext(item['filename'])[0] + '.json' has_meta = os.path.isfile(meta) and os.stat(meta).st_size > 0 if ('card-no-preview.png' in item['preview'] or not has_meta) and os.path.isfile(item['filename']): @@ -629,6 +631,7 @@ def create_ui(): with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor: for fn in candidates: executor.submit(atomic_civit_search_metadata, fn, res, rehash) + atomic_civit_search_metadata(None, res, rehash) t1 = time.time() log.debug(f'CivitAI search metadata: items={i} time={t1-t0:.2f}') txt = '
'.join([r for r in res if len(r) > 0]) diff --git a/modules/video.py b/modules/video.py index d9e40a27f..bd26aac04 100644 --- a/modules/video.py +++ b/modules/video.py @@ -1,6 +1,7 @@ import os import threading import numpy as np +from PIL import Image from modules import shared, errors from modules.images_namegen import FilenameGenerator # pylint: disable=unused-import @@ -83,3 +84,25 @@ def save_video(p, images, filename = None, video_type: str = 'none', duration: f else: save_video_atomic(images, filename, video_type, duration, loop, interpolate, scale, pad, change) return filename + + +def get_video_params(filepath: str, capture: bool = False): + import cv2 + from modules.control.util import decode_fourcc + video = cv2.VideoCapture(filepath) + if not video.isOpened(): + msg = f'Video open failed: path="{filepath}"' + shared.log.error(msg) + raise RuntimeError(msg) + frames = int(video.get(cv2.CAP_PROP_FRAME_COUNT)) + fps = round(video.get(cv2.CAP_PROP_FPS), 2) + duration = round(float(frames) / fps, 2) + w, h = int(video.get(cv2.CAP_PROP_FRAME_WIDTH)), int(video.get(cv2.CAP_PROP_FRAME_HEIGHT)) + codec = decode_fourcc(video.get(cv2.CAP_PROP_FOURCC)) + frame = None + if capture: + _status, frame = video.read() + frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + frame = Image.fromarray(frame) + video.release() + return frames, fps, duration, w, h, codec, frame From 032bd46de2fd3a9c4e4d4410799273bfc2297cfd Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 20:53:46 -0400 Subject: [PATCH 104/122] improve mp4 download Signed-off-by: Vladimir Mandic --- modules/lora/networks.py | 15 ++++++++++++++- modules/modelloader.py | 8 ++++---- modules/shared.py | 2 +- 3 files changed, 19 insertions(+), 6 deletions(-) diff --git a/modules/lora/networks.py b/modules/lora/networks.py index ecbaebce8..4903e5a6f 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -276,7 +276,8 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non while len(lora_cache) > shared.opts.lora_in_memory_limit: name = next(iter(lora_cache)) - lora_cache.pop(name, None) + lora = lora_cache.pop(name, None) + del lora if not skip_lora_load and len(diffuser_loaded) > 0: shared.log.debug(f'Network load: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') @@ -535,6 +536,10 @@ def network_deactivate(include=[], exclude=[]): modules[name] = list(component.named_modules()) active_components.append(name) total = sum(len(x) for x in modules.values()) + if shared.opts.lora_apply_gpu: + device = devices.device + else: + device = devices.cpu if len(previously_loaded_networks) > 0 and debug: pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=deactivate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) task = pbar.add_task(description='', total=total) @@ -560,6 +565,10 @@ def network_deactivate(include=[], exclude=[]): applied_layers.append(network_layer_name) del batch_updown, batch_ex_bias module.network_current_names = () + try: + module.to(device) + except Exception: + pass if task is not None: pbar.update(task, advance=1, description=f'networks={len(previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') @@ -626,6 +635,10 @@ def network_activate(include=[], exclude=[]): applied_bias += 1 del batch_updown, batch_ex_bias module.network_current_names = wanted_names + try: + module.to(device) + except Exception: + pass if task is not None: pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size}') diff --git a/modules/modelloader.py b/modules/modelloader.py index 412ef0aa1..a132b765a 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -90,7 +90,7 @@ def download_civit_preview(model_path: str, preview_url: str): is_json = preview_file.lower().endswith('.json') if is_json: shared.log.warning(f'CivitAI download: url="{preview_url}" skip json') - return '' + return 'CivitAI download error: JSON file' if os.path.exists(preview_file): return '' res = f'CivitAI download: url={preview_url} file="{preview_file}"' @@ -101,15 +101,15 @@ def download_civit_preview(model_path: str, preview_url: str): img = None shared.state.begin('CivitAI') if pbar is None: - pbar = p.Progress(p.TextColumn('[cyan]{task.description}'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), console=shared.console) + pbar = p.Progress(p.TextColumn('[cyan]Download'), p.DownloadColumn(), p.BarColumn(), p.TaskProgressColumn(), p.TimeRemainingColumn(), p.TimeElapsedColumn(), p.TransferSpeedColumn(), p.TextColumn('[yellow]{task.description}'), console=shared.console) try: with open(preview_file, 'wb') as f: with pbar: - task = pbar.add_task(description="Download starting", total=total_size) + task = pbar.add_task(description=preview_file, total=total_size) for data in r.iter_content(block_size): written = written + len(data) f.write(data) - pbar.update(task, advance=block_size, description="Downloading") + pbar.update(task, advance=block_size) if written < 1024: # min threshold os.remove(preview_file) raise ValueError(f'removed invalid download: bytes={written}') diff --git a/modules/shared.py b/modules/shared.py index 06b72790b..0235e2e31 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -926,7 +926,7 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "lora_preferred_name": OptionInfo("filename", "LoRA preferred name", gr.Radio, {"choices": ["filename", "alias"], "visible": False}), "lora_add_hashes_to_infotext": OptionInfo(False, "LoRA add hash info to metadata"), "lora_fuse_diffusers": OptionInfo(True, "LoRA fuse directly to model"), - "lora_apply_gpu": OptionInfo(True, "LoRA load directly on GPU"), + "lora_apply_gpu": OptionInfo(False, "LoRA load directly on GPU"), "lora_legacy": OptionInfo(not native, "LoRA load using legacy method"), "lora_force_diffusers": OptionInfo(False if not cmd_opts.use_openvino else True, "LoRA load using Diffusers method"), "lora_maybe_diffusers": OptionInfo(False, "LoRA load using Diffusers method for selected models"), From b5031a5ebae1cea27a29f03137d2fab9e2cebf14 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 31 Mar 2025 23:30:15 -0400 Subject: [PATCH 105/122] lora modularize code Signed-off-by: Vladimir Mandic --- modules/api/endpoints.py | 8 +- modules/api/gallery.py | 2 +- modules/extra_networks.py | 5 +- modules/infotext.py | 8 +- modules/lora/extra_networks_lora.py | 10 +- modules/lora/lora_apply.py | 217 ++++++ modules/lora/lora_common.py | 21 + modules/lora/lora_load.py | 285 ++++++++ ...network_overrides.py => lora_overrides.py} | 0 modules/lora/network.py | 6 +- modules/lora/networks.py | 650 ++---------------- modules/model_auraflow.py | 2 +- modules/model_flux.py | 3 +- modules/model_kolors.py | 2 +- modules/model_lumina.py | 3 +- modules/model_meissonic.py | 2 +- modules/model_omnigen.py | 9 +- modules/model_pixart.py | 2 +- modules/model_sana.py | 3 +- modules/model_sd3.py | 2 +- modules/model_stablecascade.py | 1 + modules/sd_models.py | 8 +- modules/ui_extensions.py | 1 - modules/ui_extra_networks_lora.py | 10 +- modules/ui_gallery.py | 4 +- modules/zluda_hijacks.py | 2 +- webui.py | 4 +- 27 files changed, 634 insertions(+), 636 deletions(-) create mode 100644 modules/lora/lora_apply.py create mode 100644 modules/lora/lora_common.py create mode 100644 modules/lora/lora_load.py rename modules/lora/{network_overrides.py => lora_overrides.py} (100%) diff --git a/modules/api/endpoints.py b/modules/api/endpoints.py index ee30fac44..87362b4c9 100644 --- a/modules/api/endpoints.py +++ b/modules/api/endpoints.py @@ -41,10 +41,10 @@ def get_embeddings(): return {"loaded": convert_embeddings(db.word_embeddings), "skipped": convert_embeddings(db.skipped_embeddings)} def get_loras(): - from modules.lora import network, networks + from modules.lora import network, lora_load def create_lora_json(obj: network.NetworkOnDisk): return { "name": obj.name, "alias": obj.alias, "path": obj.filename, "metadata": obj.metadata } - return [create_lora_json(obj) for obj in networks.available_networks.values()] + return [create_lora_json(obj) for obj in lora_load.available_networks.values()] def get_extra_networks(page: Optional[str] = None, name: Optional[str] = None, filename: Optional[str] = None, title: Optional[str] = None, fullname: Optional[str] = None, hash: Optional[str] = None): # pylint: disable=redefined-builtin res = [] @@ -134,8 +134,8 @@ def post_refresh_vae(): return shared.refresh_vaes() def post_refresh_loras(): - from modules.lora import networks - return networks.list_available_networks() + from modules.lora import lora_load + return lora_load.list_available_networks() def get_extensions_list(): from modules import extensions diff --git a/modules/api/gallery.py b/modules/api/gallery.py index e56bb00ac..e1add81af 100644 --- a/modules/api/gallery.py +++ b/modules/api/gallery.py @@ -74,7 +74,7 @@ def register_api(app: FastAPI): # register api manager = ConnectionManager() def get_video_thumbnail(filepath): - from modules.ui_control_helpers import get_video_params + from modules.video import get_video_params try: stat = os.stat(filepath) frames, fps, duration, width, height, codec, frame = get_video_params(filepath, capture=True) diff --git a/modules/extra_networks.py b/modules/extra_networks.py index e882b113c..d8b638b85 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -19,8 +19,9 @@ def register_default_extra_networks(): from modules.ui_extra_networks_styles import ExtraNetworkStyles register_extra_network(ExtraNetworkStyles()) if not shared.opts.lora_legacy: - from modules.lora.networks import extra_network_lora - register_extra_network(extra_network_lora) + from modules.lora import lora_common, extra_networks_lora + lora_common.extra_network_lora = extra_networks_lora.ExtraNetworkLora() + register_extra_network(lora_common.extra_network_lora) if shared.opts.hypernetwork_enabled: from modules.ui_extra_networks_hypernet import ExtraNetworkHypernet register_extra_network(ExtraNetworkHypernet()) diff --git a/modules/infotext.py b/modules/infotext.py index 78c1fd92e..497879d31 100644 --- a/modules/infotext.py +++ b/modules/infotext.py @@ -31,7 +31,7 @@ def unquote(text): # disabled by default can be enabled if needed def check_lora(params): try: - import modules.lora.networks as networks + from modules.lora import lora_load from modules.errors import log # pylint: disable=redefined-outer-name except Exception: return @@ -39,14 +39,14 @@ def check_lora(params): found = [] missing = [] for l in loras: - lora = networks.available_network_hash_lookup.get(l, None) + lora = lora_load.available_network_hash_lookup.get(l, None) if lora is not None: found.append(lora.name) else: missing.append(l) loras = [s.strip() for s in params.get('LoRA networks', '').split(',')] for l in loras: - lora = networks.available_network_aliases.get(l, None) + lora = lora_load.available_network_aliases.get(l, None) if lora is not None: found.append(lora.name) else: @@ -54,7 +54,7 @@ def check_lora(params): # networks.available_network_aliases.get(name, None) loras = re_lora.findall(params.get('Prompt', '')) for l in loras: - lora = networks.available_network_aliases.get(l, None) + lora = lora_load.available_network_aliases.get(l, None) if lora is not None: found.append(lora.name) else: diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index a775ac91f..7227680ce 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -2,7 +2,7 @@ from typing import List import os import re import numpy as np -from modules.lora import networks, network_overrides +from modules.lora import networks, lora_overrides, lora_load from modules import extra_networks, shared, sd_models @@ -156,13 +156,13 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access debug_log(f'Network load: type=LoRA include={include} exclude={exclude} requested={requested} fn={fn}') - force_diffusers = network_overrides.check_override() + force_diffusers = lora_overrides.check_override() if force_diffusers: has_changed = False # diffusers handle their own loading if len(exclude) == 0: - networks.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load only on first call + lora_load.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load only on first call else: - networks.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load + lora_load.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load has_changed = self.changed(requested, include, exclude) if has_changed: networks.network_deactivate(include, exclude) @@ -180,7 +180,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if shared.native: networks.previously_loaded_networks = networks.loaded_networks.copy() debug_log(f'Network load: type=LoRA active={[n.name for n in networks.previously_loaded_networks]} deactivate') - if shared.native and len(networks.diffuser_loaded) > 0: + if shared.native and len(lora_load.diffuser_loaded) > 0: if not (shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled is True): if hasattr(shared.sd_model, "unfuse_lora"): try: diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py new file mode 100644 index 000000000..7e865a167 --- /dev/null +++ b/modules/lora/lora_apply.py @@ -0,0 +1,217 @@ +from typing import Union +import re +import time +import torch +import diffusers.models.lora +from modules.lora.lora_common import timer, debug, loaded_networks, previously_loaded_networks, extra_network_lora +from modules import shared, devices, errors, model_quant + + +bnb = None +re_network_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") + + +def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], network_layer_name: str, wanted_names: tuple): + global bnb # pylint: disable=W0603 + backup_size = 0 + if len(loaded_networks) > 0 and network_layer_name is not None and any([net.modules.get(network_layer_name, None) for net in loaded_networks]): # noqa: C419 # pylint: disable=R1729 + t0 = time.time() + + weights_backup = getattr(self, "network_weights_backup", None) + bias_backup = getattr(self, "network_bias_backup", None) + if weights_backup is not None or bias_backup is not None: + if (shared.opts.lora_fuse_diffusers and not isinstance(weights_backup, bool)) or (not shared.opts.lora_fuse_diffusers and isinstance(weights_backup, bool)): # invalidate so we can change direct/backup on-the-fly + weights_backup = None + bias_backup = None + self.network_weights_backup = weights_backup + self.network_bias_backup = bias_backup + + if weights_backup is None and wanted_names != (): # pylint: disable=C1803 + weight = getattr(self, 'weight', None) + self.network_weights_backup = None + if getattr(weight, "quant_type", None) in ['nf4', 'fp4']: + if bnb is None: + bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) + if bnb is not None: + with devices.inference_context(): + if shared.opts.lora_fuse_diffusers: + self.network_weights_backup = True + else: + self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,) + self.quant_state = weight.quant_state + self.quant_type = weight.quant_type + self.blocksize = weight.blocksize + else: + if shared.opts.lora_fuse_diffusers: + self.network_weights_backup = True + else: + weights_backup = weight.clone() + self.network_weights_backup = weights_backup.to(devices.cpu) + else: + if shared.opts.lora_fuse_diffusers: + self.network_weights_backup = True + else: + self.network_weights_backup = weight.clone().to(devices.cpu) + + if bias_backup is None: + if getattr(self, 'bias', None) is not None: + if shared.opts.lora_fuse_diffusers: + self.network_bias_backup = True + else: + bias_backup = self.bias.clone() + bias_backup = bias_backup.to(devices.cpu) + + if getattr(self, 'network_weights_backup', None) is not None: + backup_size += self.network_weights_backup.numel() * self.network_weights_backup.element_size() if isinstance(self.network_weights_backup, torch.Tensor) else 0 + if getattr(self, 'network_bias_backup', None) is not None: + backup_size += self.network_bias_backup.numel() * self.network_bias_backup.element_size() if isinstance(self.network_bias_backup, torch.Tensor) else 0 + timer.backup += time.time() - t0 + return backup_size + + +def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], network_layer_name: str, use_previous: bool = False): + if shared.opts.diffusers_offload_mode == "none": + try: + self.to(devices.device) + except Exception: + pass + batch_updown = None + batch_ex_bias = None + loaded = loaded_networks if not use_previous else previously_loaded_networks + for net in loaded: + module = net.modules.get(network_layer_name, None) + if module is None: + continue + try: + t0 = time.time() + try: + weight = self.weight.to(devices.device) + except Exception: + weight = self.weight + + updown, ex_bias = module.calc_updown(weight) + if updown is not None: + if batch_updown is not None: + batch_updown += updown.to(batch_updown.device) + else: + batch_updown = updown.to(devices.device) + if ex_bias is not None: + if batch_ex_bias: + batch_ex_bias += ex_bias.to(batch_ex_bias.device) + else: + batch_ex_bias = ex_bias.to(devices.device) + timer.calc += time.time() - t0 + + if shared.opts.diffusers_offload_mode == "sequential": + t0 = time.time() + if batch_updown is not None: + batch_updown = batch_updown.to(devices.cpu) + if batch_ex_bias is not None: + batch_ex_bias = batch_ex_bias.to(devices.cpu) + t1 = time.time() + timer.move += t1 - t0 + except RuntimeError as e: + extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 + if debug: + module_name = net.modules.get(network_layer_name, None) + shared.log.error(f'LoRA apply weight name="{net.name}" module="{module_name}" layer="{network_layer_name}" {e}') + errors.display(e, 'LoRA') + raise RuntimeError('LoRA apply weight') from e + continue + return batch_updown, batch_ex_bias + + +def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, device: torch.device = devices.device): + if lora_weights is None: + return None + if deactivate: + lora_weights *= -1 + if model_weights is None: # weights are used if provided-from-backup else use self.weight + model_weights = self.weight + # TODO lora: add other quantization types + weight = None + if self.__class__.__name__ == 'Linear4bit' and bnb is not None: + try: + dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(devices.device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) + new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) + weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize, requires_grad=False) + # weight._quantize(devices.device) # TODO force imediate quantization + except Exception as e: + shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') + else: + try: + new_weight = model_weights.to(devices.device) + lora_weights.to(devices.device) + except Exception: + new_weight = model_weights + lora_weights # try without device cast + weight = torch.nn.Parameter(new_weight, requires_grad=False) + try: + # weight.to(device=device) # TODO required since quantization happens only during .to call, not during params creation + pass + except Exception: + pass # may fail if weights is meta tensor + return weight + + +def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False, device: torch.device = devices.device): + weights_backup = getattr(self, "network_weights_backup", False) + bias_backup = getattr(self, "network_bias_backup", False) + device = device or devices.device + if not isinstance(weights_backup, bool): # remove previous backup if we switched settings + weights_backup = True + if not isinstance(bias_backup, bool): + bias_backup = True + if not weights_backup and not bias_backup: + return + t0 = time.time() + + if weights_backup: + if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 + updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable + if updown is not None: + weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate, device=device) + if weight is not None: + self.weight = weight + + if bias_backup: + if ex_bias is not None: + bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, device=device) + if bias is not None: + self.bias = bias + + if hasattr(self, "qweight") and hasattr(self, "freeze"): + self.freeze() + + timer.apply += time.time() - t0 + + +def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, device: torch.device, deactivate: bool = False): + weights_backup = getattr(self, "network_weights_backup", None) + bias_backup = getattr(self, "network_bias_backup", None) + if weights_backup is None and bias_backup is None: + return + t0 = time.time() + + if weights_backup is not None: + self.weight = None + if updown is not None and len(weights_backup.shape) == 4 and weights_backup.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 + updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable + if updown is not None: + weight = network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate, device=device) + if weight is not None: + self.weight = weight + else: + self.weight = torch.nn.Parameter(weights_backup.to(device), requires_grad=False) + + if bias_backup is not None: + self.bias = None + if ex_bias is not None: + bias = network_add_weights(self, model_weights=weights_backup, lora_weights=ex_bias, deactivate=deactivate, device=device) + if bias: + self.weight = bias + else: + self.bias = torch.nn.Parameter(bias_backup.to(device), requires_grad=False) + + if hasattr(self, "qweight") and hasattr(self, "freeze"): + self.freeze() + + timer.apply += time.time() - t0 diff --git a/modules/lora/lora_common.py b/modules/lora/lora_common.py new file mode 100644 index 000000000..a6b15ae13 --- /dev/null +++ b/modules/lora/lora_common.py @@ -0,0 +1,21 @@ +from typing import List +import os +from modules.lora import lora_timers +from modules.lora import network_lora, network_hada, network_ia3, network_oft, network_lokr, network_full, network_norm, network_glora + + +timer = lora_timers.Timer() +debug = os.environ.get('SD_LORA_DEBUG', None) is not None +module_types = [ + network_lora.ModuleTypeLora(), + network_hada.ModuleTypeHada(), + network_ia3.ModuleTypeIa3(), + network_oft.ModuleTypeOFT(), + network_lokr.ModuleTypeLokr(), + network_full.ModuleTypeFull(), + network_norm.ModuleTypeNorm(), + network_glora.ModuleTypeGLora(), +] +loaded_networks: List = [] # no type due to circular import +previously_loaded_networks: List = [] # no type due to circular import +extra_network_lora = None # initialized in extra_networks.py diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py new file mode 100644 index 000000000..43efc6f04 --- /dev/null +++ b/modules/lora/lora_load.py @@ -0,0 +1,285 @@ +from typing import Union +import os +import time +import concurrent +from modules import shared, errors, devices, sd_models, sd_models_compile, files_cache +from modules.lora import network, lora_overrides, lora_convert +from modules.lora.lora_common import timer, debug, module_types, loaded_networks + + +diffuser_loaded = [] +diffuser_scales = [] +lora_cache = {} +available_networks = {} +available_network_aliases = {} +forbidden_network_aliases = {} +available_network_hash_lookup = {} + + +def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_default_multiplier) -> Union[network.Network, None]: + t0 = time.time() + name = name.replace(".", "_") + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') + if not shared.native: + return None + if not hasattr(shared.sd_model, 'load_lora_weights'): + shared.log.error(f'Network load: type=LoRA class={shared.sd_model.__class__} does not implement load lora') + return None + try: + shared.sd_model.load_lora_weights(network_on_disk.filename, adapter_name=name) + except Exception as e: + if 'already in use' in str(e): + pass + else: + if 'The following keys have not been correctly renamed' in str(e): + shared.log.error(f'Network load: type=LoRA name="{name}" diffusers unsupported format') + else: + shared.log.error(f'Network load: type=LoRA name="{name}" {e}') + if debug: + errors.display(e, "LoRA") + return None + if name not in diffuser_loaded: + diffuser_loaded.append(name) + diffuser_scales.append(lora_scale) + net = network.Network(name, network_on_disk) + net.mtime = os.path.getmtime(network_on_disk.filename) + timer.activate += time.time() - t0 + return net + + +def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: + if not shared.sd_loaded: + return None + + cached = lora_cache.get(name, None) + if debug: + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') + if cached is not None: + return cached + net = network.Network(name, network_on_disk) + net.mtime = os.path.getmtime(network_on_disk.filename) + sd = sd_models.read_state_dict(network_on_disk.filename, what='network') + if shared.sd_model_type == 'f1': # if kohya flux lora, convert state_dict + sd = lora_convert._convert_kohya_flux_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + if shared.sd_model_type == 'sd3': # if kohya flux lora, convert state_dict + try: + sd = lora_convert._convert_kohya_sd3_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + except ValueError: # EAFP for diffusers PEFT keys + pass + lora_convert.assign_network_names_to_compvis_modules(shared.sd_model) + keys_failed_to_match = {} + matched_networks = {} + bundle_embeddings = {} + dtypes = [] + convert = lora_convert.KeyConvert() + for key_network, weight in sd.items(): + parts = key_network.split('.') + if parts[0] == "bundle_emb": + emb_name, vec_name = parts[1], key_network.split(".", 2)[-1] + emb_dict = bundle_embeddings.get(emb_name, {}) + emb_dict[vec_name] = weight + bundle_embeddings[emb_name] = emb_dict + continue + if len(parts) > 5: # messy handler for diffusers peft lora + key_network_without_network_parts = '_'.join(parts[:-2]) + if not key_network_without_network_parts.startswith('lora_'): + key_network_without_network_parts = 'lora_' + key_network_without_network_parts + network_part = '.'.join(parts[-2:]).replace('lora_A', 'lora_down').replace('lora_B', 'lora_up') + else: + key_network_without_network_parts, network_part = key_network.split(".", 1) + key, sd_module = convert(key_network_without_network_parts) + if sd_module is None: + keys_failed_to_match[key_network] = key + continue + if key not in matched_networks: + matched_networks[key] = network.NetworkWeights(network_key=key_network, sd_key=key, w={}, sd_module=sd_module) + matched_networks[key].w[network_part] = weight + if weight.dtype not in dtypes: + dtypes.append(weight.dtype) + network_types = [] + for key, weights in matched_networks.items(): + net_module = None + for nettype in module_types: + net_module = nettype.create_module(net, weights) + if net_module is not None: + network_types.append(nettype.__class__.__name__) + break + if net_module is None: + shared.log.error(f'LoRA unhandled: name={name} key={key} weights={weights.w.keys()}') + else: + net.modules[key] = net_module + if len(keys_failed_to_match) > 0: + shared.log.warning(f'Network load: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') + if debug: + shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') + else: + shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} dtypes={dtypes} direct={shared.opts.lora_fuse_diffusers}') + if len(matched_networks) == 0: + return None + lora_cache[name] = net + net.bundle_embeddings = bundle_embeddings + return net + + +def maybe_recompile_model(names, te_multipliers): + recompile_model = False + skip_lora_load = False + if shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled: + if len(names) == len(shared.compiled_model_state.lora_model): + for i, name in enumerate(names): + if shared.compiled_model_state.lora_model[ + i] != f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}": + recompile_model = True + shared.compiled_model_state.lora_model = [] + break + if not recompile_model: + skip_lora_load = True + if len(loaded_networks) > 0 and debug: + shared.log.debug('Model Compile: Skipping LoRa loading') + return recompile_model, skip_lora_load + else: + recompile_model = True + shared.compiled_model_state.lora_model = [] + if recompile_model: + backup_cuda_compile = shared.opts.cuda_compile + backup_scheduler = getattr(shared.sd_model, "scheduler", None) + sd_models.unload_model_weights(op='model') + shared.opts.cuda_compile = [] + sd_models.reload_model_weights(op='model') + shared.opts.cuda_compile = backup_cuda_compile + if backup_scheduler is not None: + shared.sd_model.scheduler = backup_scheduler + return recompile_model, skip_lora_load + + +def list_available_networks(): + t0 = time.time() + available_networks.clear() + available_network_aliases.clear() + forbidden_network_aliases.clear() + available_network_hash_lookup.clear() + forbidden_network_aliases.update({"none": 1, "Addams": 1}) + if not os.path.exists(shared.cmd_opts.lora_dir): + shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') + + def add_network(filename): + if not os.path.isfile(filename): + return + name = os.path.splitext(os.path.basename(filename))[0] + name = name.replace('.', '_') + try: + entry = network.NetworkOnDisk(name, filename) + available_networks[entry.name] = entry + if entry.alias in available_network_aliases: + forbidden_network_aliases[entry.alias.lower()] = 1 + if shared.opts.lora_preferred_name == 'filename': + available_network_aliases[entry.name] = entry + else: + available_network_aliases[entry.alias] = entry + if entry.shorthash: + available_network_hash_lookup[entry.shorthash] = entry + except OSError as e: # should catch FileNotFoundError and PermissionError etc. + shared.log.error(f'LoRA: filename="{filename}" {e}') + + candidates = sorted(files_cache.list_files(shared.cmd_opts.lora_dir, ext_filter=[".pt", ".ckpt", ".safetensors"])) + with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: + for fn in candidates: + executor.submit(add_network, fn) + t1 = time.time() + timer.list = t1 - t0 + shared.log.info(f'Available LoRAs: path="{shared.cmd_opts.lora_dir}" items={len(available_networks)} folders={len(forbidden_network_aliases)} time={t1 - t0:.2f}') + + +def network_download(name): + from huggingface_hub import hf_hub_download + if os.path.exists(name): + return network.NetworkOnDisk(name, name) + parts = name.split('/') + if len(parts) >= 5 and parts[1] == 'huggingface.co': + repo_id = f'{parts[2]}/{parts[3]}' + filename = '/'.join(parts[4:]) + fn = hf_hub_download(repo_id=repo_id, filename=filename, cache_dir=shared.opts.hfcache_dir) + return network.NetworkOnDisk(name, fn) + return None + + +def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=None): + networks_on_disk: list[network.NetworkOnDisk] = [available_network_aliases.get(name, None) for name in names] + if any(x is None for x in networks_on_disk): + list_available_networks() + networks_on_disk: list[network.NetworkOnDisk] = [available_network_aliases.get(name, None) for name in names] + for i in range(len(names)): + if names[i].startswith('/'): + networks_on_disk[i] = network_download(names[i]) + failed_to_load_networks = [] + recompile_model, skip_lora_load = maybe_recompile_model(names, te_multipliers) + + loaded_networks.clear() + diffuser_loaded.clear() + diffuser_scales.clear() + t0 = time.time() + + for i, (network_on_disk, name) in enumerate(zip(networks_on_disk, names)): + net = None + if network_on_disk is not None: + shorthash = getattr(network_on_disk, 'shorthash', '').lower() + if debug: + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') + try: + if recompile_model: + shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}") + if shared.opts.lora_force_diffusers or lora_overrides.check_override(shorthash): # OpenVINO only works with Diffusers LoRa loading + net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier) + else: + net = load_safetensors(name, network_on_disk) + if net is not None: + net.mentioned_name = name + network_on_disk.read_hash() + except Exception as e: + shared.log.error(f'Network load: type=LoRA file="{network_on_disk.filename}" {e}') + if debug: + errors.display(e, 'LoRA') + continue + if net is None: + failed_to_load_networks.append(name) + shared.log.error(f'Network load: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') + continue + if hasattr(shared.sd_model, 'embedding_db'): + shared.sd_model.embedding_db.load_diffusers_embedding(None, net.bundle_embeddings) + net.te_multiplier = te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier + net.unet_multiplier = unet_multipliers[i] if unet_multipliers else shared.opts.extra_networks_default_multiplier + net.dyn_dim = dyn_dims[i] if dyn_dims else shared.opts.extra_networks_default_multiplier + loaded_networks.append(net) + + while len(lora_cache) > shared.opts.lora_in_memory_limit: + name = next(iter(lora_cache)) + lora_cache.pop(name, None) + + if not skip_lora_load and len(diffuser_loaded) > 0: + shared.log.debug(f'Network load: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') + try: + t0 = time.time() + shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) + if shared.opts.lora_fuse_diffusers and not lora_overrides.check_fuse(): + shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling + shared.sd_model.unload_lora_weights() + timer.activate += time.time() - t0 + except Exception as e: + shared.log.error(f'Network load: type=LoRA {e}') + if debug: + errors.display(e, 'LoRA') + + if len(loaded_networks) > 0 and debug: + shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in loaded_networks]} cache={list(lora_cache)}') + + if recompile_model: + shared.log.info("Network load: type=LoRA recompiling model") + backup_lora_model = shared.compiled_model_state.lora_model + if 'Model' in shared.opts.cuda_compile: + shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) + shared.compiled_model_state.lora_model = backup_lora_model + + if len(loaded_networks) > 0: + devices.torch_gc() + + timer.load = time.time() - t0 diff --git a/modules/lora/network_overrides.py b/modules/lora/lora_overrides.py similarity index 100% rename from modules/lora/network_overrides.py rename to modules/lora/lora_overrides.py diff --git a/modules/lora/network.py b/modules/lora/network.py index c4768d9ad..f6d93009c 100644 --- a/modules/lora/network.py +++ b/modules/lora/network.py @@ -91,8 +91,10 @@ class NetworkOnDisk: self.set_hash(hashes.sha256(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or '') def get_alias(self): - import modules.lora.networks as networks - return self.name if shared.opts.lora_preferred_name == "filename" or self.alias.lower() in networks.forbidden_network_aliases else self.alias + if shared.opts.lora_preferred_name == "filename": + return self.name + else: + return self.alias class Network: # LoraModule diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 4903e5a6f..44dad1afd 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -1,583 +1,12 @@ -from typing import Union, List from contextlib import nullcontext -import os -import re import time -import concurrent -import torch -import diffusers.models.lora import rich.progress as rp - -from modules.lora import lora_timers, network, lora_convert, network_overrides -from modules.lora import network_lora, network_hada, network_ia3, network_oft, network_lokr, network_full, network_norm, network_glora -from modules.lora.extra_networks_lora import ExtraNetworkLora -from modules import shared, devices, sd_models, sd_models_compile, errors, files_cache, model_quant +from modules.lora.lora_common import timer, debug, loaded_networks, previously_loaded_networks +from modules.lora.lora_apply import network_apply_weights, network_apply_direct, network_backup_weights, network_calc_weights +from modules import shared, devices, sd_models -debug = os.environ.get('SD_LORA_DEBUG', None) is not None -extra_network_lora = ExtraNetworkLora() -available_networks = {} -available_network_aliases = {} -loaded_networks: List[network.Network] = [] -previously_loaded_networks: List[network.Network] = [] applied_layers: list[str] = [] -bnb = None -lora_cache = {} -diffuser_loaded = [] -diffuser_scales = [] -available_network_hash_lookup = {} -forbidden_network_aliases = {} -re_network_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") -timer = lora_timers.Timer() -module_types = [ - network_lora.ModuleTypeLora(), - network_hada.ModuleTypeHada(), - network_ia3.ModuleTypeIa3(), - network_oft.ModuleTypeOFT(), - network_lokr.ModuleTypeLokr(), - network_full.ModuleTypeFull(), - network_norm.ModuleTypeNorm(), - network_glora.ModuleTypeGLora(), -] - -# section: load networks from disk - -def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_default_multiplier) -> Union[network.Network, None]: - t0 = time.time() - name = name.replace(".", "_") - shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') - if not shared.native: - return None - if not hasattr(shared.sd_model, 'load_lora_weights'): - shared.log.error(f'Network load: type=LoRA class={shared.sd_model.__class__} does not implement load lora') - return None - try: - shared.sd_model.load_lora_weights(network_on_disk.filename, adapter_name=name) - except Exception as e: - if 'already in use' in str(e): - pass - else: - if 'The following keys have not been correctly renamed' in str(e): - shared.log.error(f'Network load: type=LoRA name="{name}" diffusers unsupported format') - else: - shared.log.error(f'Network load: type=LoRA name="{name}" {e}') - if debug: - errors.display(e, "LoRA") - return None - if name not in diffuser_loaded: - diffuser_loaded.append(name) - diffuser_scales.append(lora_scale) - net = network.Network(name, network_on_disk) - net.mtime = os.path.getmtime(network_on_disk.filename) - timer.activate += time.time() - t0 - return net - - -def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: - if not shared.sd_loaded: - return None - - cached = lora_cache.get(name, None) - if debug: - shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') - if cached is not None: - return cached - net = network.Network(name, network_on_disk) - net.mtime = os.path.getmtime(network_on_disk.filename) - sd = sd_models.read_state_dict(network_on_disk.filename, what='network') - if shared.sd_model_type == 'f1': # if kohya flux lora, convert state_dict - sd = lora_convert._convert_kohya_flux_lora_to_diffusers(sd) or sd # pylint: disable=protected-access - if shared.sd_model_type == 'sd3': # if kohya flux lora, convert state_dict - try: - sd = lora_convert._convert_kohya_sd3_lora_to_diffusers(sd) or sd # pylint: disable=protected-access - except ValueError: # EAFP for diffusers PEFT keys - pass - lora_convert.assign_network_names_to_compvis_modules(shared.sd_model) - keys_failed_to_match = {} - matched_networks = {} - bundle_embeddings = {} - convert = lora_convert.KeyConvert() - for key_network, weight in sd.items(): - parts = key_network.split('.') - if parts[0] == "bundle_emb": - emb_name, vec_name = parts[1], key_network.split(".", 2)[-1] - emb_dict = bundle_embeddings.get(emb_name, {}) - emb_dict[vec_name] = weight - bundle_embeddings[emb_name] = emb_dict - continue - if len(parts) > 5: # messy handler for diffusers peft lora - key_network_without_network_parts = '_'.join(parts[:-2]) - if not key_network_without_network_parts.startswith('lora_'): - key_network_without_network_parts = 'lora_' + key_network_without_network_parts - network_part = '.'.join(parts[-2:]).replace('lora_A', 'lora_down').replace('lora_B', 'lora_up') - else: - key_network_without_network_parts, network_part = key_network.split(".", 1) - key, sd_module = convert(key_network_without_network_parts) - if sd_module is None: - keys_failed_to_match[key_network] = key - continue - if key not in matched_networks: - matched_networks[key] = network.NetworkWeights(network_key=key_network, sd_key=key, w={}, sd_module=sd_module) - matched_networks[key].w[network_part] = weight - network_types = [] - for key, weights in matched_networks.items(): - net_module = None - for nettype in module_types: - net_module = nettype.create_module(net, weights) - if net_module is not None: - network_types.append(nettype.__class__.__name__) - break - if net_module is None: - shared.log.error(f'LoRA unhandled: name={name} key={key} weights={weights.w.keys()}') - else: - net.modules[key] = net_module - if len(keys_failed_to_match) > 0: - shared.log.warning(f'Network load: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') - if debug: - shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') - else: - shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} direct={shared.opts.lora_fuse_diffusers}') - if len(matched_networks) == 0: - return None - lora_cache[name] = net - net.bundle_embeddings = bundle_embeddings - return net - - -def maybe_recompile_model(names, te_multipliers): - recompile_model = False - skip_lora_load = False - if shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled: - if len(names) == len(shared.compiled_model_state.lora_model): - for i, name in enumerate(names): - if shared.compiled_model_state.lora_model[ - i] != f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}": - recompile_model = True - shared.compiled_model_state.lora_model = [] - break - if not recompile_model: - skip_lora_load = True - if len(loaded_networks) > 0 and debug: - shared.log.debug('Model Compile: Skipping LoRa loading') - return recompile_model, skip_lora_load - else: - recompile_model = True - shared.compiled_model_state.lora_model = [] - if recompile_model: - backup_cuda_compile = shared.opts.cuda_compile - backup_scheduler = getattr(shared.sd_model, "scheduler", None) - sd_models.unload_model_weights(op='model') - shared.opts.cuda_compile = [] - sd_models.reload_model_weights(op='model') - shared.opts.cuda_compile = backup_cuda_compile - if backup_scheduler is not None: - shared.sd_model.scheduler = backup_scheduler - return recompile_model, skip_lora_load - - -def list_available_networks(): - t0 = time.time() - available_networks.clear() - available_network_aliases.clear() - forbidden_network_aliases.clear() - available_network_hash_lookup.clear() - forbidden_network_aliases.update({"none": 1, "Addams": 1}) - if not os.path.exists(shared.cmd_opts.lora_dir): - shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') - - def add_network(filename): - if not os.path.isfile(filename): - return - name = os.path.splitext(os.path.basename(filename))[0] - name = name.replace('.', '_') - try: - entry = network.NetworkOnDisk(name, filename) - available_networks[entry.name] = entry - if entry.alias in available_network_aliases: - forbidden_network_aliases[entry.alias.lower()] = 1 - if shared.opts.lora_preferred_name == 'filename': - available_network_aliases[entry.name] = entry - else: - available_network_aliases[entry.alias] = entry - if entry.shorthash: - available_network_hash_lookup[entry.shorthash] = entry - except OSError as e: # should catch FileNotFoundError and PermissionError etc. - shared.log.error(f'LoRA: filename="{filename}" {e}') - - candidates = sorted(files_cache.list_files(shared.cmd_opts.lora_dir, ext_filter=[".pt", ".ckpt", ".safetensors"])) - with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - for fn in candidates: - executor.submit(add_network, fn) - t1 = time.time() - timer.list = t1 - t0 - shared.log.info(f'Available LoRAs: path="{shared.cmd_opts.lora_dir}" items={len(available_networks)} folders={len(forbidden_network_aliases)} time={t1 - t0:.2f}') - - -def network_download(name): - from huggingface_hub import hf_hub_download - if os.path.exists(name): - return network.NetworkOnDisk(name, name) - parts = name.split('/') - if len(parts) >= 5 and parts[1] == 'huggingface.co': - repo_id = f'{parts[2]}/{parts[3]}' - filename = '/'.join(parts[4:]) - fn = hf_hub_download(repo_id=repo_id, filename=filename, cache_dir=shared.opts.hfcache_dir) - return network.NetworkOnDisk(name, fn) - return None - - -def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=None): - networks_on_disk: list[network.NetworkOnDisk] = [available_network_aliases.get(name, None) for name in names] - if any(x is None for x in networks_on_disk): - list_available_networks() - networks_on_disk: list[network.NetworkOnDisk] = [available_network_aliases.get(name, None) for name in names] - for i in range(len(names)): - if names[i].startswith('/'): - networks_on_disk[i] = network_download(names[i]) - failed_to_load_networks = [] - recompile_model, skip_lora_load = maybe_recompile_model(names, te_multipliers) - - loaded_networks.clear() - diffuser_loaded.clear() - diffuser_scales.clear() - t0 = time.time() - - for i, (network_on_disk, name) in enumerate(zip(networks_on_disk, names)): - net = None - if network_on_disk is not None: - shorthash = getattr(network_on_disk, 'shorthash', '').lower() - if debug: - shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') - try: - if recompile_model: - shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier}") - if shared.opts.lora_force_diffusers or network_overrides.check_override(shorthash): # OpenVINO only works with Diffusers LoRa loading - net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier) - else: - net = load_safetensors(name, network_on_disk) - if net is not None: - net.mentioned_name = name - network_on_disk.read_hash() - except Exception as e: - shared.log.error(f'Network load: type=LoRA file="{network_on_disk.filename}" {e}') - if debug: - errors.display(e, 'LoRA') - continue - if net is None: - failed_to_load_networks.append(name) - shared.log.error(f'Network load: type=LoRA name="{name}" detected={network_on_disk.sd_version if network_on_disk is not None else None} failed') - continue - if hasattr(shared.sd_model, 'embedding_db'): - shared.sd_model.embedding_db.load_diffusers_embedding(None, net.bundle_embeddings) - net.te_multiplier = te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier - net.unet_multiplier = unet_multipliers[i] if unet_multipliers else shared.opts.extra_networks_default_multiplier - net.dyn_dim = dyn_dims[i] if dyn_dims else shared.opts.extra_networks_default_multiplier - loaded_networks.append(net) - - while len(lora_cache) > shared.opts.lora_in_memory_limit: - name = next(iter(lora_cache)) - lora = lora_cache.pop(name, None) - del lora - - if not skip_lora_load and len(diffuser_loaded) > 0: - shared.log.debug(f'Network load: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') - try: - t0 = time.time() - shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) - if shared.opts.lora_fuse_diffusers and not network_overrides.check_fuse(): - shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling - shared.sd_model.unload_lora_weights() - timer.activate += time.time() - t0 - except Exception as e: - shared.log.error(f'Network load: type=LoRA {e}') - if debug: - errors.display(e, 'LoRA') - - if len(loaded_networks) > 0 and debug: - shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in loaded_networks]} cache={list(lora_cache)}') - - if recompile_model: - shared.log.info("Network load: type=LoRA recompiling model") - backup_lora_model = shared.compiled_model_state.lora_model - if 'Model' in shared.opts.cuda_compile: - shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) - shared.compiled_model_state.lora_model = backup_lora_model - - if len(loaded_networks) > 0: - devices.torch_gc() - - timer.load = time.time() - t0 - - -# section: process loaded networks - -def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], network_layer_name: str, wanted_names: tuple): - global bnb # pylint: disable=W0603 - backup_size = 0 - if len(loaded_networks) > 0 and network_layer_name is not None and any([net.modules.get(network_layer_name, None) for net in loaded_networks]): # noqa: C419 # pylint: disable=R1729 - t0 = time.time() - - weights_backup = getattr(self, "network_weights_backup", None) - bias_backup = getattr(self, "network_bias_backup", None) - if weights_backup is not None or bias_backup is not None: - if (shared.opts.lora_fuse_diffusers and not isinstance(weights_backup, bool)) or (not shared.opts.lora_fuse_diffusers and isinstance(weights_backup, bool)): # invalidate so we can change direct/backup on-the-fly - weights_backup = None - bias_backup = None - self.network_weights_backup = weights_backup - self.network_bias_backup = bias_backup - - if weights_backup is None and wanted_names != (): # pylint: disable=C1803 - weight = getattr(self, 'weight', None) - self.network_weights_backup = None - if getattr(weight, "quant_type", None) in ['nf4', 'fp4']: - if bnb is None: - bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) - if bnb is not None: - with devices.inference_context(): - if shared.opts.lora_fuse_diffusers: - self.network_weights_backup = True - else: - self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,) - self.quant_state = weight.quant_state - self.quant_type = weight.quant_type - self.blocksize = weight.blocksize - else: - if shared.opts.lora_fuse_diffusers: - self.network_weights_backup = True - else: - weights_backup = weight.clone() - self.network_weights_backup = weights_backup.to(devices.cpu) - else: - if shared.opts.lora_fuse_diffusers: - self.network_weights_backup = True - else: - self.network_weights_backup = weight.clone().to(devices.cpu) - - if bias_backup is None: - if getattr(self, 'bias', None) is not None: - if shared.opts.lora_fuse_diffusers: - self.network_bias_backup = True - else: - bias_backup = self.bias.clone() - bias_backup = bias_backup.to(devices.cpu) - - if getattr(self, 'network_weights_backup', None) is not None: - backup_size += self.network_weights_backup.numel() * self.network_weights_backup.element_size() if isinstance(self.network_weights_backup, torch.Tensor) else 0 - if getattr(self, 'network_bias_backup', None) is not None: - backup_size += self.network_bias_backup.numel() * self.network_bias_backup.element_size() if isinstance(self.network_bias_backup, torch.Tensor) else 0 - timer.backup += time.time() - t0 - return backup_size - - -def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], network_layer_name: str, use_previous: bool = False): - if shared.opts.diffusers_offload_mode == "none": - try: - self.to(devices.device) - except Exception: - pass - batch_updown = None - batch_ex_bias = None - loaded = loaded_networks if not use_previous else previously_loaded_networks - for net in loaded: - module = net.modules.get(network_layer_name, None) - if module is None: - continue - try: - t0 = time.time() - try: - weight = self.weight.to(devices.device) - except Exception: - weight = self.weight - - updown, ex_bias = module.calc_updown(weight) - if updown is not None: - if batch_updown is not None: - batch_updown += updown.to(batch_updown.device) - else: - batch_updown = updown.to(devices.device) - if ex_bias is not None: - if batch_ex_bias: - batch_ex_bias += ex_bias.to(batch_ex_bias.device) - else: - batch_ex_bias = ex_bias.to(devices.device) - timer.calc += time.time() - t0 - - if shared.opts.diffusers_offload_mode == "sequential": - t0 = time.time() - if batch_updown is not None: - batch_updown = batch_updown.to(devices.cpu) - if batch_ex_bias is not None: - batch_ex_bias = batch_ex_bias.to(devices.cpu) - t1 = time.time() - timer.move += t1 - t0 - except RuntimeError as e: - extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 - if debug: - module_name = net.modules.get(network_layer_name, None) - shared.log.error(f'LoRA apply weight name="{net.name}" module="{module_name}" layer="{network_layer_name}" {e}') - errors.display(e, 'LoRA') - raise RuntimeError('LoRA apply weight') from e - continue - return batch_updown, batch_ex_bias - - -def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, device: torch.device = None): - if lora_weights is None: - return None - if deactivate: - lora_weights *= -1 - if model_weights is None: # weights are used if provided-from-backup else use self.weight - model_weights = self.weight - # TODO lora: add other quantization types - weight = None - device = device or devices.device - if self.__class__.__name__ == 'Linear4bit' and bnb is not None: - try: - dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) - new_weight = dequant_weight.to(device) + lora_weights.to(device) - weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) - except Exception as e: - shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') - else: - try: - new_weight = model_weights.to(device) + lora_weights.to(device) - except Exception: - new_weight = model_weights + lora_weights # try without device cast - weight = torch.nn.Parameter(new_weight, requires_grad=False) - try: - # weight = weight.to(device=devices.device) # required since quantization happens only during .to call, not during params creation - pass - except Exception: - pass # may fail if weights is meta tensor - return weight - - -def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False, device: torch.device = None): - weights_backup = getattr(self, "network_weights_backup", False) - bias_backup = getattr(self, "network_bias_backup", False) - if not isinstance(weights_backup, bool): # remove previous backup if we switched settings - weights_backup = True - if not isinstance(bias_backup, bool): - bias_backup = True - if not weights_backup and not bias_backup: - return - device = device or devices.device - t0 = time.time() - - if weights_backup: - if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 - updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable - if updown is not None: - weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate, device=device) - if weight is not None: - self.weight = weight - - if bias_backup: - if ex_bias is not None: - bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, device=device) - if bias is not None: - self.bias = bias - - if hasattr(self, "qweight") and hasattr(self, "freeze"): - self.freeze() - - timer.apply += time.time() - t0 - - -def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, orig_device: torch.device, deactivate: bool = False): - weights_backup = getattr(self, "network_weights_backup", None) - bias_backup = getattr(self, "network_bias_backup", None) - if weights_backup is None and bias_backup is None: - return - t0 = time.time() - - if weights_backup is not None: - self.weight = None - if updown is not None and len(weights_backup.shape) == 4 and weights_backup.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 - updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable - if updown is not None: - weight = network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate) - if weight is not None: - self.weight = weight - else: - self.weight = torch.nn.Parameter(weights_backup.to(device=orig_device), requires_grad=False) - - if bias_backup is not None: - self.bias = None - if ex_bias is not None: - bias = network_add_weights(self, model_weights=weights_backup, lora_weights=ex_bias, deactivate=deactivate) - if bias: - self.weight = bias - else: - self.bias = torch.nn.Parameter(bias_backup.to(device=orig_device), requires_grad=False) - - if hasattr(self, "qweight") and hasattr(self, "freeze"): - self.freeze() - - timer.apply += time.time() - t0 - - -def network_deactivate(include=[], exclude=[]): - if not shared.opts.lora_fuse_diffusers or shared.opts.lora_force_diffusers: - return - t0 = time.time() - sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility - if shared.opts.diffusers_offload_mode == "sequential": - sd_models.disable_offload(sd_model) - sd_models.move_model(sd_model, device=devices.cpu) - modules = {} - - components = include if len(include) > 0 else ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer'] - components = [x for x in components if x not in exclude] - active_components = [] - for name in components: - component = getattr(sd_model, name, None) - if component is not None and hasattr(component, 'named_modules'): - modules[name] = list(component.named_modules()) - active_components.append(name) - total = sum(len(x) for x in modules.values()) - if shared.opts.lora_apply_gpu: - device = devices.device - else: - device = devices.cpu - if len(previously_loaded_networks) > 0 and debug: - pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=deactivate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) - task = pbar.add_task(description='', total=total) - else: - task = None - pbar = nullcontext() - with devices.inference_context(), pbar: - applied_layers.clear() - for component in modules.keys(): - orig_device = getattr(sd_model, component, None).device - for _, module in modules[component]: - network_layer_name = getattr(module, 'network_layer_name', None) - if shared.state.interrupted or network_layer_name is None: - if task is not None: - pbar.update(task, advance=1) - continue - batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, use_previous=True) - if shared.opts.lora_fuse_diffusers: - network_apply_direct(module, batch_updown, batch_ex_bias, deactivate=True) - else: - network_apply_weights(module, batch_updown, batch_ex_bias, orig_device, deactivate=True) - if batch_updown is not None or batch_ex_bias is not None: - applied_layers.append(network_layer_name) - del batch_updown, batch_ex_bias - module.network_current_names = () - try: - module.to(device) - except Exception: - pass - if task is not None: - pbar.update(task, advance=1, description=f'networks={len(previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') - - timer.deactivate = time.time() - t0 - if debug and len(previously_loaded_networks) > 0: - shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') - modules.clear() - if shared.opts.diffusers_offload_mode == "sequential": - sd_models.set_diffuser_offload(sd_model, op="model") def network_activate(include=[], exclude=[]): @@ -604,10 +33,7 @@ def network_activate(include=[], exclude=[]): pbar = nullcontext() applied_weight = 0 applied_bias = 0 - if shared.opts.lora_apply_gpu: - device = devices.device - else: - device = devices.cpu + device = devices.device if shared.opts.lora_apply_gpu else devices.cpu with devices.inference_context(), pbar: wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) if len(loaded_networks) > 0 else () applied_layers.clear() @@ -624,21 +50,18 @@ def network_activate(include=[], exclude=[]): backup_size += network_backup_weights(module, network_layer_name, wanted_names) batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name) if shared.opts.lora_fuse_diffusers: - network_apply_direct(module, batch_updown, batch_ex_bias, device) + network_apply_direct(module, batch_updown, batch_ex_bias, device=device) else: - network_apply_weights(module, batch_updown, batch_ex_bias, orig_device) + network_apply_weights(module, batch_updown, batch_ex_bias, device=orig_device) if batch_updown is not None or batch_ex_bias is not None: applied_layers.append(network_layer_name) + # module.to(device) # TODO maybe if batch_updown is not None: applied_weight += 1 if batch_ex_bias is not None: applied_bias += 1 del batch_updown, batch_ex_bias module.network_current_names = wanted_names - try: - module.to(device) - except Exception: - pass if task is not None: pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size}') @@ -646,8 +69,65 @@ def network_activate(include=[], exclude=[]): pbar.remove_task(task) # hide progress bar for no action timer.activate += time.time() - t0 if debug and len(loaded_networks) > 0: - shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} device={device} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} device={device} time={timer.summary}') modules.clear() if len(loaded_networks) > 0 and (applied_weight > 0 or applied_bias > 0): if shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") + + +def network_deactivate(include=[], exclude=[]): + if not shared.opts.lora_fuse_diffusers or shared.opts.lora_force_diffusers: + return + t0 = time.time() + sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility + if shared.opts.diffusers_offload_mode == "sequential": + sd_models.disable_offload(sd_model) + sd_models.move_model(sd_model, device=devices.cpu) + modules = {} + + components = include if len(include) > 0 else ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer'] + components = [x for x in components if x not in exclude] + active_components = [] + for name in components: + component = getattr(sd_model, name, None) + if component is not None and hasattr(component, 'named_modules'): + modules[name] = list(component.named_modules()) + active_components.append(name) + total = sum(len(x) for x in modules.values()) + device = devices.device if shared.opts.lora_apply_gpu else devices.cpu + if len(previously_loaded_networks) > 0 and debug: + pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=deactivate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) + task = pbar.add_task(description='', total=total) + else: + task = None + pbar = nullcontext() + with devices.inference_context(), pbar: + applied_layers.clear() + for component in modules.keys(): + orig_device = getattr(sd_model, component, None).device + for _, module in modules[component]: + network_layer_name = getattr(module, 'network_layer_name', None) + if shared.state.interrupted or network_layer_name is None: + if task is not None: + pbar.update(task, advance=1) + continue + batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, use_previous=True) + if shared.opts.lora_fuse_diffusers: + network_apply_direct(module, batch_updown, batch_ex_bias, device=device, deactivate=True) + else: + network_apply_weights(module, batch_updown, batch_ex_bias, device=orig_device, deactivate=True) + if batch_updown is not None or batch_ex_bias is not None: + # module.to(device) # TODO maybe + applied_layers.append(network_layer_name) + del batch_updown, batch_ex_bias + module.network_current_names = () + if task is not None: + pbar.update(task, advance=1, description=f'networks={len(previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') + + timer.deactivate = time.time() - t0 + if debug and len(previously_loaded_networks) > 0: + shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + modules.clear() + if shared.opts.diffusers_offload_mode == "sequential": + sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/model_auraflow.py b/modules/model_auraflow.py index 83320040b..790095cae 100644 --- a/modules/model_auraflow.py +++ b/modules/model_auraflow.py @@ -17,5 +17,5 @@ def load_auraflow(checkpoint_info, diffusers_load_config={}): cache_dir = shared.opts.diffusers_dir, **diffusers_load_config, ) - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_flux.py b/modules/model_flux.py index 74ba5c8dd..12ad0a471 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -342,6 +342,5 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch vae = None for k in kwargs.keys(): kwargs[k] = None - devices.torch_gc() - + devices.torch_gc(force=True) return pipe diff --git a/modules/model_kolors.py b/modules/model_kolors.py index 932763b4b..b6c35c85a 100644 --- a/modules/model_kolors.py +++ b/modules/model_kolors.py @@ -23,5 +23,5 @@ def load_kolors(_checkpoint_info, diffusers_load_config={}): **diffusers_load_config, ) pipe.vae.config.force_upcast = True - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_lumina.py b/modules/model_lumina.py index f9b09f23b..f19fcd7da 100644 --- a/modules/model_lumina.py +++ b/modules/model_lumina.py @@ -21,7 +21,7 @@ def load_lumina(_checkpoint_info, diffusers_load_config={}): cache_dir = shared.opts.diffusers_dir, **diffusers_load_config, ) - devices.torch_gc() + devices.torch_gc(force=True) return pipe @@ -40,4 +40,5 @@ def load_lumina2(checkpoint_info, diffusers_load_config={}): if ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): kwargs['text_encoder'] = transformers.AutoModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype, **quant_args) sd_model = diffusers.Lumina2Text2ImgPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config, **quant_args, **kwargs) + devices.torch_gc(force=True) return sd_model diff --git a/modules/model_meissonic.py b/modules/model_meissonic.py index 69ceab458..d705a32d9 100644 --- a/modules/model_meissonic.py +++ b/modules/model_meissonic.py @@ -33,5 +33,5 @@ def load_meissonic(checkpoint_info, diffusers_load_config={}): diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["meissonic"] = PipelineMeissonic diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["meissonic"] = PipelineMeissonicImg2Img diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["meissonic"] = PipelineMeissonicInpaint - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_omnigen.py b/modules/model_omnigen.py index a08ad4ed5..b7b6e3546 100644 --- a/modules/model_omnigen.py +++ b/modules/model_omnigen.py @@ -20,12 +20,5 @@ def load_omnigen(checkpoint_info, diffusers_load_config={}): # pylint: disable=u if shared.opts.diffusers_eval: pipe.model.eval() pipe.vae.to(devices.device, dtype=devices.dtype) - devices.torch_gc() - - # register - # from diffusers import pipelines - # pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["omnigen"] = pipe.__class__ - # pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["omnigen"] = pipe.__class__ - # pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["omnigen"] = pipe.__class__ - + devices.torch_gc(force=True) return pipe diff --git a/modules/model_pixart.py b/modules/model_pixart.py index c017cc468..0757a1216 100644 --- a/modules/model_pixart.py +++ b/modules/model_pixart.py @@ -26,5 +26,5 @@ def load_pixart(checkpoint_info, diffusers_load_config={}): **kwargs, **diffusers_load_config, ) - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_sana.py b/modules/model_sana.py index a31985ab6..7f39f17e0 100644 --- a/modules/model_sana.py +++ b/modules/model_sana.py @@ -73,6 +73,5 @@ def load_sana(checkpoint_info, kwargs={}): pipe.transformer.eval() t1 = time.time() shared.log.debug(f'Load model: type=Sana target={devices.dtype} te={pipe.text_encoder.dtype} transformer={pipe.transformer.dtype} vae={pipe.vae.dtype} time={t1-t0:.2f}') - - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 5b8006c2a..e3774b291 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -156,5 +156,5 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): config=config, **kwargs, ) - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_stablecascade.py b/modules/model_stablecascade.py index 0c767d33b..fc143e8e7 100644 --- a/modules/model_stablecascade.py +++ b/modules/model_stablecascade.py @@ -155,6 +155,7 @@ def load_cascade_combined(checkpoint_info, diffusers_load_config): latent_dim_scale=sd_model.decoder_pipe.config.latent_dim_scale, ) + devices.torch_gc(force=True) shared.log.debug(f'StableCascade combined: {sd_model.__class__.__name__}') return sd_model diff --git a/modules/sd_models.py b/modules/sd_models.py index b500c51b7..2d4451db1 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1053,10 +1053,10 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model', def clear_caches(): # shared.log.debug('Cache clear') if not shared.opts.lora_legacy: - from modules.lora import networks - networks.loaded_networks.clear() - networks.previously_loaded_networks.clear() - networks.lora_cache.clear() + from modules.lora import lora_common, lora_load + lora_common.loaded_networks.clear() + lora_common.previously_loaded_networks.clear() + lora_load.lora_cache.clear() from modules import prompt_parser_diffusers prompt_parser_diffusers.cache.clear() diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index ef40092ec..435ff95ae 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -437,7 +437,6 @@ def create_html(search_text, sort_column): def create_ui(): - import modules.ui extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "user", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all", visible=False) extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False, container=False) extensions_update_list = gr.Text(elem_id="extensions_update_list", visible=False, container=False) diff --git a/modules/ui_extra_networks_lora.py b/modules/ui_extra_networks_lora.py index 9dd1b3573..194f16b41 100644 --- a/modules/ui_extra_networks_lora.py +++ b/modules/ui_extra_networks_lora.py @@ -1,8 +1,8 @@ import os import json import concurrent -import modules.lora.networks as networks from modules import shared, ui_extra_networks +from modules.lora import lora_load debug = os.environ.get('SD_LORA_DEBUG', None) is not None @@ -14,7 +14,7 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): self.list_time = 0 def refresh(self): - networks.list_available_networks() + lora_load.list_available_networks() @staticmethod def get_tags(l, info): @@ -78,9 +78,9 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): return clean_tags def create_item(self, name): - l = networks.available_networks.get(name) + l = lora_load.available_networks.get(name) if l is None: - shared.log.warning(f'Networks: type=lora registered={len(list(networks.available_networks))} file="{name}" not registered') + shared.log.warning(f'Networks: type=lora registered={len(list(lora_load.available_networks))} file="{name}" not registered') return None try: # path, _ext = os.path.splitext(l.filename) @@ -111,7 +111,7 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def list_items(self): items = [] with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - future_items = {executor.submit(self.create_item, net): net for net in networks.available_networks} + future_items = {executor.submit(self.create_item, net): net for net in lora_load.available_networks} for future in concurrent.futures.as_completed(future_items): item = future.result() if item is not None: diff --git a/modules/ui_gallery.py b/modules/ui_gallery.py index a1f317caa..40bcace03 100644 --- a/modules/ui_gallery.py +++ b/modules/ui_gallery.py @@ -3,7 +3,7 @@ from datetime import datetime from urllib.parse import unquote import gradio as gr from PIL import Image -from modules import shared, ui_symbols, ui_common, images, ui_control_helpers +from modules import shared, ui_symbols, ui_common, images, video from modules.ui_components import ToolButton def read_media(fn): @@ -13,7 +13,7 @@ def read_media(fn): return [[], None, '', '', f'Media not found: {fn}'] stat = os.stat(fn) if fn.lower().endswith('.mp4'): - frames, fps, duration, w, h, codec, _frame = ui_control_helpers.get_video_params(fn) + frames, fps, duration, w, h, codec, _frame = video.get_video_params(fn) geninfo = '' log = f'''

Video {w} x {h} diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index ea60c846a..431c359b8 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -1,7 +1,7 @@ from functools import wraps import torch import torch._dynamo.device_interface -from modules import rocm, zluda, shared +from modules import shared, rocm, zluda # pylint: disable=unused-import MEM_BUS_WIDTH = { diff --git a/webui.py b/webui.py index ecfafe3f1..e44c5c2b9 100644 --- a/webui.py +++ b/webui.py @@ -88,8 +88,8 @@ def initialize(): timer.startup.record("models") if not shared.opts.lora_legacy: - import modules.lora.networks as lora_networks - lora_networks.list_available_networks() + from modules.lora import lora_load + lora_load.list_available_networks() timer.startup.record("lora") shared.prompt_styles.reload() From c208175c0f02b78070bf49e30ea529580ebc0286 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 00:05:15 -0400 Subject: [PATCH 106/122] fix flux partial quantization Signed-off-by: Vladimir Mandic --- modules/model_flux.py | 24 +++++++++++------------- 1 file changed, 11 insertions(+), 13 deletions(-) diff --git a/modules/model_flux.py b/modules/model_flux.py index 12ad0a471..3dfe83ff2 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -109,16 +109,14 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu def load_quants(kwargs, repo_id, cache_dir, allow_quant): try: - quant_args = model_quant.create_config(allow=allow_quant) - if not quant_args: - return kwargs if 'transformer' not in kwargs and (('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization) or ('Transformer' in shared.opts.bnb_quantization or 'Transformer' in shared.opts.torchao_quantization or 'Transformer' in shared.opts.quanto_quantization)): - kwargs['transformer'] = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) - quant_args = model_quant.create_config(allow=allow_quant, module='TE') - if not quant_args: - return kwargs + quant_args = model_quant.create_config(allow=allow_quant, module='Transformer') + if quant_args: + kwargs['transformer'] = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) if 'text_encoder_2' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization or 'TE' in shared.opts.quanto_quantization): - kwargs['text_encoder_2'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) + quant_args = model_quant.create_config(allow=allow_quant, module='TE') + if quant_args: + kwargs['text_encoder_2'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, torch_dtype=devices.dtype, **quant_args) except Exception as e: shared.log.error(f'Quantization: {e}') errors.display(e, 'Quantization:') @@ -209,9 +207,9 @@ def load_transformer(file_path): # triggered by opts.sd_unet change def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change - quant = model_quant.get_quant(checkpoint_info.path) + prequantized = model_quant.get_quant(checkpoint_info.path) repo_id = sd_models.path_to_repo(checkpoint_info.name) - shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" te="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') + shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" te="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={prequantized} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') debug(f'Load model: type=FLUX config={diffusers_load_config}') modelloader.hf_login() @@ -267,7 +265,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch errors.display(e, 'FLUX VAE:') # load quantized components if any - if quant == 'nf4': + if prequantized == 'nf4': try: from modules.model_flux_nf4 import load_flux_nf4 _transformer, _text_encoder = load_flux_nf4(checkpoint_info) @@ -279,7 +277,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.log.error(f"Load model: type=FLUX failed to load NF4 components: {e}") if debug: errors.display(e, 'FLUX NF4:') - if quant == 'qint8' or quant == 'qint4': + if prequantized == 'qint8' or prequantized == 'qint4': try: _transformer, _text_encoder = load_flux_quanto(checkpoint_info) if _transformer is not None: @@ -325,7 +323,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch except Exception: pass - allow_quant = 'gguf' not in (sd_unet.loaded_unet or '') and (quant is None or quant == 'none') + allow_quant = 'gguf' not in (sd_unet.loaded_unet or '') and (prequantized is None or prequantized == 'none') fn = checkpoint_info.path if (fn is None) or (not os.path.exists(fn) or os.path.isdir(fn)): kwargs = load_quants(kwargs, repo_id, cache_dir=shared.opts.diffusers_dir, allow_quant=allow_quant) From 6430f7006f7b933b24f32f935306fc5654a9b41b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 13:12:00 -0400 Subject: [PATCH 107/122] add monitor cli option and finish lora refactor Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + installer.py | 3 +- launch.py | 23 +++++--- modules/cmd_args.py | 1 + modules/lora/extra_networks_lora.py | 63 ++++++++++++---------- modules/lora/lora_apply.py | 84 ++++++++++++----------------- modules/lora/lora_load.py | 54 +++++++++---------- modules/lora/networks.py | 49 +++++++++-------- modules/memstats.py | 14 ++--- modules/model_quant.py | 22 +++++--- modules/processing_diffusers.py | 6 +-- modules/sd_models.py | 1 - modules/sd_offload.py | 4 +- 13 files changed, 169 insertions(+), 156 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 9af36d7bd..0ca2ae7ec 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -122,6 +122,7 @@ Plus... - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - add Sage Attention support - **Other** + - new command line option `--monitor PERIOD` to monitor CPU and GPU memory ever n seconds - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method - **upscale**: new experimental support for `libvips` upscaling - **quantization**: add support for `optimum-quanto` on-the-fly quantization during load for all models diff --git a/installer.py b/installer.py index 289b2d5b3..7154b5b3e 100644 --- a/installer.py +++ b/installer.py @@ -517,7 +517,7 @@ def check_python(supported_minors=[9, 10, 11, 12], reason=None): log.error(f"Python version incompatible: {sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro} required 3.{supported_minors}") if reason is not None: log.error(reason) - if not args.ignore: + if not args.ignore and not args.experimental: sys.exit(1) if int(sys.version_info.minor) == 12: os.environ.setdefault('SETUPTOOLS_USE_DISTUTILS', 'local') # hack for python 3.11 setuptools @@ -1492,6 +1492,7 @@ def add_args(parser): group_log.add_argument("--log", type=str, default=os.environ.get("SD_LOG", None), help="Set log file, default: %(default)s") group_log.add_argument('--debug', default=os.environ.get("SD_DEBUG",False), action='store_true', help="Run installer with debug logging, default: %(default)s") group_log.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") + group_log.add_argument("--monitor", default=os.environ.get("SD_PROFILE", 0), help="Run memory monitor, default: %(default)s") group_log.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help="Mount API docs, default: %(default)s") group_log.add_argument("--api-log", default=os.environ.get("SD_APILOG", True), action='store_true', help="Log all API requests") diff --git a/launch.py b/launch.py index c80840036..d00b9ef22 100755 --- a/launch.py +++ b/launch.py @@ -150,10 +150,14 @@ def run_extension_installer(ext_dir): # compatbility function installer.run_extension_installer(ext_dir) -def get_memory_stats(): - from modules.memstats import ram_stats - res = ram_stats() - return f'{res["used"]}/{res["total"]}' +def get_memory_stats(detailed:bool=False): + from modules.memstats import ram_stats, memory_stats + if not detailed: + res = ram_stats() + return f'{res["used"]}/{res["total"]}' + else: + res = memory_stats() + return res def start_server(immediate=True, server=None): @@ -260,6 +264,8 @@ def main(): get_custom_args() uv, instance = start_server(immediate=True, server=None) + t_server = time.time() + t_monitor = time.time() while True: try: alive = uv.thread.is_alive() @@ -267,8 +273,13 @@ def main(): except Exception: alive = False requests = 0 - if round(time.time()) % 120 == 0: - installer.log.debug(f'Server: alive={alive} requests={requests} memory={get_memory_stats()} {instance.state.status()}') + t_current = time.time() + if t_current - t_server > 120: + installer.log.trace(f'Server: alive={alive} requests={requests} memory={get_memory_stats()} {instance.state.status()}') + t_server = t_current + if float(args.monitor) > 0 and t_current - t_monitor > float(args.monitor): + installer.log.trace(f'Monitor: {get_memory_stats(detailed=True)}') + t_monitor = t_current if not alive: if uv is not None and uv.wants_restart: installer.log.info('Server restarting...') diff --git a/modules/cmd_args.py b/modules/cmd_args.py index 5e5e21054..a8dec6748 100644 --- a/modules/cmd_args.py +++ b/modules/cmd_args.py @@ -37,6 +37,7 @@ def main_args(): group_diag.add_argument("--no-hashing", default=os.environ.get("SD_NOHASHING", False), action='store_true', help="Disable hashing of checkpoints, default: %(default)s") group_diag.add_argument("--no-metadata", default=os.environ.get("SD_NOMETADATA", False), action='store_true', help="Disable reading of metadata from models, default: %(default)s") group_diag.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") + group_diag.add_argument("--monitor", default=os.environ.get("SD_PROFILE", 0), help="Run memory monitor, default: %(default)s") group_http = parser.add_argument_group('HTTP') group_http.add_argument('--theme', type=str, default=os.environ.get("SD_THEME", None), help='Override UI theme') diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 7227680ce..fe0e15de5 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -3,7 +3,8 @@ import os import re import numpy as np from modules.lora import networks, lora_overrides, lora_load -from modules import extra_networks, shared, sd_models +from modules.lora import lora_common as l +from modules import extra_networks, shared debug = os.environ.get('SD_LORA_DEBUG', None) is not None @@ -39,7 +40,7 @@ def prompt(p): if shared.opts.lora_apply_tags == 0: return all_tags = [] - for loaded in networks.loaded_networks: + for loaded in l.loaded_networks: page = [en for en in shared.extra_networks if en.name == 'lora'][0] item = page.create_item(loaded.name) tags = (item or {}).get("tags", {}) @@ -69,12 +70,12 @@ def prompt(p): def infotext(p): - names = [i.name for i in networks.loaded_networks] + names = [i.name for i in l.loaded_networks] if len(names) > 0: p.extra_generation_params["LoRA networks"] = ", ".join(names) if shared.opts.lora_add_hashes_to_infotext: network_hashes = [] - for item in networks.loaded_networks: + for item in l.loaded_networks: if not item.network_on_disk.shorthash: continue network_hashes.append(item.network_on_disk.shorthash) @@ -113,6 +114,19 @@ def parse(p, params_list, step=0): return names, te_multipliers, unet_multipliers, dyn_dims +def unload_diffusers(): + if hasattr(shared.sd_model, "unfuse_lora"): + try: + shared.sd_model.unfuse_lora() + except Exception: + pass + if hasattr(shared.sd_model, "unload_lora_weights"): + try: + shared.sd_model.unload_lora_weights() # fails for non-CLIP models + except Exception: + pass + + class ExtraNetworkLora(extra_networks.ExtraNetwork): def __init__(self): @@ -131,11 +145,11 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): key = f'{",".join(include)}:{",".join(exclude)}' loaded = sd_model.loaded_loras.get(key, []) # shared.log.trace(f'Network load: type=LoRA key="{key}" requested={requested} loaded={loaded}') - if (len(requested) == 0) or (len(requested) != len(loaded)): + if len(requested) != len(loaded): sd_model.loaded_loras[key] = requested return True - for r, l in zip(requested, loaded): - if r != l: + for req, load in zip(requested, loaded): + if req != load: sd_model.loaded_loras[key] = requested return True return False @@ -160,40 +174,35 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if force_diffusers: has_changed = False # diffusers handle their own loading if len(exclude) == 0: + shared.state.begin('LoRA') lora_load.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load only on first call + shared.state.end() else: lora_load.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load has_changed = self.changed(requested, include, exclude) if has_changed: - networks.network_deactivate(include, exclude) + shared.state.begin('LoRA') + if len(l.previously_loaded_networks) > 0: + shared.log.info(f'Network unload: type=LoRA apply={[n.name for n in l.previously_loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"}') + networks.network_deactivate(include, exclude) networks.network_activate(include, exclude) - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) # TODO lora: required for flux to reapply offload after lora has been applied, but fails with oom - debug_log(f'Network load: type=LoRA previous={[n.name for n in networks.previously_loaded_networks]} current={[n.name for n in networks.loaded_networks]} changed') + if len(exclude) > 0: # only update on last activation + l.previously_loaded_networks = l.loaded_networks.copy() + shared.state.end() + debug_log(f'Network load: type=LoRA previous={[n.name for n in l.previously_loaded_networks]} current={[n.name for n in l.loaded_networks]} changed') - if len(networks.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or force_diffusers) and step == 0: + if len(l.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or force_diffusers) and step == 0: infotext(p) prompt(p) if (has_changed or force_diffusers) and len(include) == 0: # print only once - shared.log.info(f'Network load: type=LoRA apply={[n.name for n in networks.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') + shared.log.info(f'Network load: type=LoRA apply={[n.name for n in l.loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={l.timer.summary}') def deactivate(self, p): - if shared.native: - networks.previously_loaded_networks = networks.loaded_networks.copy() - debug_log(f'Network load: type=LoRA active={[n.name for n in networks.previously_loaded_networks]} deactivate') if shared.native and len(lora_load.diffuser_loaded) > 0: if not (shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled is True): - if hasattr(shared.sd_model, "unfuse_lora"): - try: - shared.sd_model.unfuse_lora() - except Exception: - pass - if hasattr(shared.sd_model, "unload_lora_weights"): - try: - shared.sd_model.unload_lora_weights() # fails for non-CLIP models - except Exception: - pass - if self.active and networks.debug: - shared.log.debug(f"Network end: type=LoRA time={networks.timer.summary}") + unload_diffusers() + if self.active and l.debug: + shared.log.debug(f"Network end: type=LoRA time={l.timer.summary}") if self.errors: for k, v in self.errors.items(): shared.log.error(f'LoRA: name="{k}" errors={v}') diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 7e865a167..e3dae287a 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -3,7 +3,7 @@ import re import time import torch import diffusers.models.lora -from modules.lora.lora_common import timer, debug, loaded_networks, previously_loaded_networks, extra_network_lora +from modules.lora import lora_common as l from modules import shared, devices, errors, model_quant @@ -14,7 +14,7 @@ re_network_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], network_layer_name: str, wanted_names: tuple): global bnb # pylint: disable=W0603 backup_size = 0 - if len(loaded_networks) > 0 and network_layer_name is not None and any([net.modules.get(network_layer_name, None) for net in loaded_networks]): # noqa: C419 # pylint: disable=R1729 + if len(l.loaded_networks) > 0 and network_layer_name is not None and any([net.modules.get(network_layer_name, None) for net in l.loaded_networks]): # noqa: C419 # pylint: disable=R1729 t0 = time.time() weights_backup = getattr(self, "network_weights_backup", None) @@ -33,25 +33,15 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n if bnb is None: bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: - with devices.inference_context(): - if shared.opts.lora_fuse_diffusers: - self.network_weights_backup = True - else: - self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,) - self.quant_state = weight.quant_state - self.quant_type = weight.quant_type - self.blocksize = weight.blocksize - else: if shared.opts.lora_fuse_diffusers: self.network_weights_backup = True else: - weights_backup = weight.clone() - self.network_weights_backup = weights_backup.to(devices.cpu) - else: - if shared.opts.lora_fuse_diffusers: - self.network_weights_backup = True + self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,) + self.quant_state, self.quant_type, self.blocksize = weight.quant_state, weight.quant_type, weight.blocksize else: - self.network_weights_backup = weight.clone().to(devices.cpu) + self.network_weights_backup = weight.clone().to(devices.cpu) if not shared.opts.lora_fuse_diffusers else True + else: + self.network_weights_backup = weight.clone().to(devices.cpu) if not shared.opts.lora_fuse_diffusers else True if bias_backup is None: if getattr(self, 'bias', None) is not None: @@ -65,7 +55,7 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n backup_size += self.network_weights_backup.numel() * self.network_weights_backup.element_size() if isinstance(self.network_weights_backup, torch.Tensor) else 0 if getattr(self, 'network_bias_backup', None) is not None: backup_size += self.network_bias_backup.numel() * self.network_bias_backup.element_size() if isinstance(self.network_bias_backup, torch.Tensor) else 0 - timer.backup += time.time() - t0 + l.timer.backup += time.time() - t0 return backup_size @@ -77,7 +67,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. pass batch_updown = None batch_ex_bias = None - loaded = loaded_networks if not use_previous else previously_loaded_networks + loaded = l.loaded_networks if not use_previous else l.previously_loaded_networks for net in loaded: module = net.modules.get(network_layer_name, None) if module is None: @@ -88,8 +78,8 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. weight = self.weight.to(devices.device) except Exception: weight = self.weight - updown, ex_bias = module.calc_updown(weight) + del module if updown is not None: if batch_updown is not None: batch_updown += updown.to(batch_updown.device) @@ -100,8 +90,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. batch_ex_bias += ex_bias.to(batch_ex_bias.device) else: batch_ex_bias = ex_bias.to(devices.device) - timer.calc += time.time() - t0 - + l.timer.calc += time.time() - t0 if shared.opts.diffusers_offload_mode == "sequential": t0 = time.time() if batch_updown is not None: @@ -109,10 +98,10 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if batch_ex_bias is not None: batch_ex_bias = batch_ex_bias.to(devices.cpu) t1 = time.time() - timer.move += t1 - t0 + l.timer.move += t1 - t0 except RuntimeError as e: - extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 - if debug: + l.extra_network_lora.errors[net.name] = l.extra_network_lora.errors.get(net.name, 0) + 1 + if l.debug: module_name = net.modules.get(network_layer_name, None) shared.log.error(f'LoRA apply weight name="{net.name}" module="{module_name}" layer="{network_layer_name}" {e}') errors.display(e, 'LoRA') @@ -121,7 +110,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. return batch_updown, batch_ex_bias -def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, device: torch.device = devices.device): +def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, bias: bool = False): if lora_weights is None: return None if deactivate: @@ -135,20 +124,25 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(devices.device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize, requires_grad=False) - # weight._quantize(devices.device) # TODO force imediate quantization + # TODO lora: maybe force imediate quantization + # weight._quantize(devices.device) / weight.to(device=device) except Exception as e: shared.log.error(f'Network load: type=LoRA quant=bnb cls={self.__class__.__name__} type={self.quant_type} blocksize={self.blocksize} state={vars(self.quant_state)} weight={self.weight} bias={lora_weights} {e}') else: try: new_weight = model_weights.to(devices.device) + lora_weights.to(devices.device) - except Exception: + except Exception as e: + shared.log.warning(f'Network load: {e}') new_weight = model_weights + lora_weights # try without device cast + del model_weights + del lora_weights weight = torch.nn.Parameter(new_weight, requires_grad=False) - try: - # weight.to(device=device) # TODO required since quantization happens only during .to call, not during params creation - pass - except Exception: - pass # may fail if weights is meta tensor + del new_weight # without this its a massive memory leak + if weight is not None: + if not bias: + self.weight = weight + else: + self.bias = weight return weight @@ -166,22 +160,18 @@ def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if weights_backup: if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 - updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable + updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - weight = network_add_weights(self, lora_weights=updown, deactivate=deactivate, device=device) - if weight is not None: - self.weight = weight + network_add_weights(self, lora_weights=updown, deactivate=deactivate, bias=False) if bias_backup: if ex_bias is not None: - bias = network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, device=device) - if bias is not None: - self.bias = bias + network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, bias=True) if hasattr(self, "qweight") and hasattr(self, "freeze"): self.freeze() - timer.apply += time.time() - t0 + l.timer.apply += time.time() - t0 def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, device: torch.device, deactivate: bool = False): @@ -194,24 +184,20 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn if weights_backup is not None: self.weight = None if updown is not None and len(weights_backup.shape) == 4 and weights_backup.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 - updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable + updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - weight = network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate, device=device) - if weight is not None: - self.weight = weight + network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate, bias=False) else: self.weight = torch.nn.Parameter(weights_backup.to(device), requires_grad=False) if bias_backup is not None: self.bias = None if ex_bias is not None: - bias = network_add_weights(self, model_weights=weights_backup, lora_weights=ex_bias, deactivate=deactivate, device=device) - if bias: - self.weight = bias + network_add_weights(self, model_weights=bias_backup, lora_weights=ex_bias, deactivate=deactivate, bias=True) else: self.bias = torch.nn.Parameter(bias_backup.to(device), requires_grad=False) if hasattr(self, "qweight") and hasattr(self, "freeze"): self.freeze() - timer.apply += time.time() - t0 + l.timer.apply += time.time() - t0 diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py index 43efc6f04..110cc0e46 100644 --- a/modules/lora/lora_load.py +++ b/modules/lora/lora_load.py @@ -4,7 +4,7 @@ import time import concurrent from modules import shared, errors, devices, sd_models, sd_models_compile, files_cache from modules.lora import network, lora_overrides, lora_convert -from modules.lora.lora_common import timer, debug, module_types, loaded_networks +from modules.lora import lora_common as l diffuser_loaded = [] @@ -35,7 +35,7 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ shared.log.error(f'Network load: type=LoRA name="{name}" diffusers unsupported format') else: shared.log.error(f'Network load: type=LoRA name="{name}" {e}') - if debug: + if l.debug: errors.display(e, "LoRA") return None if name not in diffuser_loaded: @@ -43,7 +43,7 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ diffuser_scales.append(lora_scale) net = network.Network(name, network_on_disk) net.mtime = os.path.getmtime(network_on_disk.filename) - timer.activate += time.time() - t0 + l.timer.activate += time.time() - t0 return net @@ -52,19 +52,19 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: return None cached = lora_cache.get(name, None) - if debug: + if l.debug: shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') if cached is not None: return cached net = network.Network(name, network_on_disk) net.mtime = os.path.getmtime(network_on_disk.filename) sd = sd_models.read_state_dict(network_on_disk.filename, what='network') - if shared.sd_model_type == 'f1': # if kohya flux lora, convert state_dict - sd = lora_convert._convert_kohya_flux_lora_to_diffusers(sd) or sd # pylint: disable=protected-access - if shared.sd_model_type == 'sd3': # if kohya flux lora, convert state_dict + if shared.sd_model_type == 'f1': # if kohya flux lora, convert state_dict + sd = lora_convert._convert_kohya_flux_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + if shared.sd_model_type == 'sd3': # if kohya flux lora, convert state_dict try: - sd = lora_convert._convert_kohya_sd3_lora_to_diffusers(sd) or sd # pylint: disable=protected-access - except ValueError: # EAFP for diffusers PEFT keys + sd = lora_convert._convert_kohya_sd3_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + except ValueError: # EAFP for diffusers PEFT keys pass lora_convert.assign_network_names_to_compvis_modules(shared.sd_model) keys_failed_to_match = {} @@ -72,6 +72,7 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: bundle_embeddings = {} dtypes = [] convert = lora_convert.KeyConvert() + device = devices.device if shared.opts.lora_apply_gpu else devices.cpu for key_network, weight in sd.items(): parts = key_network.split('.') if parts[0] == "bundle_emb": @@ -99,7 +100,7 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: network_types = [] for key, weights in matched_networks.items(): net_module = None - for nettype in module_types: + for nettype in l.module_types: net_module = nettype.create_module(net, weights) if net_module is not None: network_types.append(nettype.__class__.__name__) @@ -110,10 +111,10 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: net.modules[key] = net_module if len(keys_failed_to_match) > 0: shared.log.warning(f'Network load: type=LoRA name="{name}" type={set(network_types)} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}') - if debug: + if l.debug: shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') else: - shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} dtypes={dtypes} direct={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} device={device} dtypes={dtypes} direct={shared.opts.lora_fuse_diffusers}') if len(matched_networks) == 0: return None lora_cache[name] = net @@ -134,7 +135,7 @@ def maybe_recompile_model(names, te_multipliers): break if not recompile_model: skip_lora_load = True - if len(loaded_networks) > 0 and debug: + if len(l.loaded_networks) > 0 and l.debug: shared.log.debug('Model Compile: Skipping LoRa loading') return recompile_model, skip_lora_load else: @@ -178,7 +179,7 @@ def list_available_networks(): available_network_aliases[entry.alias] = entry if entry.shorthash: available_network_hash_lookup[entry.shorthash] = entry - except OSError as e: # should catch FileNotFoundError and PermissionError etc. + except OSError as e: # should catch FileNotFoundError and PermissionError etc. shared.log.error(f'LoRA: filename="{filename}" {e}') candidates = sorted(files_cache.list_files(shared.cmd_opts.lora_dir, ext_filter=[".pt", ".ckpt", ".safetensors"])) @@ -186,7 +187,7 @@ def list_available_networks(): for fn in candidates: executor.submit(add_network, fn) t1 = time.time() - timer.list = t1 - t0 + l.timer.list = t1 - t0 shared.log.info(f'Available LoRAs: path="{shared.cmd_opts.lora_dir}" items={len(available_networks)} folders={len(forbidden_network_aliases)} time={t1 - t0:.2f}') @@ -214,7 +215,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non failed_to_load_networks = [] recompile_model, skip_lora_load = maybe_recompile_model(names, te_multipliers) - loaded_networks.clear() + l.loaded_networks.clear() diffuser_loaded.clear() diffuser_scales.clear() t0 = time.time() @@ -223,7 +224,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non net = None if network_on_disk is not None: shorthash = getattr(network_on_disk, 'shorthash', '').lower() - if debug: + if l.debug: shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') try: if recompile_model: @@ -237,7 +238,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non network_on_disk.read_hash() except Exception as e: shared.log.error(f'Network load: type=LoRA file="{network_on_disk.filename}" {e}') - if debug: + if l.debug: errors.display(e, 'LoRA') continue if net is None: @@ -249,7 +250,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non net.te_multiplier = te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier net.unet_multiplier = unet_multipliers[i] if unet_multipliers else shared.opts.extra_networks_default_multiplier net.dyn_dim = dyn_dims[i] if dyn_dims else shared.opts.extra_networks_default_multiplier - loaded_networks.append(net) + l.loaded_networks.append(net) while len(lora_cache) > shared.opts.lora_in_memory_limit: name = next(iter(lora_cache)) @@ -261,16 +262,16 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non t0 = time.time() shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) if shared.opts.lora_fuse_diffusers and not lora_overrides.check_fuse(): - shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling + shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # diffusers with fuse uses fixed scale since later apply does the scaling shared.sd_model.unload_lora_weights() - timer.activate += time.time() - t0 + l.timer.activate += time.time() - t0 except Exception as e: shared.log.error(f'Network load: type=LoRA {e}') - if debug: + if l.debug: errors.display(e, 'LoRA') - if len(loaded_networks) > 0 and debug: - shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in loaded_networks]} cache={list(lora_cache)}') + if len(l.loaded_networks) > 0 and l.debug: + shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in l.loaded_networks]} cache={list(lora_cache)}') if recompile_model: shared.log.info("Network load: type=LoRA recompiling model") @@ -279,7 +280,4 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) shared.compiled_model_state.lora_model = backup_lora_model - if len(loaded_networks) > 0: - devices.torch_gc() - - timer.load = time.time() - t0 + l.timer.load = time.time() - t0 diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 44dad1afd..a36ba5631 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -1,7 +1,7 @@ from contextlib import nullcontext import time import rich.progress as rp -from modules.lora.lora_common import timer, debug, loaded_networks, previously_loaded_networks +from modules.lora import lora_common as l from modules.lora.lora_apply import network_apply_weights, network_apply_direct, network_backup_weights, network_calc_weights from modules import shared, devices, sd_models @@ -11,7 +11,7 @@ applied_layers: list[str] = [] def network_activate(include=[], exclude=[]): t0 = time.time() - sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility + sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) if shared.opts.diffusers_offload_mode == "sequential": sd_models.disable_offload(sd_model) sd_models.move_model(sd_model, device=devices.cpu) @@ -25,7 +25,7 @@ def network_activate(include=[], exclude=[]): active_components.append(name) modules[name] = list(component.named_modules()) total = sum(len(x) for x in modules.values()) - if len(loaded_networks) > 0: + if len(l.loaded_networks) > 0: pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=activate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) task = pbar.add_task(description='' , total=total) else: @@ -33,9 +33,9 @@ def network_activate(include=[], exclude=[]): pbar = nullcontext() applied_weight = 0 applied_bias = 0 - device = devices.device if shared.opts.lora_apply_gpu else devices.cpu + device = devices.device if shared.opts.lora_apply_gpu or shared.opts.diffusers_offload_mode == 'none' else devices.cpu with devices.inference_context(), pbar: - wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) if len(loaded_networks) > 0 else () + wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in l.loaded_networks) if len(l.loaded_networks) > 0 else () applied_layers.clear() backup_size = 0 for component in modules.keys(): @@ -55,32 +55,32 @@ def network_activate(include=[], exclude=[]): network_apply_weights(module, batch_updown, batch_ex_bias, device=orig_device) if batch_updown is not None or batch_ex_bias is not None: applied_layers.append(network_layer_name) - # module.to(device) # TODO maybe - if batch_updown is not None: - applied_weight += 1 - if batch_ex_bias is not None: - applied_bias += 1 + applied_weight += 1 if batch_updown is not None else 0 + applied_bias += 1 if batch_ex_bias is not None else 0 + batch_updown, batch_ex_bias = None, None del batch_updown, batch_ex_bias module.network_current_names = wanted_names if task is not None: - pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size}') + bs = round(backup_size/1024/1024/1024, 2) if backup_size > 0 else None + pbar.update(task, advance=1, description=f'networks={len(l.loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={bs} device={device}') if task is not None and len(applied_layers) == 0: pbar.remove_task(task) # hide progress bar for no action - timer.activate += time.time() - t0 - if debug and len(loaded_networks) > 0: - shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} device={device} time={timer.summary}') + l.timer.activate += time.time() - t0 + if l.debug and len(l.loaded_networks) > 0: + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={round(backup_size/1024/1024/1024, 2)} fuse={shared.opts.lora_fuse_diffusers} device={device} time={l.timer.summary}') modules.clear() - if len(loaded_networks) > 0 and (applied_weight > 0 or applied_bias > 0): - if shared.opts.diffusers_offload_mode == "sequential": - sd_models.set_diffuser_offload(sd_model, op="model") + if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential": + sd_models.set_diffuser_offload(sd_model, op="model") def network_deactivate(include=[], exclude=[]): if not shared.opts.lora_fuse_diffusers or shared.opts.lora_force_diffusers: return + if len(l.previously_loaded_networks) == 0: + return t0 = time.time() - sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility + sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) if shared.opts.diffusers_offload_mode == "sequential": sd_models.disable_offload(sd_model) sd_models.move_model(sd_model, device=devices.cpu) @@ -96,7 +96,7 @@ def network_deactivate(include=[], exclude=[]): active_components.append(name) total = sum(len(x) for x in modules.values()) device = devices.device if shared.opts.lora_apply_gpu else devices.cpu - if len(previously_loaded_networks) > 0 and debug: + if len(l.previously_loaded_networks) > 0 and l.debug: pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=deactivate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) task = pbar.add_task(description='', total=total) else: @@ -118,16 +118,15 @@ def network_deactivate(include=[], exclude=[]): else: network_apply_weights(module, batch_updown, batch_ex_bias, device=orig_device, deactivate=True) if batch_updown is not None or batch_ex_bias is not None: - # module.to(device) # TODO maybe applied_layers.append(network_layer_name) del batch_updown, batch_ex_bias module.network_current_names = () if task is not None: - pbar.update(task, advance=1, description=f'networks={len(previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') + pbar.update(task, advance=1, description=f'networks={len(l.previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') - timer.deactivate = time.time() - t0 - if debug and len(previously_loaded_networks) > 0: - shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + l.timer.deactivate = time.time() - t0 + if l.debug and len(l.previously_loaded_networks) > 0: + shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in l.previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={l.timer.summary}') modules.clear() - if shared.opts.diffusers_offload_mode == "sequential": + if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/memstats.py b/modules/memstats.py index 160492069..4fd6206f1 100644 --- a/modules/memstats.py +++ b/modules/memstats.py @@ -56,15 +56,17 @@ def memory_stats(): fail_once = True mem.update({ 'ram': { 'error': str(e) } }) try: - s = torch.cuda.mem_get_info() - gpu = { 'used': gb(s[1] - s[0]), 'total': gb(s[1]) } - s = dict(torch.cuda.memory_stats()) - if s.get('num_ooms', 0) > 0: + free, total = torch.cuda.mem_get_info() + gpu = { 'used': gb(total - free), 'total': gb(total) } + stats = dict(torch.cuda.memory_stats()) + if stats.get('num_ooms', 0) > 0: shared.state.oom = True mem.update({ 'gpu': gpu, - 'retries': s.get('num_alloc_retries', 0), - 'oom': s.get('num_ooms', 0) + 'active': gb(stats.get('active_bytes.all.current', 0)), + 'peak': gb(stats.get('active_bytes.all.peak', 0)), + 'retries': stats.get('num_alloc_retries', 0), + 'oom': stats.get('num_ooms', 0), }) return mem except Exception: diff --git a/modules/model_quant.py b/modules/model_quant.py index 4078dd01a..40a2ce3f8 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -3,7 +3,7 @@ import sys import copy import time import diffusers -from installer import install, log, setup_logging +from installer import installed, install, log, setup_logging ao = None @@ -116,7 +116,9 @@ def load_torchao(msg='', silent=False): global ao # pylint: disable=global-statement if ao is not None: return ao - install('torchao==0.8.0', quiet=True) + if not installed('torchao'): + install('torchao==0.8.0', quiet=True) + log.warning('Quantization: torchao installed please restart') try: import torchao ao = torchao @@ -140,9 +142,11 @@ def load_bnb(msg='', silent=False): global bnb # pylint: disable=global-statement if bnb is not None: return bnb - if devices.backend == 'cuda': - # forcing a version will uninstall the multi-backend-refactor branch of bnb - install('bitsandbytes==0.45.1', quiet=True) + if not installed('bitsandbytes'): + if devices.backend == 'cuda': + # forcing a version will uninstall the multi-backend-refactor branch of bnb + install('bitsandbytes==0.45.1', quiet=True) + log.warning('Quantization: bitsandbytes installed please restart') try: import bitsandbytes bnb = bitsandbytes @@ -165,7 +169,9 @@ def load_quanto(msg='', silent=False): global optimum_quanto # pylint: disable=global-statement if optimum_quanto is not None: return optimum_quanto - install('optimum-quanto==0.2.7', quiet=True) + if not installed('optimum-quanto'): + install('optimum-quanto==0.2.7', quiet=True) + log.warning('Quantization: optimum-quanto installed please restart') try: from optimum import quanto # pylint: disable=no-name-in-module optimum_quanto = quanto @@ -190,7 +196,9 @@ def load_nncf(msg='', silent=False): global intel_nncf # pylint: disable=global-statement if intel_nncf is not None: return intel_nncf - install('nncf==2.7.0', quiet=True) + if not installed('nncf'): + install('nncf==2.7.0', quiet=True) + log.warning('Quantization: nncf installed please restart') try: import nncf intel_nncf = nncf diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index de21cf8b9..e61ca1afc 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -9,7 +9,7 @@ from modules import shared, devices, processing, sd_models, errors, sd_hijack_hy from modules.processing_helpers import resize_hires, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, save_intermediate, update_sampler, is_txt2img, is_refiner_enabled, get_job_name from modules.processing_args import set_pipeline_args from modules.onnx_impl import preprocess_pipeline as preprocess_onnx_pipeline, check_parameters_changed as olive_check_parameters_changed -from modules.lora import networks +from modules.lora import lora_common debug = shared.log.trace if os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -478,8 +478,8 @@ def process_diffusers(p: processing.StableDiffusionProcessing): return results extra_networks.deactivate(p) - timer.process.add('lora', networks.timer.total) - networks.timer.clear(complete=True) + timer.process.add('lora', lora_common.timer.total) + lora_common.timer.clear(complete=True) results = process_decode(p, output) timer.process.record('decode') diff --git a/modules/sd_models.py b/modules/sd_models.py index 2d4451db1..1cdec4afb 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1051,7 +1051,6 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model', def clear_caches(): - # shared.log.debug('Cache clear') if not shared.opts.lora_legacy: from modules.lora import lora_common, lora_load lora_common.loaded_networks.clear() diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 2e02f98a9..e4b24c17b 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -166,9 +166,7 @@ class OffloadHook(accelerate.hooks.ModelHook): keys = device_map.keys() for v in keys: if isinstance(device_map[v], int): - # int implies CUDA or XPU device, but it will break DirectML backend. - # Therefore, the type of device should be added. - device_map[v] = f"{devices.device.type}:{device_map[v]}" + device_map[v] = f"{devices.device.type}:{device_map[v]}" # int implies CUDA or XPU device, but it will break DirectML backend so we add type module = accelerate.dispatch_model(module, device_map=device_map, offload_dir=offload_dir) module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access module.balanced_offload_device_map = device_map From bd304b59b7f9f972145382f1a2dbd1be0267eca3 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 13:39:45 -0400 Subject: [PATCH 108/122] linting fixes Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 14 ++++++++++---- scripts/xyz_grid_shared.py | 2 +- webui.py | 1 + wiki | 2 +- 4 files changed, 13 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 0ca2ae7ec..a9b4f44b6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2025-03-31 +## Update for 2025-04-01 -### Highlights for 2025-03-31 +### Highlights for 2025-04-01 Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! @@ -15,7 +15,7 @@ Plus... - More quantization options and granular control - Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods -### Details for 2025-03-31 +### Details for 2025-04-01 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -107,6 +107,13 @@ Plus... - add image policy checks using `LlavaGuard` VLM to detect policy violations (and reasons) against top-10 standard harmful content categories - add banned words/expressions check against prompt variations +- **LoRA** + - enable memory cache by default + - significantly reduce memory usage + - improve performance + - improve detection of lora changes + - unload lora only when changes are detected + - refactor code for modularity - **IPEX** - add `--upgrade` to torch_command when using `--use-nightly` - add xpu to profiler @@ -131,7 +138,6 @@ Plus... - update `diffusers` and other requirements - rename vae, unet and text-encoder settings *None* to *Default* to avoid confusion - **CLI**: add `cli/api-grid.py` which can generate grids using params-from-file for x/y axis - - **LoRA** enable memory cache by default - **Samplers** add ability to set sigma adjustment for each sampler - **ModernUI** updates - **CSS** updates diff --git a/scripts/xyz_grid_shared.py b/scripts/xyz_grid_shared.py index a3f7a3da9..f624458d6 100644 --- a/scripts/xyz_grid_shared.py +++ b/scripts/xyz_grid_shared.py @@ -211,7 +211,7 @@ def apply_vae(p, x, xs): def list_lora(): import sys - lora = [v for k, v in sys.modules.items() if k == 'networks' or k == 'modules.lora.networks'][0] + lora = [v for k, v in sys.modules.items() if k == 'networks' or k == 'modules.lora.lora_load'][0] loras = [v.fullname for v in lora.available_networks.values()] return ['None'] + sorted(loras) diff --git a/webui.py b/webui.py index e44c5c2b9..c59aa5647 100644 --- a/webui.py +++ b/webui.py @@ -37,6 +37,7 @@ import modules.hypernetworks.hypernetwork import modules.script_callbacks import modules.api.middleware + if not modules.loader.initialized: timer.startup.record("libraries") import modules.sd_hijack # runs conditional load of ldm if not shared.native diff --git a/wiki b/wiki index bb4e16f8d..9408b299f 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit bb4e16f8d88d353f83f00f854852d1dd649fcd28 +Subproject commit 9408b299fffbb8efaccae968b2ac64d9216326fb From 17fbadb0e811c7ab946ae87cf8ceadde41ca617e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 14:00:39 -0400 Subject: [PATCH 109/122] reset memory stats between runs Signed-off-by: Vladimir Mandic --- modules/memstats.py | 7 +++++++ modules/processing.py | 1 + modules/sd_models.py | 3 ++- 3 files changed, 10 insertions(+), 1 deletion(-) diff --git a/modules/memstats.py b/modules/memstats.py index 4fd6206f1..5880a6d97 100644 --- a/modules/memstats.py +++ b/modules/memstats.py @@ -74,6 +74,13 @@ def memory_stats(): return mem +def reset_stats(): + try: + torch.cuda.reset_memory_stats() + except Exception: + pass + + def memory_cache(): return mem diff --git a/modules/processing.py b/modules/processing.py index 67d86021b..b1a8bd548 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -81,6 +81,7 @@ class Processed: self.all_seeds = all_seeds or p.all_seeds or [self.seed] self.all_subseeds = all_subseeds or p.all_subseeds or [self.subseed] self.infotexts = infotexts or [self.info] + memstats.reset_stats() def js(self): obj = { diff --git a/modules/sd_models.py b/modules/sd_models.py index 1cdec4afb..4e85ec74c 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1056,8 +1056,9 @@ def clear_caches(): lora_common.loaded_networks.clear() lora_common.previously_loaded_networks.clear() lora_load.lora_cache.clear() - from modules import prompt_parser_diffusers + from modules import prompt_parser_diffusers, memstats prompt_parser_diffusers.cache.clear() + memstats.reset_stats() def unload_model_weights(op='model'): From d3633be48496521013fe7bfa3451e84d943866ee Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 14:19:40 -0400 Subject: [PATCH 110/122] add job info to monitor stats Signed-off-by: Vladimir Mandic --- modules/memstats.py | 1 + 1 file changed, 1 insertion(+) diff --git a/modules/memstats.py b/modules/memstats.py index 5880a6d97..90512e870 100644 --- a/modules/memstats.py +++ b/modules/memstats.py @@ -62,6 +62,7 @@ def memory_stats(): if stats.get('num_ooms', 0) > 0: shared.state.oom = True mem.update({ + 'job': shared.state.job, 'gpu': gpu, 'active': gb(stats.get('active_bytes.all.current', 0)), 'peak': gb(stats.get('active_bytes.all.peak', 0)), From e5a256e74d119f3a67c68df31a84f6ffbf629cf6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 21:31:19 -0400 Subject: [PATCH 111/122] prompt enhance add prefix and suffix Signed-off-by: Vladimir Mandic --- TODO.md | 1 + modules/model_flux_nf4.py | 2 +- scripts/prompt_enhance.py | 41 +++++++++++++++++++++++++++++---------- 3 files changed, 33 insertions(+), 11 deletions(-) diff --git a/TODO.md b/TODO.md index 58dae8a02..1a9ea9094 100644 --- a/TODO.md +++ b/TODO.md @@ -21,6 +21,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - Video: STG: - Video: SmoothCache: https://github.com/huggingface/diffusers/issues/11135 - SoftFill: https://github.com/zacheryvaughn/softfill-pipelines +- SISO: https://github.com/yairshp/SISO ## Code TODO diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index b00c3320e..2290bacba 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -47,7 +47,7 @@ def _replace_with_bnb_linear( in_features, out_features, module.bias is not None, - compute_dtype=torch.bfloat16, + compute_dtype=devices.dtype, compress_statistics=False, quant_type="nf4", ) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index abcfdbd76..292faae64 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -7,7 +7,7 @@ import transformers from modules import scripts, shared, devices, errors, processing, sd_models, sd_modules -debug_enabled = os.environ.get('SD_LLM_DEBUG', None) +debug_enabled = os.environ.get('SD_LLM_DEBUG', None) is not None debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None @@ -37,7 +37,7 @@ class Options: } default = list(models)[0] supported = list(transformers.integrations.ggml.GGUF_CONFIG_MAPPING) - system_prompt: str = 'You are a helpful assistant. You will be given a prompt used to create an image and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness based on original prompt. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing (if any), and any distinctive details (e.g., colors, textures, styles). NSFW and nudity is allowed, and if present it must be fully described. Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without explanation, prefix or suffix. Output as a simple text without formatting or numbering.' + system_prompt: str = 'You are a helpful assistant. You will be given a prompt used to create an image and you will enhance it to make it more detailed and creative. Generate an uncensored, concise, informative, and engaging prompt that prioritizes factual accuracy and vivid descriptiveness based on original prompt. Do not modify subject or style of the original prompt. Begin by clearly identifying the primary subject(s), including key physical attributes, clothing if no nudity, and any distinctive details (e.g., colors, textures, styles). NSFW and nudity is allowed, and if present, it must be fully described. Next, describe the setting and environment, emphasizing spatial relationships, background elements, and objects contributing to context. Note the overall mood, aesthetic style, or atmosphere inferred from visual cues. Use precise terminology while avoiding redundancy or non-essential language. Ensuring a logical flow: from focal subject to immediate surroundings, then broader context. Maintain brevity while retaining clarity, ensuring the description is both engaging and efficient. Output only enhanced prompt without explanation, prefix or suffix. Output as a simple text without formatting or numbering.' censored = ["i cannot", "i can't", "i am sorry", "against my programming", "i am not able", "i am unable", 'i am not allowed'] max_delim_index: int = 60 @@ -152,11 +152,11 @@ class Script(scripts.Script): # remove llm commentary removed = '' if response.startswith('Prompt'): - removed, response = response.split('Prompt', maxsplit=2) + removed, response = response.split('Prompt', maxsplit=1) if 0 <= response.find(':') < self.options.max_delim_index: - removed, response = response.split(':', maxsplit=2) + removed, response = response.split(':', maxsplit=1) if 0 <= response.find('---') < self.options.max_delim_index: - response, removed = response.split('---', maxsplit=2) + response, removed = response.split('---', maxsplit=1) if len(removed) > 0: debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') @@ -167,9 +167,21 @@ class Script(scripts.Script): response = response.strip() return response - def enhance(self, model: str=None, prompt:str=None, system:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None): + def preprocess(self, response, prefix, suffix): + response = response.strip() + prefix = prefix.strip() + suffix = suffix.strip() + if len(prefix) > 0: + response = f'{prefix} {response}' + if len(suffix) > 0: + response = f'{response} {suffix}' + return response + + def enhance(self, model: str=None, prompt:str=None, system:str=None, prefix:str=None, suffix:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None): model = model or self.options.default prompt = prompt or self.prompt.value + prefix = prefix or '' + suffix = suffix or '' system = system or self.options.system_prompt tokens = tokens or self.options.max_tokens penalty = penalty or self.options.repetition_penalty @@ -235,6 +247,7 @@ class Script(scripts.Script): is_censored = self.censored(response) if not is_censored: response = self.clean(response) + response = self.preprocess(response, prefix, suffix) shared.log.info(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') if debug: shared.log.trace(f'Prompt enhance: sample={sample} tokens={tokens} temperature={temperature} penalty={penalty}') @@ -246,9 +259,11 @@ class Script(scripts.Script): return prompt return response - def apply(self, prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty): + def apply(self, prompt, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty): response = self.enhance( prompt=prompt, + prefix=prompt_prefix, + suffix=prompt_suffix, model=llm_model, system=prompt_system, sample=do_sample, @@ -306,6 +321,10 @@ class Script(scripts.Script): repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) gr.HTML('
') with gr.Accordion('Input', open=False, elem_id='prompt_enhance_system_prompt'): + with gr.Row(): + prompt_prefix = gr.Textbox(label='Prompt prefix', value='', placeholder='Optional prompt prefix', interactive=True, lines=2, elem_id='prompt_enhance_prefix') + with gr.Row(): + prompt_suffix = gr.Textbox(label='Prompt suffix', value='', placeholder='Optional prompt suffix', interactive=True, lines=2, elem_id='prompt_enhance_suffix') with gr.Row(): prompt_system = gr.Textbox(label='System prompt', value=self.options.system_prompt, interactive=True, lines=4, elem_id='prompt_enhance_system') with gr.Accordion('Output', open=True, elem_id='prompt_enhance_system_prompt'): @@ -316,15 +335,15 @@ class Script(scripts.Script): clear_btn.click(fn=lambda: '', inputs=[], outputs=[prompt_output]) copy_btn = gr.Button(value='Set prompt', elem_id='prompt_enhance_copy', variant='secondary') copy_btn.click(fn=lambda x: x, inputs=[prompt_output], outputs=[self.prompt]) - apply_btn.click(fn=self.apply, inputs=[self.prompt, apply_prompt, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty], outputs=[prompt_output, self.prompt]) - return [apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty] + apply_btn.click(fn=self.apply, inputs=[self.prompt, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty], outputs=[prompt_output, self.prompt]) + return [apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty] def after_component(self, component, **kwargs): # searching for actual ui prompt components if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt', 'video_prompt']: self.prompt = component def before_process(self, p: processing.StableDiffusionProcessing, *args, **kwargs): # pylint: disable=unused-argument - apply_auto, llm_model, prompt_system, max_tokens, do_sample, temperature, repetition_penalty = args + apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty = args if not apply_auto and not p.enhance_prompt: return if shared.state.skipped or shared.state.interrupted: @@ -336,6 +355,8 @@ class Script(scripts.Script): shared.state.begin('LLM') p.prompt = self.enhance( prompt=p.prompt, + prefix=prompt_prefix, + suffix=prompt_suffix, model=llm_model, system=prompt_system, sample=do_sample, From 79d665e5dabe152a407c9e6c9c31dcaff98a01da Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 22:16:32 -0400 Subject: [PATCH 112/122] prompt enhance lora extract/recombine Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 27 ++++++++++++++++++--------- 1 file changed, 18 insertions(+), 9 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 292faae64..bbc52b0be 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -8,7 +8,7 @@ from modules import scripts, shared, devices, errors, processing, sd_models, sd_ debug_enabled = os.environ.get('SD_LLM_DEBUG', None) is not None -debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None +debug_log = shared.log.trace if debug_enabled else lambda *args, **kwargs: None @dataclass @@ -79,8 +79,7 @@ class Script(scripts.Script): gguf_args = {} if model_type is not None and model_file is not None and len(model_type) > 2 and len(model_file) > 2: - if debug: - shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') + debug_log(f'Prompt enhance: gguf supported={self.options.supported}') if model_type not in self.options.supported: shared.log.error(f'Prompt enhance: name="{name}" repo="{model_repo}" fn="{model_file}" type={model_type} gguf not supported') shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') @@ -112,7 +111,7 @@ class Script(scripts.Script): pretrained_model_name_or_path=model_repo, cache_dir=shared.opts.hfcache_dir, ) - if debug: + if debug_enabled: modules = sd_modules.get_model_stats(self.llm) + sd_modules.get_model_stats(self.tokenizer) for m in modules: shared.log.trace(f'Prompt enhance: {m}') @@ -158,7 +157,7 @@ class Script(scripts.Script): if 0 <= response.find('---') < self.options.max_delim_index: response, removed = response.split('---', maxsplit=1) if len(removed) > 0: - debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') + debug_log(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') # remove bullets and lists lines = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line).strip() for line in response.splitlines()] @@ -167,7 +166,7 @@ class Script(scripts.Script): response = response.strip() return response - def preprocess(self, response, prefix, suffix): + def post(self, response, prefix, suffix, networks): response = response.strip() prefix = prefix.strip() suffix = suffix.strip() @@ -175,8 +174,16 @@ class Script(scripts.Script): response = f'{prefix} {response}' if len(suffix) > 0: response = f'{response} {suffix}' + if len(networks) > 0: + response = f'{response} {" ".join(networks)}' return response + def extract(self, prompt): + pattern = r'(<.*?>)' + matches = re.findall(pattern, prompt) + filtered = re.sub(pattern, '', prompt) + return filtered, matches + def enhance(self, model: str=None, prompt:str=None, system:str=None, prefix:str=None, suffix:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None): model = model or self.options.default prompt = prompt or self.prompt.value @@ -193,6 +200,8 @@ class Script(scripts.Script): if self.llm is None: shared.log.error('Prompt enhance: model not loaded') return prompt + prompt, networks = self.extract(prompt) + debug_log(f'Prompt enhance: networks={networks}') chat_template = [ { "role": "system", "content": system }, { "role": "user", "content": prompt }, @@ -226,7 +235,7 @@ class Script(scripts.Script): if shared.opts.diffusers_offload_mode != 'none': sd_models.move_model(self.llm, devices.cpu) devices.torch_gc() - if debug: + if debug_enabled: raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) shared.log.trace(f'Prompt enhance: raw="{raw_response}"') outputs_cropped = outputs[:, input_len:] @@ -247,9 +256,9 @@ class Script(scripts.Script): is_censored = self.censored(response) if not is_censored: response = self.clean(response) - response = self.preprocess(response, prefix, suffix) + response = self.post(response, prefix, suffix, networks) shared.log.info(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') - if debug: + if debug_enabled: shared.log.trace(f'Prompt enhance: sample={sample} tokens={tokens} temperature={temperature} penalty={penalty}') shared.log.trace(f'Prompt enhance: prompt="{prompt}"') shared.log.trace(f'Prompt enhance: response="{response}"') From 49ba25825ac6d82197c230a42a0b4e1e4790342b Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Wed, 2 Apr 2025 21:28:20 +0900 Subject: [PATCH 113/122] change log --- CHANGELOG.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a9b4f44b6..2dafb0fe8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -125,9 +125,9 @@ Plus... - add `--upgrade` to torch_command when using `--use-nightly` - disable fp16 for gfx1102 (rx 7600 and rx 7500 series) gpus - **ZLUDA** - - add `torch.compile` support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - - add Flash Attention 2 support under [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) - - add Sage Attention support + - [triton for ZLUDA v3.9.2](https://github.com/vladmandic/sdnext/wiki/ZLUDA#how-to-enable-triton) + - `torch.compile` is now available + - Flash Attention 2 is now available - **Other** - new command line option `--monitor PERIOD` to monitor CPU and GPU memory ever n seconds - **upscale**: new [asymmetric vae v2](https://huggingface.co/Heasterian/AsymmetricAutoencoderKLUpscaler_v2) upscaling method From 73cd4727e658b12f1871c5f7a1872bdd701608df Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 2 Apr 2025 08:59:25 -0400 Subject: [PATCH 114/122] fix paste incorrect float to int cast Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 13 +++++++------ extensions-builtin/sdnext-modernui | 2 +- modules/generation_parameters_copypaste.py | 3 +++ modules/ui_control.py | 1 + modules/ui_img2img.py | 1 + modules/ui_txt2img.py | 1 + 6 files changed, 14 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 2dafb0fe8..dfec2e7ef 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,11 +1,11 @@ # Change Log for SD.Next -## Update for 2025-04-01 +## Update for 2025-04-02 -### Highlights for 2025-04-01 +### Highlights for 2025-04-02 -Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both T2V and I2V workflows -And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB* and more! +Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both *T2V* and *I2V* workflows +And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB*, and more! Also, support for new models: **CogView-4**, **SANA 1.5**, @@ -13,9 +13,9 @@ Plus... - New **Prompt Enhance** using LLM, - New **CLiP** models, improvements to **remote VAE**, additional wiki/docs/guides - More quantization options and granular control -- Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods +- Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods, c) much lower LoRA memory usage -### Details for 2025-04-01 +### Details for 2025-04-02 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -172,6 +172,7 @@ Plus... - fix extra networks cover and inline views - fix token counter error style with modernui - fix sampler metadata when using default sampler + - fix paste incorrect float to int cast - improve lora compatibility with balanced offload ## Update for 2025-02-28 diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index ed8291c74..770db0076 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit ed8291c74faea33c51b25df379f48704e36bef37 +Subproject commit 770db007688d5be9df0def02af64a1fe6449c04e diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index 613c09ff2..788edb1e5 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -218,6 +218,9 @@ def connect_paste(button, local_paste_fields, input_comp, override_settings_comp else: try: valtype = type(output.value) + if hasattr(output, "step") and type(output.step) == float: + valtype = float + debug(f'Paste: "{key}"="{v}" type={valtype} var={vars(output)}') if valtype == bool and v == "False": val = False else: diff --git a/modules/ui_control.py b/modules/ui_control.py index 6aa6d0002..b10f2838e 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -645,6 +645,7 @@ def create_ui(_blocks: gr.Blocks=None): (mask_controls[5], "Mask dilate"), (mask_controls[6], "Mask auto"), # advanced + (cfg_scale, "Guidance scale"), (cfg_scale, "CFG scale"), (cfg_end, "CFG end"), (clip_skip, "Clip skip"), diff --git a/modules/ui_img2img.py b/modules/ui_img2img.py index e6de84f6e..294065498 100644 --- a/modules/ui_img2img.py +++ b/modules/ui_img2img.py @@ -255,6 +255,7 @@ def create_ui(): (subseed, "Variation seed"), (subseed_strength, "Variation strength"), # advanced + (cfg_scale, "Guidance scale"), (cfg_scale, "CFG scale"), (cfg_end, "CFG end"), (image_cfg_scale, "Image CFG scale"), diff --git a/modules/ui_txt2img.py b/modules/ui_txt2img.py index 42f57fc26..63d47920d 100644 --- a/modules/ui_txt2img.py +++ b/modules/ui_txt2img.py @@ -116,6 +116,7 @@ def create_ui(): (subseed, "Variation seed"), (subseed_strength, "Variation strength"), # advanced + (cfg_scale, "Guidance scale"), (cfg_scale, "CFG scale"), (cfg_end, "CFG end"), (clip_skip, "Clip skip"), From 760b41e99fea3f1ed69c66f7e23e341016c363e6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 2 Apr 2025 09:33:09 -0400 Subject: [PATCH 115/122] update requirements and changelog Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 16 ++++++++++------ installer.py | 2 +- modules/shared.py | 2 +- requirements.txt | 6 +++--- 4 files changed, 15 insertions(+), 11 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index dfec2e7ef..b73256a5d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,16 +4,20 @@ ### Highlights for 2025-04-02 -Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both *T2V* and *I2V* workflows +Time for another major release with ~120 commits and [ChangeLog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) that spans several pages! + +*Highlights?* +Video...Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1, Latte1** in both *T2V* and *I2V* workflows And combined with *on-the-fly quantization*, support for *Local/Tiny/Remote* VAE, acceleration modules such as *FasterCache or PAB*, and more! +Models...And support for new models: **CogView-4**, **SANA 1.5**, -Also, support for new models: **CogView-4**, **SANA 1.5**, - -Plus... +*Plus...* - New **Prompt Enhance** using LLM, - New **CLiP** models, improvements to **remote VAE**, additional wiki/docs/guides - More quantization options and granular control -- Pretty big performance updates to a) Any model using DiT based architecture: new caching methods, b) ZLUDA: new attention methods, c) much lower LoRA memory usage +- Pretty big performance updates to a) Any model using DiT based architecture due to new caching methods, b) ZLUDA with new attention methods, c) LoRA with much lower memory usage + +[ReadMe](https://github.com/vladmandic/automatic/blob/master/README.md) | [ChangeLog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) | [Docs](https://vladmandic.github.io/sdnext-docs/) | [WiKi](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.com/invite/sd-next-federal-batch-inspectors-1101998836328697867) ### Details for 2025-04-02 @@ -153,7 +157,7 @@ Plus... - updated [ZLUDA](https://github.com/vladmandic/sdnext/wiki/ZLUDA) guide - updated [OpenVINO](https://github.com/vladmandic/sdnext/wiki/OpenVINO) guide - updated [AMD-ROCm](https://github.com/vladmandic/sdnext/wiki/AMD-ROCm) guide - - upte [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide + - updated [Intel-ARC](https://github.com/vladmandic/sdnext/wiki/Intel-ARC) guide - **Fixes** - fix installer not starting when older version of `rich` is installed - fix circular imports when debug flags are enabled diff --git a/installer.py b/installer.py index 7154b5b3e..36c84e549 100644 --- a/installer.py +++ b/installer.py @@ -538,7 +538,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all or args.skip_git or args.experimental: return - sha = '75d7e5cc459f66a53652445d5b281054b297680d' # diffusers commit hash + sha = 'e5c6027ef89ec1a2800c0421599da89d4820f2e4' # diffusers commit hash pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' diff --git a/modules/shared.py b/modules/shared.py index 0235e2e31..dd8f8aad7 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -407,7 +407,7 @@ options_templates.update(options_section(('sd', "Models & Loading"), { "sd_checkpoint_autodownload": OptionInfo(True, "Model auto-download on demand"), "stream_load": OptionInfo(False, "Model load using streams", gr.Checkbox), "diffusers_eval": OptionInfo(True, "Force model eval", gr.Checkbox, {"visible": False }), - "diffusers_to_gpu": OptionInfo(False, "Load model directly to GPU"), + "diffusers_to_gpu": OptionInfo(False, "Model Load model direct to GPU"), "disable_accelerate": OptionInfo(False, "Disable accelerate", gr.Checkbox, {"visible": False }), "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles(), "visible": False}, refresh=refresh_checkpoints), "sd_checkpoint_cache": OptionInfo(0, "Cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": not native }), diff --git a/requirements.txt b/requirements.txt index 30a4e6065..f1de3abd7 100644 --- a/requirements.txt +++ b/requirements.txt @@ -41,18 +41,18 @@ torchsde==0.2.6 antlr4-python3-runtime==4.9.3 requests==2.32.3 tqdm==4.67.1 -accelerate==1.5.2 +accelerate==1.6.0 opencv-contrib-python-headless==4.9.0.80 einops==0.4.1 gradio==3.43.2 -huggingface_hub==0.29.3 +huggingface_hub==0.30.1 numexpr==2.8.8 numpy==1.26.4 numba==0.59.1 protobuf==4.25.3 pytorch_lightning==1.9.4 tokenizers==0.21.1 -transformers==4.50.1 +transformers==4.50.3 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 From 8e3ef400141f4428dd7cedf79cf4cf2dd6c08243 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Wed, 2 Apr 2025 23:19:41 +0900 Subject: [PATCH 116/122] zluda llm/vlm temp fix --- modules/zluda_hijacks.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index 431c359b8..bbbec7a81 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -1,7 +1,7 @@ from functools import wraps import torch import torch._dynamo.device_interface -from modules import shared, rocm, zluda # pylint: disable=unused-import +from modules import shared, zluda # pylint: disable=unused-import MEM_BUS_WIDTH = { @@ -18,6 +18,13 @@ MEM_BUS_WIDTH = { } +_topk = torch.topk +def topk(input: torch.Tensor, *args, **kwargs): # pylint: disable=redefined-builtin + device = input.device + values, indices = _topk(input.cpu(), *args, **kwargs) + return torch.return_types.topk((values.to(device), indices.to(device),)) + + class DeviceProperties: PROPERTIES_OVERRIDE = {"regs_per_multiprocessor": 65535, "gcnArchName": "UNKNOWN ARCHITECTURE"} internal: torch._C._CudaDeviceProperties @@ -42,6 +49,7 @@ def torch__C__cuda_getCurrentRawStream(device): def do_hijack(): + torch.topk = topk if zluda.default_agent is not None: DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = zluda.default_agent.name torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access From 4457ea2e51f40f2b707cd6c2137f81ff46215aff Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 2 Apr 2025 17:22:45 -0400 Subject: [PATCH 117/122] mark legacy scripts Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + TODO.md | 33 ++++----- scripts/allegrovideo.py | 2 +- scripts/cogvideo.py | 2 +- scripts/flux_prompt_enhance.py | 2 +- scripts/hunyuanvideo.py | 2 +- scripts/legacy_allegrovideo.py | 131 +++++++++++++++++++++++++++++++++ scripts/ltxvideo.py | 2 +- scripts/mochivideo.py | 2 +- 9 files changed, 154 insertions(+), 23 deletions(-) create mode 100644 scripts/legacy_allegrovideo.py diff --git a/CHANGELOG.md b/CHANGELOG.md index b73256a5d..5e6f1fb3e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,7 @@ Models...And support for new models: **CogView-4**, **SANA 1.5**, *Plus...* - New **Prompt Enhance** using LLM, +- New pipelines such as **InfiniteYou** - New **CLiP** models, improvements to **remote VAE**, additional wiki/docs/guides - More quantization options and granular control - Pretty big performance updates to a) Any model using DiT based architecture due to new caching methods, b) ZLUDA with new attention methods, c) LoRA with much lower memory usage diff --git a/TODO.md b/TODO.md index 1a9ea9094..c6cef03da 100644 --- a/TODO.md +++ b/TODO.md @@ -6,10 +6,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ### Issues/Limitations -- Video: Hunyuan Video I2V: requires `transformers==4.47.1` -- Video: CogVideoX 1.5 5B T2V/I2V: all-gray output -- Video: Allegro T2V: all-gray output -- Video: Latte1 T2V: garbage output +N/A ## Future Candidates @@ -25,17 +22,19 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Code TODO -- control: support scripts via api -- enable ROCm for windows when available -- fc: autodetect distilled based on model -- fc: autodetect tensor format based on model -- hypertile: vae breaks when using non-standard sizes -- infotext: handle using regex instead -- lora: add other quantization types -- lora: force-reloading entire model as loading transformers only leads to massive memory usage -- lora: required for flux to reapply offload after lora has been applied, but fails with oom -- lora: support pre-quantized flux -- model loader: implement model in-memory caching -- modernui: monkey-patch for missing tabs.select event -- processing: remove duplicate mask params +> pnpm lint | grep W0511 | awk -F'TODO ' '{print "- "$NF}' | sed 's/ (fixme)//g' + +- install: enable ROCm for windows when available - resize image: enable full VAE mode for resize-latent +- infotext: handle using regex instead +- fc: autodetect tensor format based on model +- fc: autodetect distilled based on model +- processing: remove duplicate mask params +- model loader: implement model in-memory caching +- hypertile: vae breaks when using non-standard sizes +- model load: force-reloading entire model as loading transformers only leads to massive memory usage +- lora: add other quantization types +- lora: maybe force imediate quantization +- modules/lora/lora_extract.py:185:9: W0511: TODO: lora support pre-quantized flux +- control: support scripts via api +- modernui: monkey-patch for missing tabs.select event diff --git a/scripts/allegrovideo.py b/scripts/allegrovideo.py index fdcd52adf..cf35500fb 100644 --- a/scripts/allegrovideo.py +++ b/scripts/allegrovideo.py @@ -31,7 +31,7 @@ def hijack_encode_prompt(*args, **kwargs): class Script(scripts.Script): def title(self): - return 'Video: Allegro' + return 'Video: Allegro (Legacy)' def show(self, is_img2img): return not is_img2img if shared.native else False diff --git a/scripts/cogvideo.py b/scripts/cogvideo.py index c18b4eb2f..de3c7736c 100644 --- a/scripts/cogvideo.py +++ b/scripts/cogvideo.py @@ -22,7 +22,7 @@ debug = (os.environ.get('SD_LOAD_DEBUG', None) is not None) or (os.environ.get(' class Script(scripts.Script): def title(self): - return 'Video: CogVideoX' + return 'Video: CogVideoX (Legacy)' def show(self, is_img2img): return shared.native diff --git a/scripts/flux_prompt_enhance.py b/scripts/flux_prompt_enhance.py index abfbeae6d..0ab087e1b 100644 --- a/scripts/flux_prompt_enhance.py +++ b/scripts/flux_prompt_enhance.py @@ -27,7 +27,7 @@ class Script(scripts.Script): prompt: gr.Textbox = None def title(self): - return 'Prompt enhance' + return 'Flux Prompt enhance (Legacy)' def show(self, is_img2img): return shared.native diff --git a/scripts/hunyuanvideo.py b/scripts/hunyuanvideo.py index dfd33e8ca..c39cec688 100644 --- a/scripts/hunyuanvideo.py +++ b/scripts/hunyuanvideo.py @@ -60,7 +60,7 @@ def hijack_encode_prompt(*args, **kwargs): class Script(scripts.Script): def title(self): - return 'Video: Hunyuan Video' + return 'Video: Hunyuan Video (Legacy)' def show(self, is_img2img): return not is_img2img if shared.native else False diff --git a/scripts/legacy_allegrovideo.py b/scripts/legacy_allegrovideo.py new file mode 100644 index 000000000..cf35500fb --- /dev/null +++ b/scripts/legacy_allegrovideo.py @@ -0,0 +1,131 @@ +import time +import gradio as gr +import transformers +import diffusers +from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer + + +repo_id = 'rhymes-ai/Allegro' + + +def hijack_decode(*args, **kwargs): + t0 = time.time() + vae: diffusers.AutoencoderKLAllegro = shared.sd_model.vae + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) + res = shared.sd_model.vae.orig_decode(*args, **kwargs) + t1 = time.time() + timer.process.add('vae', t1-t0) + shared.log.debug(f'Video: vae={vae.__class__.__name__} time={t1-t0:.2f}') + return res + + +def hijack_encode_prompt(*args, **kwargs): + t0 = time.time() + res = shared.sd_model.vae.orig_encode_prompt(*args, **kwargs) + t1 = time.time() + timer.process.add('te', t1-t0) + shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + return res + + +class Script(scripts.Script): + def title(self): + return 'Video: Allegro (Legacy)' + + def show(self, is_img2img): + return not is_img2img if shared.native else False + + # return signature is array of gradio components + def ui(self, is_img2img): + with gr.Row(): + gr.HTML('  Allegro Video
') + with gr.Row(): + num_frames = gr.Slider(label='Frames', minimum=4, maximum=88, step=1, value=22) + with gr.Row(): + override_scheduler = gr.Checkbox(label='Override scheduler', value=True) + with gr.Row(): + from modules.ui_sections import create_video_inputs + video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') + return [num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] + + def run(self, p: processing.StableDiffusionProcessing, num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument + # set params + num_frames = int(num_frames) + p.width = 8 * int(p.width // 8) + p.height = 8 * int(p.height // 8) + p.do_not_save_grid = True + p.ops.append('video') + + # load model + if shared.sd_model.__class__ != diffusers.AllegroPipeline: + sd_models.unload_model_weights() + t0 = time.time() + quant_args = model_quant.create_config() + transformer = diffusers.AllegroTransformer3DModel.from_pretrained( + repo_id, + subfolder="transformer", + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args + ) + shared.log.debug(f'Video: module={transformer.__class__.__name__}') + text_encoder = transformers.T5EncoderModel.from_pretrained( + repo_id, + subfolder="text_encoder", + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **quant_args + ) + shared.log.debug(f'Video: module={text_encoder.__class__.__name__}') + shared.sd_model = diffusers.AllegroPipeline.from_pretrained( + repo_id, + # transformer=transformer, + # text_encoder=text_encoder, + cache_dir=shared.opts.hfcache_dir, + torch_dtype=devices.dtype, + **quant_args + ) + t1 = time.time() + shared.log.debug(f'Video: load cls={shared.sd_model.__class__.__name__} repo="{repo_id}" dtype={devices.dtype} time={t1-t0:.2f}') + sd_models.set_diffuser_options(shared.sd_model) + shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) + shared.sd_model.sd_model_hash = None + shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode + shared.sd_model.vae.orig_encode_prompt = shared.sd_model.encode_prompt + shared.sd_model.vae.decode = hijack_decode + shared.sd_model.encode_prompt = hijack_encode_prompt + shared.sd_model.vae.enable_tiling() + # shared.sd_model.vae.enable_slicing() + + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + devices.torch_gc(force=True) + + processing.fix_seed(p) + if override_scheduler: + p.sampler_name = 'Default' + p.steps = 100 + p.task_args['num_frames'] = num_frames + p.task_args['output_type'] = 'pil' + p.task_args['clean_caption'] = False + + p.all_prompts, p.all_negative_prompts = shared.prompt_styles.apply_styles_to_prompts([p.prompt], [p.negative_prompt], p.styles, [p.seed]) + p.task_args['prompt'] = p.all_prompts[0] + p.task_args['negative_prompt'] = p.all_negative_prompts[0] + + # w = shared.sd_model.transformer.config.sample_width * shared.sd_model.vae_scale_factor_spatial + # h = shared.sd_model.transformer.config.sample_height * shared.sd_model.vae_scale_factor_spatial + # n = shared.sd_model.transformer.config.sample_frames * shared.sd_model.vae_scale_factor_temporal + + # run processing + t0 = time.time() + shared.state.disable_preview = True + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={num_frames}') + processed = processing.process_images(p) + shared.state.disable_preview = False + t1 = time.time() + if processed is not None and len(processed.images) > 0: + shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') + if video_type != 'None': + images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) + return processed diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py index 09f30360e..697e64021 100644 --- a/scripts/ltxvideo.py +++ b/scripts/ltxvideo.py @@ -52,7 +52,7 @@ def hijack_encode_prompt(*args, **kwargs): class Script(scripts.Script): def title(self): - return 'Video: LTX Video' + return 'Video: LTX Video (Legacy)' def show(self, is_img2img): return shared.native diff --git a/scripts/mochivideo.py b/scripts/mochivideo.py index e2c193c24..1e7ba5541 100644 --- a/scripts/mochivideo.py +++ b/scripts/mochivideo.py @@ -10,7 +10,7 @@ repo_id = 'genmo/mochi-1-preview' class Script(scripts.Script): def title(self): - return 'Video: Mochi.1 Video' + return 'Video: Mochi.1 Video (Legacy)' def show(self, is_img2img): return not is_img2img if shared.native else False From 6c29c3fb4d1dc107f7da0ed7261723d42e01f134 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 07:04:59 -0400 Subject: [PATCH 118/122] edit built-in styles Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 7 ++--- modules/ui_extra_networks.py | 52 +++++++++++++++++++----------------- 2 files changed, 32 insertions(+), 27 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5e6f1fb3e..9d95d2717 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2025-04-02 +## Update for 2025-04-03 -### Highlights for 2025-04-02 +### Highlights for 2025-04-03 Time for another major release with ~120 commits and [ChangeLog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) that spans several pages! @@ -20,7 +20,7 @@ Models...And support for new models: **CogView-4**, **SANA 1.5**, [ReadMe](https://github.com/vladmandic/automatic/blob/master/README.md) | [ChangeLog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) | [Docs](https://vladmandic.github.io/sdnext-docs/) | [WiKi](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.com/invite/sd-next-federal-batch-inspectors-1101998836328697867) -### Details for 2025-04-02 +### Details for 2025-04-03 - **Video tab** - see [Video Wiki](https://github.com/vladmandic/sdnext/wiki/Video) for details! @@ -178,6 +178,7 @@ Models...And support for new models: **CogView-4**, **SANA 1.5**, - fix token counter error style with modernui - fix sampler metadata when using default sampler - fix paste incorrect float to int cast + - do not allow edit of built-in styles - improve lora compatibility with balanced offload ## Update for 2025-02-28 diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index f96974d28..a73f33a56 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -175,9 +175,9 @@ class ExtraNetworksPage: img = None try: img = Image.open(f) - except Exception: + except Exception as e: img = None - shared.log.warning(f'Extra network removing invalid image: {f}') + shared.log.warning(f'Network removing invalid: image={f} {e}') try: if img is None: img = None @@ -189,9 +189,9 @@ class ExtraNetworksPage: img.close() created += 1 except Exception as e: - shared.log.warning(f'Extra network error creating thumbnail: {f} {e}') + shared.log.warning(f'Network create thumbnail={f} {e}') if created > 0: - shared.log.info(f"Network thumbnails: {self.name} created={created}") + shared.log.info(f'Network thumbnails: {self.name} created={created}') self.missing_thumbs.clear() def create_items(self, tabname): @@ -221,7 +221,7 @@ class ExtraNetworksPage: return self.patch(self.html, tabname) self_name_id = self.name.replace(" ", "_") if skip: - return f"

Extra network page not ready
Click refresh to try again
" + return f"
Network page not ready
Click refresh to try again
" subdirs = {} allowed_folders = [os.path.abspath(x) for x in self.allowed_directories_for_previews() if os.path.exists(x)] for parentdir, dirs in {d: files_cache.walk(d, cached=True, recurse=files_cache.not_hidden) for d in allowed_folders}.items(): @@ -239,7 +239,7 @@ class ExtraNetworksPage: if not subdir: continue subdirs[subdir] = 1 - debug(f"Networks: page='{self.name}' subfolders={list(subdirs)}") + debug(f'Networks: page="{self.name}" subfolders={list(subdirs)}') subdirs = OrderedDict(sorted(subdirs.items())) if self.name == 'model' and shared.opts.extra_network_reference_enable: subdirs['Local'] = 1 @@ -289,7 +289,7 @@ class ExtraNetworksPage: self.html += ''.join(htmls) self.page_time = time.time() self.html = f"
{subdirs_html}
{self.html}
" - shared.log.debug(f"Networks: type='{self.name}' items={len(self.items)} subfolders={len(subdirs)} tab={tabname} folders={self.allowed_directories_for_previews()} list={self.list_time:.2f} thumb={self.preview_time:.2f} desc={self.desc_time:.2f} info={self.info_time:.2f} workers={shared.max_workers}") + shared.log.debug(f'Networks: type="{self.name}" items={len(self.items)} subfolders={len(subdirs)} tab={tabname} folders={self.allowed_directories_for_previews()} list={self.list_time:.2f} thumb={self.preview_time:.2f} desc={self.desc_time:.2f} info={self.info_time:.2f} workers={shared.max_workers}') if len(self.missing_thumbs) > 0: threading.Thread(target=self.create_thumb).start() return self.patch(self.html, tabname) @@ -677,13 +677,13 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): if ui.gallery is not None: images = list(ui.gallery.temp_files) # gallery cannot be used as input component so looking at most recently registered temp files if len(images) < 1: - shared.log.warning(f'Extra network no image: item={ui.last_item.name}') + shared.log.warning(f'Network no image: item="{ui.last_item.name}"') return 'html/card-no-preview.png' try: images.sort(key=lambda f: os.path.getmtime(f), reverse=True) image = Image.open(images[0]) except Exception as e: - shared.log.error(f'Extra network error opening image: item={ui.last_item.name} {e}') + shared.log.error(f'Network error opening image: item="{ui.last_item.name}" {e}') return 'html/card-no-preview.png' fn_delete_img(image) if image.width > 512 or image.height > 512: @@ -691,9 +691,9 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): image.thumbnail((512, 512), Image.Resampling.HAMMING) try: image.save(ui.last_item.local_preview, quality=50) - shared.log.debug(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}"') + shared.log.debug(f'Networks save image: item="{ui.last_item.name}" filename="{ui.last_item.local_preview}"') except Exception as e: - shared.log.error(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}" {e}') + shared.log.error(f'Network save image: item="{ui.last_item.name}" filename="{ui.last_item.local_preview}" {e}') return image def fn_delete_img(_image): @@ -702,7 +702,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): for file in [f'{fn}{mid}{ext}' for ext in preview_extensions for mid in ['.thumb.', '.preview.', '.']]: if os.path.exists(file): os.remove(file) - shared.log.debug(f'Extra network delete image: item={ui.last_item.name} filename="{file}"') + shared.log.debug(f'Network delete image: item="{ui.last_item.name}" filename="{file}"') return 'html/card-no-preview.png' def fn_save_desc(desc): @@ -714,7 +714,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' with open(fn, 'w', encoding='utf-8') as f: f.write(desc) - shared.log.debug(f'Extra network save desc: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network save desc: item="{ui.last_item.name}" filename="{fn}"') return desc def fn_delete_desc(desc): @@ -722,7 +722,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): return desc fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' if os.path.exists(fn): - shared.log.debug(f'Extra network delete desc: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network delete desc: item="{ui.last_item.name}" filename="{fn}"') os.remove(fn) return '' return desc @@ -730,7 +730,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): def fn_save_info(info): fn = os.path.splitext(ui.last_item.filename)[0] + '.json' shared.writefile(info, fn, silent=True) - shared.log.debug(f'Extra network save info: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network save info: item="{ui.last_item.name}" filename="{fn}"') return info def fn_delete_info(info): @@ -738,14 +738,14 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): return info fn = os.path.splitext(ui.last_item.filename)[0] + '.json' if os.path.exists(fn): - shared.log.debug(f'Extra network delete info: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network delete info: item="{ui.last_item.name}" filename="{fn}"') os.remove(fn) return '' return info def fn_save_style(info, description, prompt, negative, extra, wildcards): if not isinstance(info, dict) or isinstance(info, list): - shared.log.warning(f'Extra network save style skip: item={ui.last_item.name} not a dict: {type(info)}') + shared.log.warning(f'Network save style skip: item="{ui.last_item.name}" not a dict: {type(info)}') return info if ui.last_item is None: return info @@ -753,7 +753,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): if hasattr(ui.last_item, 'type') and ui.last_item.type == 'Style': info.update(**{ 'description': description, 'prompt': prompt, 'negative': negative, 'extra': extra, 'wildcards': wildcards }) shared.writefile(info, fn, silent=True) - shared.log.debug(f'Extra network save style: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network save style: item="{ui.last_item.name}" filename="{fn}"') return info def fn_delete_style(info): @@ -761,7 +761,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): return info fn = os.path.splitext(ui.last_item.filename)[0] + '.json' if os.path.exists(fn): - shared.log.debug(f'Extra network delete style: item={ui.last_item.name} filename="{fn}"') + shared.log.debug(f'Network delete style: item="{ui.last_item.name}" filename="{fn}"') os.remove(fn) return {} return info @@ -785,6 +785,10 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): if 'modelVersions' in fullinfo: # sanitize massive objects fullinfo['modelVersions'] = [] info = fullinfo + if isinstance(info, list): + item.filename = None + shared.log.warning('Network: show details not supported for compound item') + info = None """ if prompt is not None: item.prompt = prompt @@ -812,7 +816,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): model = '' style = '' note = '' - if not os.path.exists(item.filename): + if item.filename is not None and not os.path.exists(item.filename): note = f'
Target filename: {item.filename}' if page.title == 'Model': merge = len(list(meta.get('sd_merge_models', {}))) @@ -904,7 +908,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): page.refresh_time = 0 page.refresh() page.create_page(ui.tabname) - shared.log.debug(f"Networks: refresh page='{page.title}' items={len(page.items)} tab={ui.tabname}") + shared.log.debug(f'Networks: refresh page="{page.title}" items={len(page.items)} tab={ui.tabname}') pages.append(page.html) ui.search.update(title) return pages @@ -918,7 +922,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): page.card = card_full if page.view == 'gallery' else card_list page.html = '' page.create_page(ui.tabname) - shared.log.debug(f"Networks: refresh page='{page.title}' items={len(page.items)} tab={ui.tabname} view={page.view}") + shared.log.debug(f'Networks: refresh page="{page.title}" items={len(page.items)} tab={ui.tabname} view={page.view}') pages.append(page.html) ui.search.update(title) return pages @@ -973,9 +977,9 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): } shared.writefile(item, fn, silent=True) if len(prompt) > 0: - shared.log.debug(f"Network quick save style: item={name} filename='{fn}' unparsed={shared.opts.extra_networks_unparsed}") + shared.log.debug(f'Network quick save style: item="{name}" filename="{fn}" unparsed={shared.opts.extra_networks_unparsed}') else: - shared.log.warning(f"Network quick save model: item={name} filename='{fn}' prompt is empty") + shared.log.warning(f'Network quick save model: item="{name}" filename="{fn}" prompt is empty') def ui_sort_cards(sort_order): if shared.opts.extra_networks_sort != sort_order: From 5bdc87b68af0206d3e0e5cb3804de0ad9ce803e1 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 08:22:43 -0400 Subject: [PATCH 119/122] fix server restart Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + javascript/script.js | 1 + javascript/settings.js | 5 +++-- javascript/startup.js | 9 ++++++++- modules/shared.py | 2 +- modules/ui_extra_networks.py | 5 +++++ modules/ui_loadsave.py | 7 ------- modules/ui_settings.py | 5 ++++- webui.py | 2 ++ 9 files changed, 25 insertions(+), 12 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 9d95d2717..72b8d6d30 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -178,6 +178,7 @@ Models...And support for new models: **CogView-4**, **SANA 1.5**, - fix token counter error style with modernui - fix sampler metadata when using default sampler - fix paste incorrect float to int cast + - fix server restart from ui - do not allow edit of built-in styles - improve lora compatibility with balanced offload diff --git a/javascript/script.js b/javascript/script.js index f943f4626..8c270bb63 100644 --- a/javascript/script.js +++ b/javascript/script.js @@ -65,6 +65,7 @@ function onOptionsChanged(callback) { function executeCallbacks(queue, arg) { // if (!uiLoaded) return for (const callback of queue) { + if (!callback) continue; try { callback(arg); } catch (e) { diff --git a/javascript/settings.js b/javascript/settings.js index 7a43969ac..99902fb02 100644 --- a/javascript/settings.js +++ b/javascript/settings.js @@ -26,7 +26,7 @@ async function updateOpts(json_string) { const key = Object.keys(op)[0]; const callback = op[key]; if (opts[key] && opts[key] !== settings_data.values[key]) { - log('updateOpts', key, opts[key], settings_data.values[key]); + log('updateOpt', key, opts[key], settings_data.values[key]); if (callback) callback(new_opts[key], opts[key]); } } @@ -37,7 +37,8 @@ async function updateOpts(json_string) { if (callback) callback(new_opts[key], opts[key]); } - opts = new_opts; + window.opts = new_opts; + log('updateOpts', Object.keys(new_opts).length); Object.entries(opts_metadata).forEach(([opt, meta]) => { if (!opts_tabs[meta.tab_name]) opts_tabs[meta.tab_name] = {}; if (!opts_tabs[meta.tab_name].unsaved_keys) opts_tabs[meta.tab_name].unsaved_keys = new Set(); diff --git a/javascript/startup.js b/javascript/startup.js index d8407391d..328c167b6 100644 --- a/javascript/startup.js +++ b/javascript/startup.js @@ -3,6 +3,7 @@ window.api = '/sdapi/v1'; window.subpath = ''; async function initStartup() { + const t0 = performance.now(); log('initStartup'); if (window.setupLogger) await setupLogger(); @@ -24,7 +25,11 @@ async function initStartup() { await reconnectUI(); // make sure all of the ui is ready and options are loaded - while (Object.keys(window.opts).length === 0) await sleep(50); + let t1 = performance.now(); + while ((Object.keys(window.opts).length === 0) && (t1 - t0 < 10000)) { + t1 = performance.now(); + await sleep(50); + } log('mountURL', window.opts.subpath); if (window.opts.subpath?.length > 0) { window.subpath = window.opts.subpath; @@ -43,6 +48,8 @@ async function initStartup() { setHints(); applyStyles(); initIndexDB(); + t1 = performance.now(); + log('initStartup', Math.round(1000 * (t1 - t0) / 1000000)); } onUiLoaded(initStartup); diff --git a/modules/shared.py b/modules/shared.py index dd8f8aad7..63d01446c 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -1250,7 +1250,7 @@ total_tqdm = TotalTQDM() def restart_server(restart=True): if demo is None: return - log.warning('Server shutdown requested') + log.critical('Server shutdown requested') try: sys.tracebacklimit = 0 stdout = io.StringIO() diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index a73f33a56..e16a26bbf 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -56,6 +56,9 @@ preview_map = None def init_api(): def fetch_file(filename: str = ""): + global allowed_dirs # pylint: disable=global-statement + if len(allowed_dirs) == 0: + allowed_dirs = shared.demo.allowed_paths if not os.path.exists(filename): return JSONResponse({ "error": f"file {filename}: not found" }, status_code=404) if filename.startswith('html/') or filename.startswith('models/'): @@ -473,6 +476,8 @@ def register_page(page: ExtraNetworksPage): def register_pages(): debug('EN register-pages') + shared.extra_networks.clear() + allowed_dirs.clear() from modules.ui_extra_networks_checkpoints import ExtraNetworksPageCheckpoints from modules.ui_extra_networks_vae import ExtraNetworksPageVAEs from modules.ui_extra_networks_styles import ExtraNetworksPageStyles diff --git a/modules/ui_loadsave.py b/modules/ui_loadsave.py index 1135134d2..594b44c1d 100644 --- a/modules/ui_loadsave.py +++ b/modules/ui_loadsave.py @@ -13,7 +13,6 @@ class UiLoadsave: def __init__(self, filename): self.filename = filename self.component_mapping = {} - self.finalized_ui = False self.ui_defaults_view = None # button self.ui_defaults_apply = None # button self.ui_defaults_review = None # button @@ -24,8 +23,6 @@ class UiLoadsave: self.ui_settings = self.read_from_file() def add_component(self, path, x): - """adds component to the registry of tracked components""" - assert not self.finalized_ui def apply_field(obj, field, condition=None, init_field=None): key = f"{path}/{field}" @@ -253,7 +250,6 @@ class UiLoadsave: return "Restored system defaults for user interface" def create_ui(self): - """creates ui elements for editing defaults UI, without adding any logic to them""" with gr.Row(elem_id="config_row"): self.ui_defaults_apply = gr.Button(value='Set UI defaults', elem_id="ui_defaults_apply", variant="primary") self.ui_defaults_submenu = gr.Button(value='Set UI menu states', elem_id="ui_submenu_apply", variant="primary") @@ -262,9 +258,6 @@ class UiLoadsave: self.ui_defaults_review = gr.HTML("", elem_id="ui_defaults_review") def setup_ui(self): - """adds logic to elements created with create_ui; all add_block class must be made before this""" - assert not self.finalized_ui - self.finalized_ui = True self.ui_defaults_view.click(fn=self.ui_view, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review]) self.ui_defaults_apply.click(fn=self.ui_apply, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review]) self.ui_defaults_restore.click(fn=self.ui_restore, inputs=[], outputs=[self.ui_defaults_review]) diff --git a/modules/ui_settings.py b/modules/ui_settings.py index 06d37c251..bda278a81 100644 --- a/modules/ui_settings.py +++ b/modules/ui_settings.py @@ -3,9 +3,9 @@ import gradio as gr from modules import timer, shared, paths, theme, sd_models, modelloader, ui_common, ui_loadsave, generation_parameters_copypaste, call_queue, script_callbacks +text_settings = None # holds json of entire shared.opts ui_system_tabs = None # required for system-info dummy_component = gr.Textbox(visible=False, value='dummy') -text_settings = gr.Textbox(elem_id="settings_json", value=lambda: shared.opts.dumpjson(), visible=False) loadsave = ui_loadsave.UiLoadsave(shared.cmd_opts.ui_config) quicksettings_names = {x: i for i, x in enumerate(shared.opts.quicksettings_list) if x != 'quicksettings'} quicksettings_list = [] @@ -168,6 +168,8 @@ def run_settings_single(value, key, progress=False): def create_ui(): + global text_settings # pylint: disable=global-statement + text_settings = gr.Textbox(elem_id="settings_json", elem_classes=["settings_json"], value=lambda: shared.opts.dumpjson(), visible=False) with gr.Row(elem_id="system_row"): restart_submit = gr.Button(value="Restart server", variant='primary', elem_id="restart_submit") shutdown_submit = gr.Button(value="Shutdown server", variant='primary', elem_id="shutdown_submit") @@ -203,6 +205,7 @@ def create_ui(): shared.log.debug(f'UI settings: sections={len(sections)} settings={len(list(shared.opts.data_labels))}') with gr.Tabs(elem_id="settings"): + quicksettings_list.clear() for (section_id, section_text) in sections: items = [item for item in shared.opts.data_labels.items() if item[1].section[0] == section_id] # find all items in this section hidden = section_id is None or 'hidden' in section_id.lower() or 'hidden' in section_text.lower() diff --git a/webui.py b/webui.py index c59aa5647..1996c41b5 100644 --- a/webui.py +++ b/webui.py @@ -293,6 +293,8 @@ def start_ui(): allowed_paths = [os.path.dirname(__file__)] if shared.cmd_opts.data_dir is not None and os.path.isdir(shared.cmd_opts.data_dir): allowed_paths.append(shared.cmd_opts.data_dir) + if shared.cmd_opts.models_dir is not None and os.path.isdir(shared.cmd_opts.models_dir): + allowed_paths.append(shared.cmd_opts.models_dir) if shared.cmd_opts.allowed_paths is not None: allowed_paths += [p for p in shared.cmd_opts.allowed_paths if os.path.isdir(p)] shared.log.debug(f'Root paths: {allowed_paths}') From 5c6c1465f436d2c82ca8fbaf0be5e4af5592c56b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 10:03:48 -0400 Subject: [PATCH 120/122] fix style apply params Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 1 + modules/lora/extra_networks_lora.py | 2 +- modules/lora/lora_apply.py | 4 ++-- modules/styles.py | 10 ++++++---- 4 files changed, 10 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 72b8d6d30..0e73825b3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -179,6 +179,7 @@ Models...And support for new models: **CogView-4**, **SANA 1.5**, - fix sampler metadata when using default sampler - fix paste incorrect float to int cast - fix server restart from ui + - fix style apply params - do not allow edit of built-in styles - improve lora compatibility with balanced offload diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index fe0e15de5..0f981d93a 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -205,5 +205,5 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): shared.log.debug(f"Network end: type=LoRA time={l.timer.summary}") if self.errors: for k, v in self.errors.items(): - shared.log.error(f'LoRA: name="{k}" errors={v}') + shared.log.error(f'Network: type=LoRA name="{k}" errors={v}') self.errors.clear() diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index e3dae287a..8bf327151 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -101,9 +101,9 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. l.timer.move += t1 - t0 except RuntimeError as e: l.extra_network_lora.errors[net.name] = l.extra_network_lora.errors.get(net.name, 0) + 1 + module_name = net.modules.get(network_layer_name, None) + shared.log.error(f'Network: type=LoRA name="{net.name}" module="{module_name}" layer="{network_layer_name}" apply weight: {e}') if l.debug: - module_name = net.modules.get(network_layer_name, None) - shared.log.error(f'LoRA apply weight name="{net.name}" module="{module_name}" layer="{network_layer_name}" {e}') errors.display(e, 'LoRA') raise RuntimeError('LoRA apply weight') from e continue diff --git a/modules/styles.py b/modules/styles.py index ecc3204f7..2de41c2a8 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -154,6 +154,7 @@ def apply_styles_to_extra(p, style: Style): p.original_negative = p.negative_prompt style_extra = apply_wildcards_to_prompt(style.extra, [style.wildcards], silent=True) + style_extra = ' ' + style_extra.lower() extra.update(infotext.parse(style_extra)) extra.pop('Prompt', None) extra.pop('Negative prompt', None) @@ -162,14 +163,15 @@ def apply_styles_to_extra(p, style: Style): for k, v in extra.items(): k = k.lower() k = k.replace(' ', '_') - if k in name_exclude: # exclude some fields - continue if k in name_map: # rename some fields k = name_map[k] + if k in name_exclude: # exclude some fields + continue if hasattr(p, k): orig = getattr(p, k) - if type(orig) != type(v) and orig is not None: - v = type(orig)(v) + if (type(orig) != type(v)) and (orig is not None): + if not (type(orig) == int and type(v) == float): # dont convert float to int + v = type(orig)(v) setattr(p, k, v) fields.append(f'{k}={v}') else: From 45b01dde8155d401885096914b23debcd63fe793 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 10:19:27 -0400 Subject: [PATCH 121/122] fix typo Signed-off-by: Vladimir Mandic --- scripts/xyz_grid_draw.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/xyz_grid_draw.py b/scripts/xyz_grid_draw.py index bac96bb8b..d654a2c7f 100644 --- a/scripts/xyz_grid_draw.py +++ b/scripts/xyz_grid_draw.py @@ -128,6 +128,6 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend processed_result.infotexts.insert(0, processed_result.infotexts[0]) t2 = time.time() - shared.log.info(f'XYZ grid complete: images={list_size} results={len(processed_result.images)}size={grid.size if grid is not None else None} time={t1-t0:.2f} save={t2-t1:.2f}') + shared.log.info(f'XYZ grid complete: images={list_size} results={len(processed_result.images)} size={grid.size if grid is not None else None} time={t1-t0:.2f} save={t2-t1:.2f}') p.skip_processing = True return processed_result From a5f796a16613922978e9d0d3bacd6d9b0f32a4ff Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 10:24:55 -0400 Subject: [PATCH 122/122] change field used for true cfg Signed-off-by: Vladimir Mandic --- modules/processing_diffusers.py | 2 +- modules/processing_info.py | 1 + modules/video_models/video_run.py | 2 +- 3 files changed, 3 insertions(+), 2 deletions(-) diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index e61ca1afc..644e0b276 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -71,7 +71,7 @@ def process_base(p: processing.StableDiffusionProcessing): eta=shared.opts.scheduler_eta, guidance_scale=p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, - true_cfg_scale=p.diffusers_guidance_rescale, + true_cfg_scale=p.pag_scale, denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None, denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None, num_frames=getattr(p, 'frames', 1), diff --git a/modules/processing_info.py b/modules/processing_info.py index d254168cc..4b57d859d 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -54,6 +54,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No "CFG scale": p.cfg_scale if p.cfg_scale > 1.0 else None, "CFG rescale": p.diffusers_guidance_rescale if p.diffusers_guidance_rescale > 0 else None, "CFG end": p.cfg_end if p.cfg_end < 1.0 else None, + "CFG true": p.pag_scale if p.pag_scale > 1 else None, "Clip skip": p.clip_skip if p.clip_skip > 1 else None, "Batch": f'{p.n_iter}x{p.batch_size}' if p.n_iter > 1 or p.batch_size > 1 else None, "Model": None if (not shared.opts.add_model_name_to_info) or (not shared.sd_model.sd_checkpoint_info.model_name) else shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', ''), diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 1c939936a..2b4fbcbea 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -39,7 +39,7 @@ def generate(*args, **kwargs): denoising_strength=float(init_strength), init_image=init_image, cfg_scale=float(guidance_scale), - diffusers_guidance_rescale=float(guidance_true), + pag_scale=float(guidance_true), vae_type=vae_type, vae_tile_frames=int(vae_tile_frames), override_settings=override_settings,