From efc1b95f3346fbd4da21a1ab49d083d9263fcc8a Mon Sep 17 00:00:00 2001 From: Kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 16 Jan 2024 17:10:07 +0200 Subject: [PATCH] Tiled vae, batch support, progress bar Breaks old nodes as they are, should be remade. --- __init__.py | 3 +- empty_text_embed.pt | Bin 0 -> 8966 bytes ldm/models/diffusion/ddpm_ccsr_stage1.py | 58 +- ldm/models/diffusion/ddpm_ccsr_stage2.py | 108 +-- model/q_sampler.py | 856 +++++++++++++---------- nodes.py | 81 ++- utils/tilevae.py | 719 +++++++++++++++++++ 7 files changed, 1357 insertions(+), 468 deletions(-) create mode 100644 empty_text_embed.pt create mode 100644 utils/tilevae.py diff --git a/__init__.py b/__init__.py index 5109219..2e96bd6 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,3 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS -WEB_DIRECTORY = "./web" -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] \ No newline at end of file +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/empty_text_embed.pt b/empty_text_embed.pt new file mode 100644 index 0000000000000000000000000000000000000000..dcedcdecbf8ec649a9f088f3a9e50996dce8f89b GIT binary patch literal 8966 zcmZ{K2{=~K*1xIDb21j12$d2k?>U=F5h9gR6d?)8SPCJE3?Va>D3VMiiFdE{Hb_dT zWL7H8X^;k`;nV+q|L?i?`|e%m*=O&w&)UCduQRQ+*ZJ952@3J?iHPw1A5rEL<@4~} z85FuX$YXcVW)I&j9`1A8U4vX_@AURrws;)hhQAUi;1=W;;I?h{=HMVNpFpwA0Uld| zy?orq2JHy+3)mboPt2cBTYJ6m7=GZGrmLriKmR|5h%NK+a}8Skk4`avf%W`jI{N;C z>jlSZx1GWMLN+aKB7Wj++X}}}{}#LdXt5%n^1ruuj{dU6 z3jeV8{~7&HyTzdea9`^=RE-zr7KV61Rj(2I?YSQ;)pmi|(sEoEUIBQsmaIa%JAE_$ z8VKHOr2Ro3P1H`Ng6rpO=4MhAy|{2a$nFhhjtKNJ%O2YE=OKiltElKUVAtpe~9C?HXme+RG{R0FUNdH74x}y+{tMx z853(UuIcq-&{oHRbs?%SIvCG<5L979%lql@NDaO2#*af~O~9Kp#C03!pq`XKw82Bz zFL@eNZXLv@hZb|9s?8W-!K-wBX*PW%l?m?Q{WL)<5@h(U!E`}==vgFAXY85`FUnZh zml_NzEuBnth&z1gHHYdbQ^>UMgZH?Rwnz>`VqYv+Tg*hBS0I#KK{(YtU}rLI<;pAWtxf%I)t5b4w0*uWg2(3l$ooa^ur?$c9zKr-5;w@ei)fxzuMRkSCqtIL9E#tS zggE0pkidBWSzkio&y7R0x>^Pp?{_rTRURUrWpQ74sPmFDC&Mf)ABg^b6~dN%q%YJ; zX_A^goX9GMtz8)){H>UdQ)>ge#N#l*S%mAiu)u`F?*S?`Dv+S_1{5sBpv^=J!mCQC z*Ki!8=Tyci?Hs0b#{;@CJ_ky)zrdB8LfUj>H~0IYT==9_1<(9qK#+Y7wG*0QkLNwE z@2(us`S6fY)kuJO#+QI!Y(Gpdvxh`EH~PM97jxVA4Y>FFg2KzW&|1|BvvszE^eRVs z-R~8*cFkfS!v=Lu7qxNeyE&jx`307{NQ3mPyKq%xG8<%QWs+apWzy%j1~#Am31;KO zO(RTR!H7#L*yiNIDvSj4`Dz%DSDQ zx&uNW;q(e}S`(QWMZKJ-p)+A~LMNP+Qirdi-c|ZxgfBS)uSgThPBG^6A@lju#Md14lK2$Ui2 zfdICdzXj%{IxuCs7(eZKjPG^{co&gG3$~LDDoH?Jr1z`QPS1|DHLZ8a@D}@|)qE;vtxQr4er2J^{bab;01O!*HY{ z8YGfu!xrC2X29DNe1{)49N6s-_0wzUjzV|1I#nNryT5_U=F{-Tk}#J?bZN|~0s3n< z<%G_5Y)HMLjP8Dyn8jjd5MOZsCa>nf#RKEf_VNd6lXw{vuIz@lBAsAL;+XlVl~Cdk z$2?wN4K}6XaP6Koj1Rg7{tYiUTIB|e(&R+$V{3?yJoyhhHTnB+K zt=z^jTRrRBfq%ypNIt#-Mgl}I@Io25ElXxXRY9${mcmVqkD(HFIHlXhWY|1z3K`qpzBxAg$ZHUhMcT zI^bCYksfkPsKaJhzs(oMJ(&$2H?ldxhmo0c(Gq$UMVZ{yGBobMe1@F0hf7&nGq zd05UEfumsw*a}~O56T)4Hd-`xZ6r9k7uu=Cj0UiB$b+{RD&fMK`yghnO3&VXO}8dz z!la07Mta|7@VbAUbEDN0K;P{_h=vWa<_+y^Wxz5aRJmlxC#AlEFov|9ISgej&`k_Ouw0@(s|?WP`7$@P>j0< z4~vW7tj_@Bc&{05Unydo+xW0hSQA`54w|T*9i<(scEW9g40!L@Mn#XO!7A}wj^IUA z5Vt=G{nAG{37^OAn~VgG?wi6se!d@;eU5+?JCDQWZ%tGsd+s+nk+}mqpKHVSnWdaM-2(8`YNHD~_%OB50G7$vz`*UfaJ_LAB+fF0D{H>f zsL$c_LzNV@O}2rRqE4{GGz?_hIvak-UZ=AUya)BYEkHJIfuq~5(Q1|-s@u8r_aX(r zs!b58nhsc#0yo4LK+lR{&bt@o(qDqV`DGqHf8~-0B9ExjVsS zt2{N-Gyp+f032)o0tG*VL1^_=+NvA@C!Cw0cVr&?c}1Zjb~bm8wIal>+Sbr>vlzD7 zBE?#1pu@jbSF3ap2z6we)y+ z7(8ryP6xv7(B8{4z-~o0%#ykV(Ro8mX#azTkh)^XS3L(-MjYzk*+@&us$iWG;VL;- z!?)lxsODq9AU_H2jW}>$ti1`UUILgn{~5hkz(RtgGt}oTgj7*CroSi?UQCT;PT!Y@ zD_f%J*&GCk+;0#!=SPM?7X>hUA00cD4D67#v7DHU(RGn;xRQda? ze(OH)mWbq}ofW00XH2KhRtup;z#6KfaDi6LwuN)S2$Ftj&|y0tIL|U5P=rE;Lm6!@ zt_3SzBV3z#fC(GYVOxhBAiZ$7VXB}IJ9|%T!-tcbpt3O?T4Lt;dF_{cUo-w366C3(FRX~#Rr6el`Mv9x2$Q=sWg!7 zYh&^!&w;&4P0$!Az%JlF3mf0&!sn$y5W%kircqO%B7HY#hRD)o>7S{eWdQ^iykTm5 zj=-NK;#e1c6db-*(L>oGFgP<6raSh6S!gBYchiF<$D&!0HTg_-O(PiD+d`9*2t?&Q zq~Y`r*pAjSz5bTa@HLN{iz*=KQw1KA3gJSvI$fzD&GJ%2*n93u^m;-z*JP_0WKBw; zb3R;w+r|}e$~}p?E_Y@3EOz8E%~ zP$-)ZqY9SwdP?0;k{VAr3g_|3*n+?4Xe<2R`9qItN+e&c#%Rvx8 zD~NE>hk9iVc3-dq=jvK5dJattQ*6v697HT0L-1<>|Zs_)I44jn24 zj4W2dTDlOL$JQf(!`I-l`FJ!?w&YBU*Z>poC?t)W4~nsQa5hI6ufJD^fFoH@^<+C! zFH6BGu#ei?mxJ|3d+1TA>`q$zaWX4@ z?FFN%EeDJ4h`^ma19b6P9=s1oh68seQq}cOsgi6sn;9Mjk1A!5bR*>Il)&4*4Q%a^ z%k<>&vk)nC7G{611l?;bbTNMuY_giknRr;4oh!NuybW$b?YSm6+Rg`eA5361{30Ra zc|6_Q)=hV33Q@_ri_nsElNPk-L+-Ct%#fP}G%V!6otqT4xvI0OBbBhbCJp8fCDVmt z{$8i=FbICv0(wmt1+^uhxtAaM&ey=v_Y=XN^BM*X>)`IpL;l7vOVod%7uzb94SRw#xQW4;P;GOMX0|#(ynH-Zm93+! zjU8JyNfui7FzmpCpUjWBO|-{X44>Hb!J^BTXb%;`j^OPuAQDO0ff4F@_ZRq{TMg!| zZQP=6b*5qMD{vMof^TX)5E4t_-IPsGWn9kKk2j%OngdM5!lFg@OINdVA21Z1{Gek^ z76{7vGI95{V0Q3E`Zy>StX-5~#;2)ZuO^NL+CHH5LINv{7eM7MeI_A4uZRyIacFHa6r(EkJkikRPb>R?Kaf>U64QGLM?;Gy0;RZPMaSijp!+|!x zi-L=9)0y1%eK5CH8K$}4qI^#)AwVmd9p5QLbyKq-v80B)vK%q7P3;B6n2U_z$$Cy( zs5GT#5|Mjt8SC#ZYSgvp9)^*VU7-w#SCnS}UOJv_6i(UKO-7Qo)##3_N5V z$r{CXK)aD7b}q;Ri&k@@v2O<%%#nrJC~NBWLIxk*iXxn<5d12xi^(5t@nm-;nb&j> z2k3rm8|lCcpEL;aezN)F^@!cpePn@sGzMjwQms;i)M@^#ObQaV=mx4idCsPOSi^>& zKZxT*^x??G$+$CZ8Vs#8BT}y?;&FLpFPv{lz6@P;+{J@#&;UKwJ&Kmkwa?;}G|&8%c{KDnIB!=OppxcqJ!DGYpw`EDPu zqxvfGzC8;o{?uXQk|v^Zq6D7H6Vr9|3NZJIA#aj<9=o|>0sczAPSeL?IXRPA#3WOK zH@C|Pk9`ou?8je|7|W1fcP_II2d&A}tI=qicn{i3Wle>|^1$fS0Cu*=5^-BfPQ@6KOim+a zI!>`uP%BK9_aXf-@}^az(mVlFdlrKwr~UBKdIgLZ zi^tZaBz(r5!7FHwV9vE%Wp5PRhpurHMjQ9y>xM;`(zp=h?hJw7QCa*l<0;jdqRUg6 zm(0pUZ$yVLzQiE?2TAIXWX}hF!PKU!C^{%g)`bg@ZLXd;!heM}IgI0-U+6<7FloT| zbtTrsAPxVmjN-Qo@V<=}aqB2yTl4eik)7M1iZ2(0&kNzI&9_l|;1>C|;sO~K_r`fg zlHlRw#aNJ&fodaf_49f9quI&Et?+et?+0KaZ1N zA4bC*L3D{bh-<|j;^nU?$m2d_pFf_)3lkOP{Z_QX-;R95IBt~8TNOrDygJKXPFap& zFE)^s4`$%ipWncB>nu~#E-4~#bs|hUFNyN9X4ql-2%To7kc|rykuM+yc2=;a)3VCg zM>8+7Rar>mHf6D^lLJBSWHYfFe!%=Z)=gelEW(C&L2YhD6}_%gf0|XE}HwHb^zsh>%XL5cJr$pMGpSN7xH(BvI%#=1y3K z0y~n}0D~?lGIS=}4v3KJcDL9?L3!vS-pmdL@!>}+UwD0yOANi*iNu@tWX9HF46*WH zy*CJAz0zawiqAl~DsxnE>m*NCsT0K+5~j_`$~ftz7OuUYhzXZpf^=yHC@-3B>S$ku zul#qish4^%Gpz=NKk=je(hYcc^)~bjbRz!CrFhL%Ram=l5!t7_k$fm!juYAiOuuFI z!J$4&(sLjWIiKoj)+IsmcGp*o`y>SqCZw}r4^>FJ+fMRpJr0sn_J7O0}!qa%Hb&xK% zUwI9hLxgGOTV%PSO3O%jY+1*nN@= zNsK4);v1OEtE-`Ly%oxOWuR2=3dWst7^d`H;cR-spfkqd9XluTy(g0lDtVD~$5$}k z@-=o_`rsltKJ0&Wk#!2RA$yOM;}qpn;J-Q<6Zm4$p=Ugf%s++aj*hJ@zvSZY@Ui&I zPSkY9Eq&S;Mq&430p7SdeXu@T8aIs0gu%%7FwiW=yYe|6kKaweTxg?*ils5o%$r&n z{KncaL7b&MpXf`e^G-c4#%7szlnDOHPH=F)ehxZ@ z93!8}buc`fM|TC^=WLVD!^@BB7@egeyc1vUp@60uj1C8|&jKn*<2`B9T8#!WHLQT$ zo6SdPxg<=>Da7c-9OCAoP5d7BL3>LXcIDn-96qgOm*1Ys^%Yr5LYAE-Mt)r|Gwv21 zzp8?%{sC<0?%Vj*Qw_ItH=#h}UGl_24_oqPqj>8$ykol^mE=Fd1lM@BOYI1&%iQB8 zO&MToVp@sY{sfF&EQQM=VuVq`iAzyF@l2n} z&UY%omacsyJw$_dVoDy0C>yZr`m*ri@KF@JJw!(>9%HSqFte`KkG=k38E^lav2Sm& z8tGDV!RW{ntm%>@?0XZ-p1z%eJw^m)GLPWqrz*03X&XptKEYqVmykW*kD16fh_RXi z=BQjc9`b**uoinRk|>=t#?B&!;kTZMOIs(PbwoOdzbnS9enHl*E^#bJtw?OvDVVyQ z636R%rSR#{Gxq%9m1yPJz(jS=$G+ZL@M=lK>sO_4c)?Sc`>BOp9;AaOR2y*8!9`?w zN(+`(tFhz~KWx5Z$jklt7B}^W!L_(qg#T(kJ8jNt(}N!Aq}6Z%8XK<0B=MK5k(v^( zZSq76SJ;GRjVGAkRbp&LQ85ULjj%EX8APIO8YqvbLc#}ed=|D6-6yB9+B0<6f~|Zw zPw@bT6gHDRp?plLOE7D+TpP7I`JlexAU1iF5EbWKc+_Z5P<2?A`2= z!5a`)FqteVnTXuvUfkX103%;lqo+eD8^sQ4-h_wkHAvD*F6w2fgV>+9EZ_7p;?n8MUXj&eN4mAibEiXN=i3?6eoYd~ zR}3`=W3gH20Dd`Li8j`bctt^o8$;VA9N3`>B!!SuT6yi>ayo>QtPjc>~kG zZzZ>KjiBPJFH)2`H!8LCCj5#H;hp#p^9*NZh_vxW3c)@;2Y;% zWS-4q_RB_fqTXG>(OB%yM)sb*awEohTxT0Fge?hMC|`;C*mg;6E91K ztU294r7Rir=j|qI3U7kcM^l_gSS=GO+QQZvi18l$ODyO zc3mT7b|tML@=NC8+-3`sv%wvAzM0I6l#*tQRnMa4wqf{_ng)tdCb-XV341Di6r|hD zc?Bnm@WD@Q{OP3v7C-bDA>;ix8X~|d3j9Wi>}E*+dK`cMNX4uJ17zCqy==Cj4W5T( zWPinVc(OdyhMI zpF!7DVN)mPa+JHKjc^K4RP-(@FEWm|U{?uxKeAzVJXXdou_$u-yeD(_)h;|cuA8%< zYAUO!u#PC|9Kk;^bLoJ+9jrN&NLHAqVfXhG_%7QEAA1r=X{R1$q)kE-B~8=sTh|eZ z5AO*_$bjv8dl#&i4RJ*mCqZRsF#cJ4n_~5G(y^kDDld2qs?|%${pn34OZ_9-)Q^Xz z_Hx|x!;%y~UP*o^#o&o2!k9Fq2>cGua8~JYe6nv2d%|Wp?a#Bpmib6UM)!{SJOi@n zwkVmWWy?OhpM|;l;$%a63VBp!O-@>su>8>{khkLvJ7||qwi;A{lAAWks&!^t{IiKm zY$Dt-E8&LgIHRc64$`B0kR)`?BZ<3JAl#}HcHD6!D^!vgJ^ng2;miSO?f%0|sY~GS zyN6@d=L&K+Qk>immL?moRk3atXP|=GeQ5r;61bH!VM|yhBXxH@p>8kmac&ew zWe9_)eK4uqsf)yO1$H<~@G_dZ!A<)i_5{S@+oT0V?ch_0pxG#0If*EAT)~;+G>PD$ z`52l!i`d(27|SJfV32JB-ncsh9VevXNv#Q{m$Pet-*)Vu%}Rb=i(es5QaQ#x%{>gE z=rx|JdMn~)BFujwq3d{4!#$hG0K>NE% zrvx%Gmjf|)aSQiwZUD~Tm;>ULGaBwhT_7*)l4xGE(5tS78pC`Jm$avP@x=b%mc3QTc4gl8T^`UslADqAM6Gd-S041|BCI=UczwKzo)b1y_rt4ajcVDD>WIA#@(WnsQLEe#rU)2 z`S5E9i`q<1e(<2OHg4z@5QcNyZV&_WKWy?3DdM8+3Tb*8ru#IYn~*D|aX;^D=~?{LqG)95R+35{Ns!i0xYSi3U;JwLGv+e3(XVl}ZE@WS-67!u>Ihw1V>@|RsueW7j^pT-!q zV66OSc7YP#qW{Y-2=NFA9HTERTWs|2=>`A9H2C*T1F^p|xOdERt-#m>{uVnc{=e1# za60~(`=I>y#J|XYWO)2@n0Z5jUD}4m^u3Y8|0tosq`-(cAEbNjQh(I`78gwzW b c h w') - return log + return log \ No newline at end of file diff --git a/model/q_sampler.py b/model/q_sampler.py index bf5e056..61ed3bf 100644 --- a/model/q_sampler.py +++ b/model/q_sampler.py @@ -3,13 +3,17 @@ from typing import Optional, Tuple, Dict, List, Callable import torch import numpy as np from tqdm import tqdm +import einops +import os +from PIL import Image -from ..ldm.modules.diffusionmodules.util import make_beta_schedule -from ..model.cond_fn import Guidance +from ldm.modules.diffusionmodules.util import make_beta_schedule from ..utils.align_color import ( wavelet_reconstruction, adaptive_instance_normalization ) +import comfy.utils + # https://github.com/openai/guided-diffusion/blob/main/guided_diffusion/respace.py def space_timesteps(num_timesteps, section_counts): """ @@ -32,7 +36,7 @@ def space_timesteps(num_timesteps, section_counts): """ if isinstance(section_counts, str): if section_counts.startswith("ddim"): - desired_count = int(section_counts[len("ddim") :]) + desired_count = int(section_counts[len("ddim"):]) for i in range(1, num_timesteps): if len(range(0, num_timesteps, i)) == desired_count: return set(range(0, num_timesteps, i)) @@ -80,7 +84,7 @@ def _extract_into_tensor(arr, timesteps, broadcast_shape): except: # to be compatible with mps res = torch.from_numpy(arr.astype(np.float32)).to(device=timesteps.device)[timesteps].float() - + while len(res.shape) < len(broadcast_shape): res = res[..., None] return res.expand(broadcast_shape) @@ -90,15 +94,15 @@ class SpacedSampler: """ Implementation for spaced sampling schedule proposed in IDDPM. This class is designed for sampling ControlLDM. - + https://arxiv.org/pdf/2102.09672.pdf """ - + def __init__( - self, - model: "ControlLDM", - schedule: str="linear", - var_type: str="fixed_small" + self, + model: "ControlLDM", + schedule: str = "linear", + var_type: str = "fixed_small" ) -> "SpacedSampler": self.model = model self.original_num_steps = model.num_timesteps @@ -108,7 +112,7 @@ class SpacedSampler: def make_schedule(self, num_steps: int) -> None: """ Initialize sampling parameters according to `num_steps`. - + Args: num_steps (int): Sampling steps. @@ -123,12 +127,12 @@ class SpacedSampler: ) original_alphas = 1.0 - original_betas original_alphas_cumprod = np.cumprod(original_alphas, axis=0) - + # calcualte betas for spaced sampling # https://github.com/openai/guided-diffusion/blob/main/guided_diffusion/respace.py used_timesteps = space_timesteps(self.original_num_steps, str(num_steps)) print(f"timesteps used in spaced sampler: \n\t{sorted(list(used_timesteps))}") - + betas = [] last_alpha_cumprod = 1.0 for i, alpha_cumprod in enumerate(original_alphas_cumprod): @@ -139,14 +143,14 @@ class SpacedSampler: assert len(betas) == num_steps betas = np.array(betas, dtype=np.float64) self.betas = betas - - self.timesteps = np.array(sorted(list(used_timesteps)), dtype=np.int32) # e.g. [0, 10, 20, ...] + + self.timesteps = np.array(sorted(list(used_timesteps)), dtype=np.int32) # e.g. [0, 10, 20, ...] alphas = 1.0 - betas self.alphas_cumprod = np.cumprod(alphas, axis=0) self.alphas_cumprod_prev = np.append(1.0, self.alphas_cumprod[:-1]) self.alphas_cumprod_next = np.append(self.alphas_cumprod[1:], 0.0) - assert self.alphas_cumprod_prev.shape == (num_steps, ) - + assert self.alphas_cumprod_prev.shape == (num_steps,) + # calculations for diffusion q(x_t | x_{t-1}) and others self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.alphas_cumprod) @@ -156,7 +160,7 @@ class SpacedSampler: # calculations for posterior q(x_{t-1} | x_t, x_0) self.posterior_variance = ( - betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) ) # log calculation clipped because the posterior variance is 0 at the # beginning of the diffusion chain. @@ -164,18 +168,18 @@ class SpacedSampler: np.append(self.posterior_variance[1], self.posterior_variance[1:]) ) self.posterior_mean_coef1 = ( - betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) ) self.posterior_mean_coef2 = ( - (1.0 - self.alphas_cumprod_prev) - * np.sqrt(alphas) - / (1.0 - self.alphas_cumprod) + (1.0 - self.alphas_cumprod_prev) + * np.sqrt(alphas) + / (1.0 - self.alphas_cumprod) ) - + def make_tao_schedule(self, num_steps: int) -> None: """ Initialize sampling parameters according to `num_steps`. - + Args: num_steps (int): Sampling steps. @@ -190,12 +194,12 @@ class SpacedSampler: ) original_alphas = 1.0 - original_betas original_alphas_cumprod = np.cumprod(original_alphas, axis=0) - + # calcualte betas for spaced sampling # https://github.com/openai/guided-diffusion/blob/main/guided_diffusion/respace.py used_timesteps = space_timesteps(self.original_num_steps, str(num_steps)) print(f"timesteps used in spaced sampler: \n\t{sorted(list(used_timesteps))}") - + betas = [] last_alpha_cumprod = 1.0 for i, alpha_cumprod in enumerate(original_alphas_cumprod): @@ -206,14 +210,14 @@ class SpacedSampler: assert len(betas) == num_steps betas = np.array(betas, dtype=np.float64) self.tao_betas = betas - - self.tao_timesteps = np.array(sorted(list(used_timesteps)), dtype=np.int32) # e.g. [0, 10, 20, ...] + + self.tao_timesteps = np.array(sorted(list(used_timesteps)), dtype=np.int32) # e.g. [0, 10, 20, ...] alphas = 1.0 - betas self.tao_alphas_cumprod = np.cumprod(alphas, axis=0) self.tao_alphas_cumprod_prev = np.append(1.0, self.tao_alphas_cumprod[:-1]) self.tao_alphas_cumprod_next = np.append(self.tao_alphas_cumprod[1:], 0.0) - assert self.tao_alphas_cumprod_prev.shape == (num_steps, ) - + assert self.tao_alphas_cumprod_prev.shape == (num_steps,) + # calculations for diffusion q(x_t | x_{t-1}) and others self.tao_sqrt_alphas_cumprod = np.sqrt(self.tao_alphas_cumprod) self.tao_sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.tao_alphas_cumprod) @@ -223,7 +227,7 @@ class SpacedSampler: # calculations for posterior q(x_{t-1} | x_t, x_0) self.tao_posterior_variance = ( - betas * (1.0 - self.tao_alphas_cumprod_prev) / (1.0 - self.tao_alphas_cumprod) + betas * (1.0 - self.tao_alphas_cumprod_prev) / (1.0 - self.tao_alphas_cumprod) ) # log calculation clipped because the posterior variance is 0 at the # beginning of the diffusion chain. @@ -231,19 +235,19 @@ class SpacedSampler: np.append(self.tao_posterior_variance[1], self.tao_posterior_variance[1:]) ) self.tao_posterior_mean_coef1 = ( - betas * np.sqrt(self.tao_alphas_cumprod_prev) / (1.0 - self.tao_alphas_cumprod) + betas * np.sqrt(self.tao_alphas_cumprod_prev) / (1.0 - self.tao_alphas_cumprod) ) self.tao_posterior_mean_coef2 = ( - (1.0 - self.tao_alphas_cumprod_prev) - * np.sqrt(alphas) - / (1.0 - self.tao_alphas_cumprod) + (1.0 - self.tao_alphas_cumprod_prev) + * np.sqrt(alphas) + / (1.0 - self.tao_alphas_cumprod) ) def q_sample( - self, - x_start: torch.Tensor, - t: torch.Tensor, - noise: Optional[torch.Tensor]=None + self, + x_start: torch.Tensor, + t: torch.Tensor, + noise: Optional[torch.Tensor] = None ) -> torch.Tensor: """ Implement the marginal distribution q(x_t|x_0). @@ -261,26 +265,26 @@ class SpacedSampler: noise = torch.randn_like(x_start) assert noise.shape == x_start.shape return ( - _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start - + _extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) - * noise + _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start + + _extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) + * noise ) def q_posterior_mean_variance( - self, - x_start: torch.Tensor, - x_t: torch.Tensor, - t: torch.Tensor + self, + x_start: torch.Tensor, + x_t: torch.Tensor, + t: torch.Tensor ) -> Tuple[torch.Tensor]: """ Implement the posterior distribution q(x_{t-1}|x_t, x_0). - + Args: x_start (torch.Tensor): The predicted images (NCHW) in timestep `t`. x_t (torch.Tensor): The sampled intermediate variables (NCHW) of timestep `t`. - t (torch.Tensor): Timestep (N) of `x_t`. `t` serves as an index to get + t (torch.Tensor): Timestep (N) of `x_t`. `t` serves as an index to get parameters for each timestep. - + Returns: posterior_mean (torch.Tensor): Mean of the posterior distribution. posterior_variance (torch.Tensor): Variance of the posterior distribution. @@ -288,36 +292,36 @@ class SpacedSampler: """ assert x_start.shape == x_t.shape posterior_mean = ( - _extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start - + _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t + _extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start + + _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t ) posterior_variance = _extract_into_tensor(self.posterior_variance, t, x_t.shape) posterior_log_variance_clipped = _extract_into_tensor( self.posterior_log_variance_clipped, t, x_t.shape ) assert ( - posterior_mean.shape[0] - == posterior_variance.shape[0] - == posterior_log_variance_clipped.shape[0] - == x_start.shape[0] + posterior_mean.shape[0] + == posterior_variance.shape[0] + == posterior_log_variance_clipped.shape[0] + == x_start.shape[0] ) return posterior_mean, posterior_variance, posterior_log_variance_clipped - + def q_posterior_tao_mean_variance( - self, - x_start: torch.Tensor, - x_t: torch.Tensor, - t: torch.Tensor + self, + x_start: torch.Tensor, + x_t: torch.Tensor, + t: torch.Tensor ) -> Tuple[torch.Tensor]: """ Implement the posterior distribution q(x_{t-1}|x_t, x_0). - + Args: x_start (torch.Tensor): The predicted images (NCHW) in timestep `t`. x_t (torch.Tensor): The sampled intermediate variables (NCHW) of timestep `t`. - t (torch.Tensor): Timestep (N) of `x_t`. `t` serves as an index to get + t (torch.Tensor): Timestep (N) of `x_t`. `t` serves as an index to get parameters for each timestep. - + Returns: posterior_mean (torch.Tensor): Mean of the posterior distribution. posterior_variance (torch.Tensor): Variance of the posterior distribution. @@ -325,40 +329,40 @@ class SpacedSampler: """ assert x_start.shape == x_t.shape posterior_mean = ( - _extract_into_tensor(self.tao_posterior_mean_coef1, t, x_t.shape) * x_start - + _extract_into_tensor(self.tao_posterior_mean_coef2, t, x_t.shape) * x_t + _extract_into_tensor(self.tao_posterior_mean_coef1, t, x_t.shape) * x_start + + _extract_into_tensor(self.tao_posterior_mean_coef2, t, x_t.shape) * x_t ) posterior_variance = _extract_into_tensor(self.tao_posterior_variance, t, x_t.shape) posterior_log_variance_clipped = _extract_into_tensor( self.tao_posterior_log_variance_clipped, t, x_t.shape ) assert ( - posterior_mean.shape[0] - == posterior_variance.shape[0] - == posterior_log_variance_clipped.shape[0] - == x_start.shape[0] + posterior_mean.shape[0] + == posterior_variance.shape[0] + == posterior_log_variance_clipped.shape[0] + == x_start.shape[0] ) return posterior_mean, posterior_variance, posterior_log_variance_clipped def _predict_xstart_from_eps( - self, - x_t: torch.Tensor, - t: torch.Tensor, - eps: torch.Tensor + self, + x_t: torch.Tensor, + t: torch.Tensor, + eps: torch.Tensor ) -> torch.Tensor: assert x_t.shape == eps.shape return ( - _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - - _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps + _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t + - _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps ) - + def predict_noise( - self, - x: torch.Tensor, - t: torch.Tensor, - cond: Dict[str, torch.Tensor], - cfg_scale: float, - uncond: Optional[Dict[str, torch.Tensor]] + self, + x: torch.Tensor, + t: torch.Tensor, + cond: Dict[str, torch.Tensor], + cfg_scale: float, + uncond: Optional[Dict[str, torch.Tensor]] ) -> torch.Tensor: if uncond is None or cfg_scale == 1.: model_output = self.model.apply_model(x, t, cond) @@ -367,88 +371,23 @@ class SpacedSampler: model_cond = self.model.apply_model(x, t, cond) model_uncond = self.model.apply_model(x, t, uncond) model_output = model_uncond + cfg_scale * (model_cond - model_uncond) - + if self.model.parameterization == "v": e_t = self.model.predict_eps_from_z_and_v(x, t, model_output) else: e_t = model_output return e_t - - def apply_cond_fn( - self, - x: torch.Tensor, - cond: Dict[str, torch.Tensor], - t: torch.Tensor, - index: torch.Tensor, - cond_fn: Guidance, - cfg_scale: float, - uncond: Optional[Dict[str, torch.Tensor]] - ) -> torch.Tensor: - device = x.device - t_now = int(t[0].item()) + 1 - # ----------------- predict noise and x0 ----------------- # - e_t = self.predict_noise( - x, t, cond, cfg_scale, uncond - ) - pred_x0: torch.Tensor = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) - model_mean, _, _ = self.q_posterior_mean_variance( - x_start=pred_x0, x_t=x, t=index - ) - - # apply classifier guidance for multiple times - for _ in range(cond_fn.repeat): - # ----------------- compute gradient for x0 in latent space ----------------- # - target, pred = None, None - if cond_fn.space == "latent": - target = self.model.get_first_stage_encoding( - self.model.encode_first_stage(cond_fn.target.to(device)) - ) - pred = pred_x0 - elif cond_fn.space == "rgb": - # We need to backward gradient to x0 in latent space, so it's required - # to trace the computation graph while decoding the latent. - with torch.enable_grad(): - pred_x0.requires_grad_(True) - target = cond_fn.target.to(device) - pred = self.model.decode_first_stage_with_grad(pred_x0) - else: - raise NotImplementedError(cond_fn.space) - delta_pred = cond_fn(target, pred, t_now) - - # ----------------- apply classifier guidance ----------------- # - if delta_pred is not None: - if cond_fn.space == "rgb": - # compute gradient for pred_x0 - pred.backward(delta_pred) - delta_pred_x0 = pred_x0.grad - # update prex_x0 - pred_x0 += delta_pred_x0 - # our classifier guidance is equivalent to multiply delta_pred_x0 - # by a constant and then add it to model_mean, We set the constant - # to 0.5 - model_mean += 0.5 * delta_pred_x0 - pred_x0.grad.zero_() - else: - delta_pred_x0 = delta_pred - pred_x0 += delta_pred_x0 - model_mean += 0.5 * delta_pred_x0 - else: - # means stop guidance - break - - return model_mean.detach().clone(), pred_x0.detach().clone() - + @torch.no_grad() def p_sample( - self, - x: torch.Tensor, - cond: Dict[str, torch.Tensor], - t: torch.Tensor, - index: torch.Tensor, - cfg_scale: float, - uncond: Optional[Dict[str, torch.Tensor]], - cond_fn: Optional[Guidance] + self, + x: torch.Tensor, + cond: Dict[str, torch.Tensor], + t: torch.Tensor, + index: torch.Tensor, + cfg_scale: float, + uncond: Optional[Dict[str, torch.Tensor]], ) -> torch.Tensor: # variance of posterior distribution q(x_{t-1}|x_t, x_0) model_variance = { @@ -456,23 +395,15 @@ class SpacedSampler: "fixed_small": self.posterior_variance }[self.var_type] model_variance = _extract_into_tensor(model_variance, index, x.shape) - - # mean of posterior distribution q(x_{t-1}|x_t, x_0) - if cond_fn is not None: - # apply classifier guidance - model_mean, pred_x0 = self.apply_cond_fn( - x, cond, t, index, cond_fn, - cfg_scale, uncond - ) - else: - e_t = self.predict_noise( - x, t, cond, cfg_scale, uncond - ) - pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) - model_mean, _, _ = self.q_posterior_mean_variance( - x_start=pred_x0, x_t=x, t=index - ) - + + e_t = self.predict_noise( + x, t, cond, cfg_scale, uncond + ) + pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) + model_mean, _, _ = self.q_posterior_mean_variance( + x_start=pred_x0, x_t=x, t=index + ) + # sample x_t from q(x_{t-1}|x_t, x_0) noise = torch.randn_like(x) nonzero_mask = ( @@ -480,17 +411,16 @@ class SpacedSampler: ) x_prev = model_mean + nonzero_mask * torch.sqrt(model_variance) * noise return x_prev - + @torch.no_grad() def p_sample_x0( - self, - x: torch.Tensor, - cond: Dict[str, torch.Tensor], - t: torch.Tensor, - index: torch.Tensor, - cfg_scale: float, - uncond: Optional[Dict[str, torch.Tensor]], - cond_fn: Optional[Guidance] + self, + x: torch.Tensor, + cond: Dict[str, torch.Tensor], + t: torch.Tensor, + index: torch.Tensor, + cfg_scale: float, + uncond: Optional[Dict[str, torch.Tensor]], ) -> torch.Tensor: # variance of posterior distribution q(x_{t-1}|x_t, x_0) model_variance = { @@ -498,23 +428,16 @@ class SpacedSampler: "fixed_small": self.posterior_variance }[self.var_type] model_variance = _extract_into_tensor(model_variance, index, x.shape) - + # mean of posterior distribution q(x_{t-1}|x_t, x_0) - if cond_fn is not None: - # apply classifier guidance - model_mean, pred_x0 = self.apply_cond_fn( - x, cond, t, index, cond_fn, - cfg_scale, uncond - ) - else: - e_t = self.predict_noise( - x, t, cond, cfg_scale, uncond - ) - pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) - model_mean, _, _ = self.q_posterior_mean_variance( - x_start=pred_x0, x_t=x, t=index - ) - + e_t = self.predict_noise( + x, t, cond, cfg_scale, uncond + ) + pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) + model_mean, _, _ = self.q_posterior_mean_variance( + x_start=pred_x0, x_t=x, t=index + ) + # sample x_t from q(x_{t-1}|x_t, x_0) noise = torch.randn_like(x) nonzero_mask = ( @@ -522,93 +445,111 @@ class SpacedSampler: ) x_prev = model_mean + nonzero_mask * torch.sqrt(model_variance) * noise return x_prev, pred_x0 - - @torch.no_grad() - def p_sample_tao( - self, - x: torch.Tensor, - cond: Dict[str, torch.Tensor], - t: torch.Tensor, - index: torch.Tensor, - t_max: float, - cfg_scale: float, - uncond: Optional[Dict[str, torch.Tensor]], - cond_fn: Optional[Guidance] - ) -> torch.Tensor: - - if cond_fn is not None: - # apply classifier guidance - model_mean, pred_x0 = self.apply_cond_fn( - x, cond, t, index, cond_fn, - cfg_scale, uncond - ) - else: - e_t = self.predict_noise( - x, t, cond, cfg_scale, uncond - ) - pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) - - # sample x_t from q(x_{t-1}|x_t, x_0) - noise = torch.randn_like(x) - tao_index = torch.tensor(torch.round(index * t_max),dtype=torch.int64) - x_prev = self.q_sample(pred_x0, tao_index) - - return x_prev - @torch.no_grad() - def sample_with_mixdiff_ccsr( - self, - tile_size: int, - tile_stride: int, - steps: int, - t_max: float, - t_min: float, - shape: Tuple[int], - cond_img: torch.Tensor, - positive_prompt: str, - negative_prompt: str, - x_T: Optional[torch.Tensor]=None, - cfg_scale: float=1., - cond_fn: Optional[Guidance]=None, - color_fix_type: str="none" + def p_sample_tao( + self, + x: torch.Tensor, + cond: Dict[str, torch.Tensor], + t: torch.Tensor, + index: torch.Tensor, + t_max: float, + cfg_scale: float, + uncond: Optional[Dict[str, torch.Tensor]] + ) -> torch.Tensor: + + e_t = self.predict_noise( + x, t, cond, cfg_scale, uncond + ) + pred_x0 = self._predict_xstart_from_eps(x_t=x, t=index, eps=e_t) + + # sample x_t from q(x_{t-1}|x_t, x_0) + noise = torch.randn_like(x) + tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64) + x_prev = self.q_sample(pred_x0, tao_index) + + return x_prev + + @torch.no_grad() + def sample_with_tile_ccsr( + self, + empty_text_embed: torch.Tensor, + tile_size: int, + tile_stride: int, + steps: int, + t_max: float, + t_min: float, + shape: Tuple[int], + cond_img: torch.Tensor, + positive_prompt: str, + negative_prompt: str, + x_T: Optional[torch.Tensor] = None, + cfg_scale: float = 1., + color_fix_type: str = "none" ) -> torch.Tensor: def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]: hi_list = list(range(0, h - tile_size + 1, tile_stride)) if (h - tile_size) % tile_stride != 0: hi_list.append(h - tile_size) - + wi_list = list(range(0, w - tile_size + 1, tile_stride)) if (w - tile_size) % tile_stride != 0: wi_list.append(w - tile_size) - + coords = [] for hi in hi_list: for wi in wi_list: coords.append((hi, hi + tile_size, wi, wi + tile_size)) return coords - + + def gaussian_weights(tile_width: int, tile_height: int, nbatches: int) -> torch.Tensor: + """Generates a gaussian mask of weights for tile contributions""" + from numpy import pi, exp, sqrt + import numpy as np + + latent_width = tile_width + latent_height = tile_height + + var = 0.01 + midpoint = (latent_width - 1) / 2 # -1 because index goes from 0 to latent_width - 1 + x_probs = [ + exp(-(x - midpoint) * (x - midpoint) / (latent_width * latent_width) / (2 * var)) / sqrt(2 * pi * var) + for x in range(latent_width)] + midpoint = latent_height / 2 + y_probs = [ + exp(-(y - midpoint) * (y - midpoint) / (latent_height * latent_height) / (2 * var)) / sqrt(2 * pi * var) + for y in range(latent_height)] + + weights = np.outer(y_probs, x_probs) + + return torch.tile(torch.tensor(weights, device=next(self.model.parameters()).device), (nbatches, 4, 1, 1)) + # make sampling parameters (e.g. sigmas) self.make_schedule(num_steps=steps) - + device = next(self.model.parameters()).device b, _, h, w = shape if x_T is None: img = torch.randn(shape, dtype=torch.float32, device=device) else: img = x_T - # create buffers for accumulating predicted noise of different diffusion process - noise_buffer = torch.zeros_like(img) - count = torch.zeros(shape, dtype=torch.long, device=device) + # timesteps iterator - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(self.timesteps) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) - + # q_sample for the start ts = torch.full((b,), time_range[0], device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps - 1) + # calculate the weights + tile_weights = gaussian_weights(tile_size // 8, tile_size // 8, 1) + + # create buffers for accumulating predicted noise of different diffusion process + noise_buffer = torch.zeros_like(img) + count = torch.zeros_like(img) + # predict noise for each tile tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8)) for hi, hi_end, wi, wi_end in tiles_iterator: @@ -619,43 +560,42 @@ class SpacedSampler: tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] tile_cond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] + "c_crossattn": [empty_text_embed] } tile_uncond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] + "c_crossattn": [empty_text_embed] } - # TODO: tile_cond_fn - + # predict noise for this tile tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) - + # accumulate noise - noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise - count[:, :, hi:hi_end, wi:wi_end] += 1 - - # average on noise (score) - noise_buffer.div_(count) + noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise * tile_weights + count[:, :, hi:hi_end, wi:wi_end] += tile_weights + + # fuse by tile_weights on noise (score) + noise_buffer /= count pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) - tao_index = torch.tensor(torch.round(index * t_max),dtype=torch.int64) - img = self.q_sample(pred_x0, tao_index) + tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64) + img = self.q_sample(pred_x0, tao_index) + noise_buffer.zero_() count.zero_() - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] - total_steps = len(time_range) - time_range = time_range[-int(round(total_steps*t_max)):] - total_steps_use = len(time_range) - time_range = time_range[:-int(round(total_steps*t_min))] - iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + total_steps = len(time_range) + time_range = time_range[-int(round(total_steps * t_max)):] + total_steps_use = len(time_range) + time_range = time_range[:-int(round(total_steps * t_min))] + iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + pbar = comfy.utils.ProgressBar(total_steps // 3) # sampling loop for i, step in enumerate(iterator): - ts = torch.full((b,), step, device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps_use - i - 1) - + # predict noise for each tile tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8)) for hi, hi_end, wi, wi_end in tiles_iterator: @@ -672,15 +612,169 @@ class SpacedSampler: "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] } - # TODO: tile_cond_fn - # predict noise for this tile tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) - + + # accumulate noise + noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise * tile_weights + count[:, :, hi:hi_end, wi:wi_end] += tile_weights + pbar.update(1) + # average on noise (score) + noise_buffer /= count + # sample previous latent + pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) + mean, _, _ = self.q_posterior_mean_variance( + x_start=pred_x0, x_t=img, t=index + ) + variance = { + "fixed_large": np.append(self.posterior_variance[1], self.betas[1:]), + "fixed_small": self.posterior_variance + }[self.var_type] + variance = _extract_into_tensor(variance, index, noise_buffer.shape) + + nonzero_mask = ( + (index != 0).float().view(-1, *([1] * (len(noise_buffer.shape) - 1))) + ) + img = mean + nonzero_mask * torch.sqrt(variance) * torch.randn_like(mean) + + noise_buffer.zero_() + count.zero_() + + img = pred_x0 + + img_pixel = (self.model.decode_first_stage(img) + 1) / 2 + # apply color correction (borrowed from StableSR) + if color_fix_type == "adain": + img_pixel = adaptive_instance_normalization(img_pixel, cond_img) + elif color_fix_type == "wavelet": + img_pixel = wavelet_reconstruction(img_pixel, cond_img) + else: + assert color_fix_type == "none", f"unexpected color fix type: {color_fix_type}" + return img_pixel + + @torch.no_grad() + def sample_with_mixdiff_ccsr( + self, + empty_text_embed: torch.Tensor, + tile_size: int, + tile_stride: int, + steps: int, + t_max: float, + t_min: float, + shape: Tuple[int], + cond_img: torch.Tensor, + positive_prompt: str, + negative_prompt: str, + x_T: Optional[torch.Tensor] = None, + cfg_scale: float = 1., + color_fix_type: str = "none" + ) -> torch.Tensor: + def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]: + hi_list = list(range(0, h - tile_size + 1, tile_stride)) + if (h - tile_size) % tile_stride != 0: + hi_list.append(h - tile_size) + + wi_list = list(range(0, w - tile_size + 1, tile_stride)) + if (w - tile_size) % tile_stride != 0: + wi_list.append(w - tile_size) + + coords = [] + for hi in hi_list: + for wi in wi_list: + coords.append((hi, hi + tile_size, wi, wi + tile_size)) + return coords + + # make sampling parameters (e.g. sigmas) + self.make_schedule(num_steps=steps) + + device = next(self.model.parameters()).device + b, _, h, w = shape + if x_T is None: + img = torch.randn(shape, dtype=torch.float32, device=device) + else: + img = x_T + # create buffers for accumulating predicted noise of different diffusion process + noise_buffer = torch.zeros_like(img) + count = torch.zeros(shape, dtype=torch.long, device=device) + # timesteps iterator + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + total_steps = len(self.timesteps) + iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + pbar = comfy.utils.ProgressBar(total_steps // 3) + + # q_sample for the start + ts = torch.full((b,), time_range[0], device=device, dtype=torch.long) + index = torch.full_like(ts, fill_value=total_steps - 1) + + # predict noise for each tile + tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8)) + for hi, hi_end, wi, wi_end in tiles_iterator: + tiles_iterator.set_description(f"Process tile with location ({hi} {hi_end}) ({wi} {wi_end})") + # noisy latent of this diffusion process (tile) at this step + tile_img = img[:, :, hi:hi_end, wi:wi_end] + # prepare condition for this tile + tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] + tile_cond = { + "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], + "c_crossattn": [empty_text_embed] + } + tile_uncond = { + "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], + "c_crossattn": [empty_text_embed] + } + # predict noise for this tile + tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) + + # accumulate noise + noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise + count[:, :, hi:hi_end, wi:wi_end] += 1 + pbar.update(1) + # average on noise (score) + noise_buffer.div_(count) + pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) + tao_index = torch.tensor(torch.round(index * t_max), dtype=torch.int64) + img = self.q_sample(pred_x0, tao_index) + + noise_buffer.zero_() + count.zero_() + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + total_steps = len(time_range) + time_range = time_range[-int(round(total_steps * t_max)):] + total_steps_use = len(time_range) + time_range = time_range[:-int(round(total_steps * t_min))] + iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + + # sampling loop + for i, step in enumerate(iterator): + + ts = torch.full((b,), step, device=device, dtype=torch.long) + index = torch.full_like(ts, fill_value=total_steps_use - i - 1) + + # predict noise for each tile + tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8)) + for hi, hi_end, wi, wi_end in tiles_iterator: + tiles_iterator.set_description(f"Process tile with location ({hi} {hi_end}) ({wi} {wi_end})") + # noisy latent of this diffusion process (tile) at this step + tile_img = img[:, :, hi:hi_end, wi:wi_end] + # prepare condition for this tile + tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] + tile_cond = { + "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], + "c_crossattn": [empty_text_embed] + } + tile_uncond = { + "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], + "c_crossattn": [empty_text_embed] + } + + # predict noise for this tile + tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) + # accumulate noise noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise count[:, :, hi:hi_end, wi:wi_end] += 1 - + # average on noise (score) noise_buffer.div_(count) # sample previous latent @@ -693,15 +787,15 @@ class SpacedSampler: "fixed_small": self.posterior_variance }[self.var_type] variance = _extract_into_tensor(variance, index, noise_buffer.shape) - + nonzero_mask = ( (index != 0).float().view(-1, *([1] * (len(noise_buffer.shape) - 1))) ) img = mean + nonzero_mask * torch.sqrt(variance) * torch.randn_like(mean) - + noise_buffer.zero_() count.zero_() - + img = pred_x0 # decode samples of each diffusion process img_buffer = torch.zeros_like(cond_img) @@ -720,44 +814,44 @@ class SpacedSampler: img_buffer[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] += tile_img_pixel count[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] += 1 img_buffer.div_(count) - + return img_buffer - + @torch.no_grad() def sample_with_mixdiff_control( - self, - control_imgs: torch.Tensor, - tile_size: int, - tile_stride: int, - steps: int, - tao_steps: int, - shape: Tuple[int], - cond_img: torch.Tensor, - positive_prompt: str, - negative_prompt: str, - x_T: Optional[torch.Tensor]=None, - cfg_scale: float=1., - cond_fn: Optional[Guidance]=None, - color_fix_type: str="none" + self, + empty_text_embed: torch.Tensor, + control_imgs: torch.Tensor, + tile_size: int, + tile_stride: int, + steps: int, + tao_steps: int, + shape: Tuple[int], + cond_img: torch.Tensor, + positive_prompt: str, + negative_prompt: str, + x_T: Optional[torch.Tensor] = None, + cfg_scale: float = 1., + color_fix_type: str = "none" ) -> torch.Tensor: def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int) -> Tuple[int, int, int, int]: hi_list = list(range(0, h - tile_size + 1, tile_stride)) if (h - tile_size) % tile_stride != 0: hi_list.append(h - tile_size) - + wi_list = list(range(0, w - tile_size + 1, tile_stride)) if (w - tile_size) % tile_stride != 0: wi_list.append(w - tile_size) - + coords = [] for hi in hi_list: for wi in wi_list: coords.append((hi, hi + tile_size, wi, wi + tile_size)) return coords - + # make sampling parameters (e.g. sigmas) self.make_schedule(num_steps=steps) - + device = next(self.model.parameters()).device b, _, h, w = shape if x_T is None: @@ -768,15 +862,15 @@ class SpacedSampler: noise_buffer = torch.zeros_like(img) count = torch.zeros(shape, dtype=torch.long, device=device) # timesteps iterator - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(self.timesteps) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) - + # q_sample for the start ts = torch.full((b,), time_range[0], device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps - 1) - # start point: LR + # start point: LR img = self.q_sample(control_imgs, index) # predict noise for each tile @@ -789,34 +883,32 @@ class SpacedSampler: tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] tile_cond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] + "c_crossattn": [empty_text_embed] } tile_uncond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] + "c_crossattn": [empty_text_embed] } - # TODO: tile_cond_fn - # predict noise for this tile tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) - + # accumulate noise noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise count[:, :, hi:hi_end, wi:wi_end] += 1 - + # average on noise (score) noise_buffer.div_(count) # sample previous latent pred_x0 = self._predict_xstart_from_eps(x_t=img, t=index, eps=noise_buffer) - tao_index = index - index//(tao_steps-1) + tao_index = index - index // (tao_steps - 1) img = self.q_sample(pred_x0, tao_index) - + noise_buffer.zero_() count.zero_() - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(time_range) - time_range = time_range[total_steps//(tao_steps-1):] + time_range = time_range[total_steps // (tao_steps - 1):] total_steps_use = len(time_range) # time_range = time_range[:-total_steps//(tao_steps-1)] iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) @@ -826,7 +918,7 @@ class SpacedSampler: ts = torch.full((b,), step, device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps_use - i - 1) - + # predict noise for each tile tiles_iterator = tqdm(_sliding_windows(h, w, tile_size // 8, tile_stride // 8)) for hi, hi_end, wi, wi_end in tiles_iterator: @@ -837,21 +929,19 @@ class SpacedSampler: tile_cond_img = cond_img[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] tile_cond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] + "c_crossattn": [empty_text_embed] } tile_uncond = { "c_latent": [self.model.apply_condition_encoder(tile_cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] + "c_crossattn": [empty_text_embed] } - # TODO: tile_cond_fn - # predict noise for this tile tile_noise = self.predict_noise(tile_img, ts, tile_cond, cfg_scale, tile_uncond) - + # accumulate noise noise_buffer[:, :, hi:hi_end, wi:wi_end] += tile_noise count[:, :, hi:hi_end, wi:wi_end] += 1 - + # average on noise (score) noise_buffer.div_(count) # sample previous latent @@ -864,15 +954,15 @@ class SpacedSampler: "fixed_small": self.posterior_variance }[self.var_type] variance = _extract_into_tensor(variance, index, noise_buffer.shape) - + nonzero_mask = ( (index != 0).float().view(-1, *([1] * (len(noise_buffer.shape) - 1))) ) img = mean + nonzero_mask * torch.sqrt(variance) * torch.randn_like(mean) - + noise_buffer.zero_() count.zero_() - + img = pred_x0 # decode samples of each diffusion process img_buffer = torch.zeros_like(cond_img) @@ -891,72 +981,73 @@ class SpacedSampler: img_buffer[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] += tile_img_pixel count[:, :, hi * 8:hi_end * 8, wi * 8: wi_end * 8] += 1 img_buffer.div_(count) - + return img_buffer - + @torch.no_grad() def sample_ccsr( - self, - steps: int, - t_max: float, - t_min: float, - shape: Tuple[int], - cond_img: torch.Tensor, - positive_prompt: str, - negative_prompt: str, - x_T: Optional[torch.Tensor]=None, - cfg_scale: float=1., - cond_fn: Optional[Guidance]=None, - color_fix_type: str="none" + self, + empty_text_embed: torch.Tensor, + steps: int, + t_max: float, + t_min: float, + shape: Tuple[int], + cond_img: torch.Tensor, + positive_prompt: str, + negative_prompt: str, + x_T: Optional[torch.Tensor] = None, + cfg_scale: float = 1., + color_fix_type: str = "none" ) -> torch.Tensor: self.make_schedule(num_steps=steps) # self.make_tao_schedule(num_steps=tao_steps) - + device = next(self.model.parameters()).device b = shape[0] if x_T is None: img = torch.randn(shape, device=device) else: img = x_T - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(self.timesteps) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) - + cond = { "c_latent": [self.model.apply_condition_encoder(cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] + "c_crossattn": [empty_text_embed] } uncond = { "c_latent": [self.model.apply_condition_encoder(cond_img)], - "c_crossattn": [self.model.get_learned_conditioning([negative_prompt] * b)] + "c_crossattn": [empty_text_embed] } # q_sample for the start ts = torch.full((b,), time_range[0], device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps - 1) img = self.p_sample_tao( - img, cond, ts, index=index,t_max=t_max, - cfg_scale=cfg_scale, uncond=uncond, - cond_fn=cond_fn - ) - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + img, cond, ts, index=index, t_max=t_max, + cfg_scale=cfg_scale, uncond=uncond + ) + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(time_range) - time_range = time_range[-int(round(total_steps*t_max)):] + time_range = time_range[-int(round(total_steps * t_max)):] total_steps_use = len(time_range) - time_range = time_range[:-int(round(total_steps*t_min))] + time_range = time_range[:-int(round(total_steps * t_min))] iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + pbar = comfy.utils.ProgressBar(total_steps // 3) for i, step in enumerate(iterator): + ts = torch.full((b,), step, device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps_use - i - 1) img, x0 = self.p_sample_x0( img, cond, ts, index=index, - cfg_scale=cfg_scale, uncond=uncond, - cond_fn=cond_fn + cfg_scale=cfg_scale, uncond=uncond ) - + pbar.update(1) + img = x0 img_pixel = (self.model.decode_first_stage(img) + 1) / 2 # apply color correction (borrowed from StableSR) @@ -967,35 +1058,34 @@ class SpacedSampler: else: assert color_fix_type == "none", f"unexpected color fix type: {color_fix_type}" return img_pixel - + @torch.no_grad() def sample_ccsr_stage1( - self, - steps: int, - t_max: float, - shape: Tuple[int], - cond_img: torch.Tensor, - positive_prompt: str, - negative_prompt: str, - x_T: Optional[torch.Tensor]=None, - cfg_scale: float=1., - cond_fn: Optional[Guidance]=None, - color_fix_type: str="none" + self, + steps: int, + t_max: float, + shape: Tuple[int], + cond_img: torch.Tensor, + positive_prompt: str, + negative_prompt: str, + x_T: Optional[torch.Tensor] = None, + cfg_scale: float = 1., + color_fix_type: str = "none" ) -> torch.Tensor: self.make_schedule(num_steps=steps) # self.make_tao_schedule(num_steps=tao_steps) - + device = next(self.model.parameters()).device b = shape[0] if x_T is None: img = torch.randn(shape, device=device) else: img = x_T - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(self.timesteps) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) - + cond = { "c_latent": [self.model.apply_condition_encoder(cond_img)], "c_crossattn": [self.model.get_learned_conditioning([positive_prompt] * b)] @@ -1009,14 +1099,13 @@ class SpacedSampler: ts = torch.full((b,), time_range[0], device=device, dtype=torch.long) index = torch.full_like(ts, fill_value=total_steps - 1) img = self.p_sample_tao( - img, cond, ts, index=index,t_max=t_max, - cfg_scale=cfg_scale, uncond=uncond, - cond_fn=cond_fn - ) - - time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] + img, cond, ts, index=index, t_max=t_max, + cfg_scale=cfg_scale, uncond=uncond + ) + + time_range = np.flip(self.timesteps) # [1000, 950, 900, ...] total_steps = len(time_range) - time_range = time_range[-int(round(total_steps*t_max)):] + time_range = time_range[-int(round(total_steps * t_max)):] total_steps = len(time_range) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) @@ -1025,10 +1114,9 @@ class SpacedSampler: index = torch.full_like(ts, fill_value=total_steps - i - 1) img = self.p_sample( img, cond, ts, index=index, - cfg_scale=cfg_scale, uncond=uncond, - cond_fn=cond_fn + cfg_scale=cfg_scale, uncond=uncond ) - + img_pixel = (self.model.decode_first_stage(img) + 1) / 2 # apply color correction (borrowed from StableSR) if color_fix_type == "adain": diff --git a/nodes.py b/nodes.py index 1c02b43..d008772 100644 --- a/nodes.py +++ b/nodes.py @@ -26,12 +26,21 @@ class CCSR_Upscale: "image": ("IMAGE", ), "resize_method": (s.upscale_methods, {"default": "lanczos"}), "scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 20.0, "step": 0.01}), - "steps": ("INT", {"default": 45, "min": 2, "max": 4096, "step": 1}), + "steps": ("INT", {"default": 45, "min": 3, "max": 4096, "step": 1}), "t_max": ("FLOAT", {"default": 0.6667,"min": 0, "max": 1, "step": 0.01}), "t_min": ("FLOAT", {"default": 0.3333,"min": 0, "max": 1, "step": 0.01}), + "sampling_method": ( + [ + 'ccsr', + 'ccsr_tiled_mixdiff', + 'ccsr_tiled_vae_gaussian_weights', + ], { + "default": 'ccsr_tiled_mixdiff' + }), "tile_size": ("INT", {"default": 512, "min": 1, "max": 4096, "step": 1}), "tile_stride": ("INT", {"default": 256, "min": 1, "max": 4096, "step": 1}), - "tiled": ("BOOLEAN", {"default": False}), + "vae_tile_size_encode": ("INT", {"default": 1024, "min": 2, "max": 4096, "step": 8}), + "vae_tile_size_decode": ("INT", {"default": 1024, "min": 2, "max": 4096, "step": 8}), "color_fix_type": ( [ 'none', @@ -41,6 +50,7 @@ class CCSR_Upscale: "default": 'adain' }), "keep_model_loaded": ("BOOLEAN", {"default": False}), + }, } @@ -52,12 +62,12 @@ class CCSR_Upscale: CATEGORY = "CCSR" @torch.no_grad() - def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tiled,tile_size, tile_stride, color_fix_type, keep_model_loaded): + def process(self, ccsr_model, image, resize_method, scale_by, steps, t_max, t_min, tile_size, tile_stride, color_fix_type, keep_model_loaded, vae_tile_size_encode, vae_tile_size_decode, sampling_method): comfy.model_management.unload_all_models() device = comfy.model_management.get_torch_device() config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml") + empty_text_embed = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), map_location=device) dtype = torch.float16 if comfy.model_management.should_use_fp16() and not comfy.model_management.is_device_mps(device) else torch.float32 - if not hasattr(self, "model") or self.model is None: config = OmegaConf.load(config_path) self.model = instantiate_from_config(config) @@ -69,8 +79,9 @@ class CCSR_Upscale: self.model.to(device, dtype=dtype) sampler = SpacedSampler(self.model, var_type="fixed_small") + batch_size = image.shape[0] image, = ImageScaleBy.upscale(self, image, resize_method, scale_by) - + # Assuming 'image' is a PyTorch tensor with shape [B, H, W, C] and you want to resize it. B, H, W, C = image.shape @@ -88,34 +99,52 @@ class CCSR_Upscale: resized_image = resized_image.to(device) strength = 1.0 self.model.control_scales = [strength] * 13 - cond_fn = None + height, width = resized_image.size(-2), resized_image.size(-1) shape = (1, 4, height // 8, width // 8) - x_T = torch.randn(shape, device=self.model.device, dtype=dtype) + x_T = torch.randn(shape, device=self.model.device, dtype=torch.float32) autocast_condition = dtype == torch.float16 and not comfy.model_management.is_device_mps(device) - with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): - if not tiled: - samples = sampler.sample_ccsr( - steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image, - positive_prompt="", negative_prompt="", x_T=x_T, - cfg_scale=1.0, cond_fn=cond_fn, - color_fix_type=color_fix_type - ) - else: - samples = sampler.sample_with_mixdiff_ccsr( - tile_size=tile_size, tile_stride=tile_stride, - steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image, - positive_prompt="", negative_prompt="", x_T=x_T, - cfg_scale=1.0, cond_fn=cond_fn, - color_fix_type=color_fix_type - ) + out = [] + with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): + for i in range(batch_size): + + + if sampling_method == 'ccsr_tiled_mixdiff': + print("Using tiled mixdiff") + samples = sampler.sample_with_mixdiff_ccsr( + empty_text_embed, tile_size=tile_size, tile_stride=tile_stride, + steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0), + positive_prompt="", negative_prompt="", x_T=x_T, + cfg_scale=1.0, + color_fix_type=color_fix_type + ) + elif sampling_method == 'ccsr_tiled_vae_gaussian_weights': + self.model._init_tiled_vae(encoder_tile_size=vae_tile_size_encode // 8, decoder_tile_size=vae_tile_size_decode // 8) + print("Using gaussian weights") + samples = sampler.sample_with_tile_ccsr( + empty_text_embed, tile_size=tile_size, tile_stride=tile_stride, + steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0), + positive_prompt="", negative_prompt="", x_T=x_T, + cfg_scale=1.0, + color_fix_type=color_fix_type + ) + else: + print("no tiling") + samples = sampler.sample_ccsr( + empty_text_embed, steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image[i].unsqueeze(0), + positive_prompt="", negative_prompt="", x_T=x_T, + cfg_scale=1.0, + color_fix_type=color_fix_type + ) + out.append(samples.squeeze(0)) + original_height, original_width = H, W processed_height = samples.size(2) target_width = int(processed_height * (original_width / original_height)) - - resized_back_image, = ImageScale.upscale(self, samples.permute(0, 2, 3, 1).cpu(), "lanczos", target_width, processed_height, crop="disabled") - + out_stacked = torch.stack(out, dim=0).cpu().to(torch.float32).permute(0, 2, 3, 1) + resized_back_image, = ImageScale.upscale(self, out_stacked, "lanczos", target_width, processed_height, crop="disabled") + if not keep_model_loaded: self.model = None comfy.model_management.soft_empty_cache() diff --git a/utils/tilevae.py b/utils/tilevae.py new file mode 100644 index 0000000..a2173e4 --- /dev/null +++ b/utils/tilevae.py @@ -0,0 +1,719 @@ +''' +# ------------------------------------------------------------------------ +# +# Tiled VAE +# +# Introducing a revolutionary new optimization designed to make +# the VAE work with giant images on limited VRAM! +# Say goodbye to the frustration of OOM and hello to seamless output! +# +# ------------------------------------------------------------------------ +# +# This script is a wild hack that splits the image into tiles, +# encodes each tile separately, and merges the result back together. +# +# Advantages: +# - The VAE can now work with giant images on limited VRAM +# (~10 GB for 8K images!) +# - The merged output is completely seamless without any post-processing. +# +# Drawbacks: +# - NaNs always appear in for 8k images when you use fp16 (half) VAE +# You must use --no-half-vae to disable half VAE for that giant image. +# - The gradient calculation is not compatible with this hack. It +# will break any backward() or torch.autograd.grad() that passes VAE. +# (But you can still use the VAE to generate training data.) +# +# How it works: +# 1. The image is split into tiles, which are then padded with 11/32 pixels' in the decoder/encoder. +# 2. When Fast Mode is disabled: +# 1. The original VAE forward is decomposed into a task queue and a task worker, which starts to process each tile. +# 2. When GroupNorm is needed, it suspends, stores current GroupNorm mean and var, send everything to RAM, and turns to the next tile. +# 3. After all GroupNorm means and vars are summarized, it applies group norm to tiles and continues. +# 4. A zigzag execution order is used to reduce unnecessary data transfer. +# 3. When Fast Mode is enabled: +# 1. The original input is downsampled and passed to a separate task queue. +# 2. Its group norm parameters are recorded and used by all tiles' task queues. +# 3. Each tile is separately processed without any RAM-VRAM data transfer. +# 4. After all tiles are processed, tiles are written to a result buffer and returned. +# Encoder color fix = only estimate GroupNorm before downsampling, i.e., run in a semi-fast mode. +# +# Enjoy! +# +# @Author: LI YI @ Nanyang Technological University - Singapore +# @Date: 2023-03-02 +# @License: CC BY-NC-SA 4.0 +# +# Please give me a star if you like this project! +# +# ------------------------------------------------------------------------- +''' + +import gc +import math +import sys +from time import time +from tqdm import tqdm + +import torch +import torch.version +import torch.nn.functional as F + +cpu = torch.device("cpu") +device = device_interrogate = device_gfpgan = device_esrgan = device_codeformer = torch.device("cuda") +dtype = torch.float16 +dtype_vae = torch.float16 +dtype_unet = torch.float16 +unet_needs_upcast = False + +def torch_gc(): + + if torch.cuda.is_available(): + with torch.cuda.device("cuda"): + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + + if has_mps(): + mac_specific.torch_mps_gc() + +def has_mps() -> bool: + if sys.platform != "darwin": + return False + else: + return mac_specific.has_mps + +def get_optimal_device_name(): + if torch.cuda.is_available(): + return "cuda" + + if has_mps(): + return "mps" + + return "cpu" + + +def get_optimal_device(): + return torch.device(get_optimal_device_name()) + +class NansException(Exception): + pass + +def test_for_nans(x, where): + if not torch.all(torch.isnan(x)).item(): + return + + if where == "unet": + message = "A tensor with all NaNs was produced in Unet." + + elif where == "vae": + message = "A tensor with all NaNs was produced in VAE." + + else: + message = "A tensor with all NaNs was produced." + + message += " Use --disable-nan-check commandline argument to disable this check." + + raise NansException(message) + +def get_rcmd_enc_tsize(): + if torch.cuda.is_available() and device not in ['cpu', cpu]: + total_memory = torch.cuda.get_device_properties(device).total_memory // 2**20 + if total_memory > 16*1000: ENCODER_TILE_SIZE = 3072 + elif total_memory > 12*1000: ENCODER_TILE_SIZE = 2048 + elif total_memory > 8*1000: ENCODER_TILE_SIZE = 1536 + else: ENCODER_TILE_SIZE = 960 + else: ENCODER_TILE_SIZE = 512 + return ENCODER_TILE_SIZE + + +def get_rcmd_dec_tsize(): + if torch.cuda.is_available() and device not in ['cpu', cpu]: + total_memory = torch.cuda.get_device_properties(device).total_memory // 2**20 + if total_memory > 30*1000: DECODER_TILE_SIZE = 256 + elif total_memory > 16*1000: DECODER_TILE_SIZE = 192 + elif total_memory > 12*1000: DECODER_TILE_SIZE = 128 + elif total_memory > 8*1000: DECODER_TILE_SIZE = 96 + else: DECODER_TILE_SIZE = 64 + else: DECODER_TILE_SIZE = 64 + return DECODER_TILE_SIZE + + +def inplace_nonlinearity(x): + # Test: fix for Nans + return F.silu(x, inplace=True) + +def attn_forward(self, h_): + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + # compute attention + b, c, h, w = q.shape + q = q.reshape(b, c, h*w) + q = q.permute(0, 2, 1) # b,hw,c + k = k.reshape(b, c, h*w) # b,c,hw + w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] + w_ = w_ * (int(c)**(-0.5)) + w_ = torch.nn.functional.softmax(w_, dim=2) + + # attend to values + v = v.reshape(b, c, h*w) + w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q) + # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] + h_ = torch.bmm(v, w_) + h_ = h_.reshape(b, c, h, w) + + h_ = self.proj_out(h_) + + return h_ + + +def attn2task(task_queue, net): + task_queue.append(('store_res', lambda x: x)) + task_queue.append(('pre_norm', net.norm)) + task_queue.append(('attn', lambda x, net=net: attn_forward(net, x))) + task_queue.append(['add_res', None]) + + +def resblock2task(queue, block): + """ + Turn a ResNetBlock into a sequence of tasks and append to the task queue + + @param queue: the target task queue + @param block: ResNetBlock + + """ + if block.in_channels != block.out_channels: + if block.use_conv_shortcut: + queue.append(('store_res', block.conv_shortcut)) + else: + queue.append(('store_res', block.nin_shortcut)) + else: + queue.append(('store_res', lambda x: x)) + queue.append(('pre_norm', block.norm1)) + queue.append(('silu', inplace_nonlinearity)) + queue.append(('conv1', block.conv1)) + queue.append(('pre_norm', block.norm2)) + queue.append(('silu', inplace_nonlinearity)) + queue.append(('conv2', block.conv2)) + queue.append(['add_res', None]) + + +def build_sampling(task_queue, net, is_decoder): + """ + Build the sampling part of a task queue + @param task_queue: the target task queue + @param net: the network + @param is_decoder: currently building decoder or encoder + """ + if is_decoder: + resblock2task(task_queue, net.mid.block_1) + attn2task(task_queue, net.mid.attn_1) + resblock2task(task_queue, net.mid.block_2) + resolution_iter = reversed(range(net.num_resolutions)) + block_ids = net.num_res_blocks + 1 + condition = 0 + module = net.up + func_name = 'upsample' + else: + resolution_iter = range(net.num_resolutions) + block_ids = net.num_res_blocks + condition = net.num_resolutions - 1 + module = net.down + func_name = 'downsample' + + for i_level in resolution_iter: + for i_block in range(block_ids): + resblock2task(task_queue, module[i_level].block[i_block]) + if i_level != condition: + task_queue.append((func_name, getattr(module[i_level], func_name))) + + if not is_decoder: + resblock2task(task_queue, net.mid.block_1) + attn2task(task_queue, net.mid.attn_1) + resblock2task(task_queue, net.mid.block_2) + + +def build_task_queue(net, is_decoder): + """ + Build a single task queue for the encoder or decoder + @param net: the VAE decoder or encoder network + @param is_decoder: currently building decoder or encoder + @return: the task queue + """ + task_queue = [] + task_queue.append(('conv_in', net.conv_in)) + + # construct the sampling part of the task queue + # because encoder and decoder share the same architecture, we extract the sampling part + build_sampling(task_queue, net, is_decoder) + + if not is_decoder or not net.give_pre_end: + task_queue.append(('pre_norm', net.norm_out)) + task_queue.append(('silu', inplace_nonlinearity)) + task_queue.append(('conv_out', net.conv_out)) + if is_decoder and net.tanh_out: + task_queue.append(('tanh', torch.tanh)) + + return task_queue + + +def clone_task_queue(task_queue): + """ + Clone a task queue + @param task_queue: the task queue to be cloned + @return: the cloned task queue + """ + return [[item for item in task] for task in task_queue] + + +def get_var_mean(input, num_groups, eps=1e-6): + """ + Get mean and var for group norm + """ + b, c = input.size(0), input.size(1) + channel_in_group = int(c/num_groups) + input_reshaped = input.contiguous().view(1, int(b * num_groups), channel_in_group, *input.size()[2:]) + var, mean = torch.var_mean(input_reshaped, dim=[0, 2, 3, 4], unbiased=False) + return var, mean + + +def custom_group_norm(input, num_groups, mean, var, weight=None, bias=None, eps=1e-6): + """ + Custom group norm with fixed mean and var + + @param input: input tensor + @param num_groups: number of groups. by default, num_groups = 32 + @param mean: mean, must be pre-calculated by get_var_mean + @param var: var, must be pre-calculated by get_var_mean + @param weight: weight, should be fetched from the original group norm + @param bias: bias, should be fetched from the original group norm + @param eps: epsilon, by default, eps = 1e-6 to match the original group norm + + @return: normalized tensor + """ + b, c = input.size(0), input.size(1) + channel_in_group = int(c/num_groups) + input_reshaped = input.contiguous().view( + 1, int(b * num_groups), channel_in_group, *input.size()[2:]) + + out = F.batch_norm(input_reshaped, mean, var, weight=None, bias=None, training=False, momentum=0, eps=eps) + out = out.view(b, c, *input.size()[2:]) + + # post affine transform + if weight is not None: + out *= weight.view(1, -1, 1, 1) + if bias is not None: + out += bias.view(1, -1, 1, 1) + return out + + +def crop_valid_region(x, input_bbox, target_bbox, is_decoder): + """ + Crop the valid region from the tile + @param x: input tile + @param input_bbox: original input bounding box + @param target_bbox: output bounding box + @param scale: scale factor + @return: cropped tile + """ + padded_bbox = [i * 8 if is_decoder else i//8 for i in input_bbox] + margin = [target_bbox[i] - padded_bbox[i] for i in range(4)] + return x[:, :, margin[2]:x.size(2)+margin[3], margin[0]:x.size(3)+margin[1]] + + +# ↓↓↓ https://github.com/Kahsolt/stable-diffusion-webui-vae-tile-infer ↓↓↓ + +def perfcount(fn): + def wrapper(*args, **kwargs): + ts = time() + + if torch.cuda.is_available(): + torch.cuda.reset_peak_memory_stats(device) + torch_gc() + gc.collect() + + ret = fn(*args, **kwargs) + + torch_gc() + gc.collect() + if torch.cuda.is_available(): + vram = torch.cuda.max_memory_allocated(device) / 2**20 + print(f'[Tiled VAE]: Done in {time() - ts:.3f}s, max VRAM alloc {vram:.3f} MB') + else: + print(f'[Tiled VAE]: Done in {time() - ts:.3f}s') + + return ret + return wrapper + +# ↑↑↑ https://github.com/Kahsolt/stable-diffusion-webui-vae-tile-infer ↑↑↑ + + +class GroupNormParam: + + def __init__(self): + self.var_list = [] + self.mean_list = [] + self.pixel_list = [] + self.weight = None + self.bias = None + + def add_tile(self, tile, layer): + var, mean = get_var_mean(tile, 32) + # For giant images, the variance can be larger than max float16 + # In this case we create a copy to float32 + if var.dtype == torch.float16 and var.isinf().any(): + fp32_tile = tile.float() + var, mean = get_var_mean(fp32_tile, 32) + # ============= DEBUG: test for infinite ============= + # if torch.isinf(var).any(): + # print('var: ', var) + # ==================================================== + self.var_list.append(var) + self.mean_list.append(mean) + self.pixel_list.append( + tile.shape[2]*tile.shape[3]) + if hasattr(layer, 'weight'): + self.weight = layer.weight + self.bias = layer.bias + else: + self.weight = None + self.bias = None + + def summary(self): + """ + summarize the mean and var and return a function + that apply group norm on each tile + """ + if len(self.var_list) == 0: return None + + var = torch.vstack(self.var_list) + mean = torch.vstack(self.mean_list) + max_value = max(self.pixel_list) + pixels = torch.tensor(self.pixel_list, dtype=torch.float32, device=device) / max_value + sum_pixels = torch.sum(pixels) + pixels = pixels.unsqueeze(1) / sum_pixels + var = torch.sum(var * pixels, dim=0) + mean = torch.sum(mean * pixels, dim=0) + return lambda x: custom_group_norm(x, 32, mean, var, self.weight, self.bias) + + @staticmethod + def from_tile(tile, norm): + """ + create a function from a single tile without summary + """ + var, mean = get_var_mean(tile, 32) + if var.dtype == torch.float16 and var.isinf().any(): + fp32_tile = tile.float() + var, mean = get_var_mean(fp32_tile, 32) + # if it is a macbook, we need to convert back to float16 + if var.device.type == 'mps': + # clamp to avoid overflow + var = torch.clamp(var, 0, 60000) + var = var.half() + mean = mean.half() + if hasattr(norm, 'weight'): + weight = norm.weight + bias = norm.bias + else: + weight = None + bias = None + + def group_norm_func(x, mean=mean, var=var, weight=weight, bias=bias): + return custom_group_norm(x, 32, mean, var, weight, bias, 1e-6) + return group_norm_func + + +class VAEHook: + + def __init__(self, net, tile_size, is_decoder:bool, fast_decoder:bool, fast_encoder:bool, color_fix:bool, to_gpu:bool=False): + self.net = net # encoder | decoder + self.tile_size = tile_size + self.is_decoder = is_decoder + self.fast_mode = (fast_encoder and not is_decoder) or (fast_decoder and is_decoder) + self.color_fix = color_fix and not is_decoder + self.to_gpu = to_gpu + self.pad = 11 if is_decoder else 32 # FIXME: magic number + + def __call__(self, x): + original_device = next(self.net.parameters()).device + try: + if self.to_gpu: + self.net = self.net.to(get_optimal_device()) + + B, C, H, W = x.shape + if max(H, W) <= self.pad * 2 + self.tile_size: + print("[Tiled VAE]: the input size is tiny and unnecessary to tile.") + return self.net.original_forward(x) + else: + print("Using tiled VAE Decoder") + return self.vae_tile_forward(x) + finally: + self.net = self.net.to(original_device) + + def get_best_tile_size(self, lowerbound, upperbound): + """ + Get the best tile size for GPU memory + """ + divider = 32 + while divider >= 2: + remainer = lowerbound % divider + if remainer == 0: + return lowerbound + candidate = lowerbound - remainer + divider + if candidate <= upperbound: + return candidate + divider //= 2 + return lowerbound + + def split_tiles(self, h, w): + """ + Tool function to split the image into tiles + @param h: height of the image + @param w: width of the image + @return: tile_input_bboxes, tile_output_bboxes + """ + tile_input_bboxes, tile_output_bboxes = [], [] + tile_size = self.tile_size + pad = self.pad + num_height_tiles = math.ceil((h - 2 * pad) / tile_size) + num_width_tiles = math.ceil((w - 2 * pad) / tile_size) + # If any of the numbers are 0, we let it be 1 + # This is to deal with long and thin images + num_height_tiles = max(num_height_tiles, 1) + num_width_tiles = max(num_width_tiles, 1) + + # Suggestions from https://github.com/Kahsolt: auto shrink the tile size + real_tile_height = math.ceil((h - 2 * pad) / num_height_tiles) + real_tile_width = math.ceil((w - 2 * pad) / num_width_tiles) + real_tile_height = self.get_best_tile_size(real_tile_height, tile_size) + real_tile_width = self.get_best_tile_size(real_tile_width, tile_size) + + print(f'[Tiled VAE]: split to {num_height_tiles}x{num_width_tiles} = {num_height_tiles*num_width_tiles} tiles. ' + + f'Optimal tile size {real_tile_width}x{real_tile_height}, original tile size {tile_size}x{tile_size}') + + for i in range(num_height_tiles): + for j in range(num_width_tiles): + # bbox: [x1, x2, y1, y2] + # the padding is is unnessary for image borders. So we directly start from (32, 32) + input_bbox = [ + pad + j * real_tile_width, + min(pad + (j + 1) * real_tile_width, w), + pad + i * real_tile_height, + min(pad + (i + 1) * real_tile_height, h), + ] + + # if the output bbox is close to the image boundary, we extend it to the image boundary + output_bbox = [ + input_bbox[0] if input_bbox[0] > pad else 0, + input_bbox[1] if input_bbox[1] < w - pad else w, + input_bbox[2] if input_bbox[2] > pad else 0, + input_bbox[3] if input_bbox[3] < h - pad else h, + ] + + # scale to get the final output bbox + output_bbox = [x * 8 if self.is_decoder else x // 8 for x in output_bbox] + tile_output_bboxes.append(output_bbox) + + # indistinguishable expand the input bbox by pad pixels + tile_input_bboxes.append([ + max(0, input_bbox[0] - pad), + min(w, input_bbox[1] + pad), + max(0, input_bbox[2] - pad), + min(h, input_bbox[3] + pad), + ]) + + return tile_input_bboxes, tile_output_bboxes + + @torch.no_grad() + def estimate_group_norm(self, z, task_queue, color_fix): + device = z.device + tile = z + last_id = len(task_queue) - 1 + while last_id >= 0 and task_queue[last_id][0] != 'pre_norm': + last_id -= 1 + if last_id <= 0 or task_queue[last_id][0] != 'pre_norm': + raise ValueError('No group norm found in the task queue') + # estimate until the last group norm + for i in range(last_id + 1): + task = task_queue[i] + if task[0] == 'pre_norm': + group_norm_func = GroupNormParam.from_tile(tile, task[1]) + task_queue[i] = ('apply_norm', group_norm_func) + if i == last_id: + return True + tile = group_norm_func(tile) + elif task[0] == 'store_res': + task_id = i + 1 + while task_id < last_id and task_queue[task_id][0] != 'add_res': + task_id += 1 + if task_id >= last_id: + continue + task_queue[task_id][1] = task[1](tile) + elif task[0] == 'add_res': + tile += task[1].to(device) + task[1] = None + elif color_fix and task[0] == 'downsample': + for j in range(i, last_id + 1): + if task_queue[j][0] == 'store_res': + task_queue[j] = ('store_res_cpu', task_queue[j][1]) + return True + else: + tile = task[1](tile) + try: + test_for_nans(tile, "vae") + except: + print(f'Nan detected in fast mode estimation. Fast mode disabled.') + return False + + raise IndexError('Should not reach here') + + @perfcount + @torch.no_grad() + def vae_tile_forward(self, z): + """ + Decode a latent vector z into an image in a tiled manner. + @param z: latent vector + @return: image + """ + device = next(self.net.parameters()).device + net = self.net + tile_size = self.tile_size + is_decoder = self.is_decoder + + z = z.detach() # detach the input to avoid backprop + + N, height, width = z.shape[0], z.shape[2], z.shape[3] + net.last_z_shape = z.shape + + # Split the input into tiles and build a task queue for each tile + print(f'[Tiled VAE]: input_size: {z.shape}, tile_size: {tile_size}, padding: {self.pad}') + + in_bboxes, out_bboxes = self.split_tiles(height, width) + + # Prepare tiles by split the input latents + tiles = [] + for input_bbox in in_bboxes: + tile = z[:, :, input_bbox[2]:input_bbox[3], input_bbox[0]:input_bbox[1]].cpu() + tiles.append(tile) + + num_tiles = len(tiles) + num_completed = 0 + + # Build task queues + single_task_queue = build_task_queue(net, is_decoder) + if self.fast_mode: + # Fast mode: downsample the input image to the tile size, + # then estimate the group norm parameters on the downsampled image + scale_factor = tile_size / max(height, width) + z = z.to(device) + downsampled_z = F.interpolate(z, scale_factor=scale_factor, mode='nearest-exact') + # use nearest-exact to keep statictics as close as possible + print(f'[Tiled VAE]: Fast mode enabled, estimating group norm parameters on {downsampled_z.shape[3]} x {downsampled_z.shape[2]} image') + + # ======= Special thanks to @Kahsolt for distribution shift issue ======= # + # The downsampling will heavily distort its mean and std, so we need to recover it. + std_old, mean_old = torch.std_mean(z, dim=[0, 2, 3], keepdim=True) + std_new, mean_new = torch.std_mean(downsampled_z, dim=[0, 2, 3], keepdim=True) + downsampled_z = (downsampled_z - mean_new) / std_new * std_old + mean_old + del std_old, mean_old, std_new, mean_new + # occasionally the std_new is too small or too large, which exceeds the range of float16 + # so we need to clamp it to max z's range. + downsampled_z = torch.clamp_(downsampled_z, min=z.min(), max=z.max()) + estimate_task_queue = clone_task_queue(single_task_queue) + if self.estimate_group_norm(downsampled_z, estimate_task_queue, color_fix=self.color_fix): + single_task_queue = estimate_task_queue + del downsampled_z + + task_queues = [clone_task_queue(single_task_queue) for _ in range(num_tiles)] + + # Dummy result + result = None + result_approx = None + # try: + # with devices.autocast(): + # result_approx = torch.cat([F.interpolate(cheap_approximation(x).unsqueeze(0), scale_factor=opt_f, mode='nearest-exact') for x in z], dim=0).cpu() + # except: pass + # Free memory of input latent tensor + del z + + # Task queue execution + pbar = tqdm(total=num_tiles * len(task_queues[0]), desc=f"[Tiled VAE]: Executing {'Decoder' if is_decoder else 'Encoder'} Task Queue: ") + + # execute the task back and forth when switch tiles so that we always + # keep one tile on the GPU to reduce unnecessary data transfer + forward = True + interrupted = False + #state.interrupted = interrupted + while True: + # if state.interrupted: interrupted = True ; break + + group_norm_param = GroupNormParam() + for i in range(num_tiles) if forward else reversed(range(num_tiles)): + # if state.interrupted: interrupted = True ; break + + tile = tiles[i].to(device) + input_bbox = in_bboxes[i] + task_queue = task_queues[i] + + interrupted = False + while len(task_queue) > 0: + # if state.interrupted: interrupted = True ; break + + # DEBUG: current task + # print('Running task: ', task_queue[0][0], ' on tile ', i, '/', num_tiles, ' with shape ', tile.shape) + task = task_queue.pop(0) + if task[0] == 'pre_norm': + group_norm_param.add_tile(tile, task[1]) + break + elif task[0] == 'store_res' or task[0] == 'store_res_cpu': + task_id = 0 + res = task[1](tile) + if not self.fast_mode or task[0] == 'store_res_cpu': + res = res.cpu() + while task_queue[task_id][0] != 'add_res': + task_id += 1 + task_queue[task_id][1] = res + elif task[0] == 'add_res': + tile += task[1].to(device) + task[1] = None + else: + tile = task[1](tile) + pbar.update(1) + + if interrupted: break + + # check for NaNs in the tile. + # If there are NaNs, we abort the process to save user's time + # devices.test_for_nans(tile, "vae") + + if len(task_queue) == 0: + tiles[i] = None + num_completed += 1 + if result is None: # NOTE: dim C varies from different cases, can only be inited dynamically + result = torch.zeros((N, tile.shape[1], height * 8 if is_decoder else height // 8, width * 8 if is_decoder else width // 8), device=device, requires_grad=False) + result[:, :, out_bboxes[i][2]:out_bboxes[i][3], out_bboxes[i][0]:out_bboxes[i][1]] = crop_valid_region(tile, in_bboxes[i], out_bboxes[i], is_decoder) + del tile + elif i == num_tiles - 1 and forward: + forward = False + tiles[i] = tile + elif i == 0 and not forward: + forward = True + tiles[i] = tile + else: + tiles[i] = tile.cpu() + del tile + + if interrupted: break + if num_completed == num_tiles: break + + # insert the group norm task to the head of each task queue + group_norm_func = group_norm_param.summary() + if group_norm_func is not None: + for i in range(num_tiles): + task_queue = task_queues[i] + task_queue.insert(0, ('apply_norm', group_norm_func)) + + # Done! + pbar.close() + return result if result is not None else result_approx.to(device)