hugging-apps/echo-memory
0
1import torch2from .sd_unet import Timesteps, ResnetBlock, AttentionBlock, PushBlock, DownSampler3from .sdxl_unet import SDXLUNet4from .tiler import TileWorker5from .sd_controlnet import ControlNetConditioningLayer6from collections import OrderedDict7 8 9 10class QuickGELU(torch.nn.Module):11 12 def forward(self, x: torch.Tensor):13 return x * torch.sigmoid(1.702 * x)14 15 16 17class ResidualAttentionBlock(torch.nn.Module):18 19 def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):20 super().__init__()21 22 self.attn = torch.nn.MultiheadAttention(d_model, n_head)23 self.ln_1 = torch.nn.LayerNorm(d_model)24 self.mlp = torch.nn.Sequential(OrderedDict([25 ("c_fc", torch.nn.Linear(d_model, d_model * 4)),26 ("gelu", QuickGELU()),27 ("c_proj", torch.nn.Linear(d_model * 4, d_model))28 ]))29 self.ln_2 = torch.nn.LayerNorm(d_model)30 self.attn_mask = attn_mask31 32 def attention(self, x: torch.Tensor):33 self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None34 return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]35 36 def forward(self, x: torch.Tensor):37 x = x + self.attention(self.ln_1(x))38 x = x + self.mlp(self.ln_2(x))39 return x40 41 42 43class SDXLControlNetUnion(torch.nn.Module):44 def __init__(self, global_pool=False):45 super().__init__()46 self.time_proj = Timesteps(320)47 self.time_embedding = torch.nn.Sequential(48 torch.nn.Linear(320, 1280),49 torch.nn.SiLU(),50 torch.nn.Linear(1280, 1280)51 )52 self.add_time_proj = Timesteps(256)53 self.add_time_embedding = torch.nn.Sequential(54 torch.nn.Linear(2816, 1280),55 torch.nn.SiLU(),56 torch.nn.Linear(1280, 1280)57 )58 self.control_type_proj = Timesteps(256)59 self.control_type_embedding = torch.nn.Sequential(60 torch.nn.Linear(256 * 8, 1280),61 torch.nn.SiLU(),62 torch.nn.Linear(1280, 1280)63 )64 self.conv_in = torch.nn.Conv2d(4, 320, kernel_size=3, padding=1)65 66 self.controlnet_conv_in = ControlNetConditioningLayer(channels=(3, 16, 32, 96, 256, 320))67 self.controlnet_transformer = ResidualAttentionBlock(320, 8)68 self.task_embedding = torch.nn.Parameter(torch.randn(8, 320))69 self.spatial_ch_projs = torch.nn.Linear(320, 320)70 71 self.blocks = torch.nn.ModuleList([72 # DownBlock2D73 ResnetBlock(320, 320, 1280),74 PushBlock(),75 ResnetBlock(320, 320, 1280),76 PushBlock(),77 DownSampler(320),78 PushBlock(),79 # CrossAttnDownBlock2D80 ResnetBlock(320, 640, 1280),81 AttentionBlock(10, 64, 640, 2, 2048),82 PushBlock(),83 ResnetBlock(640, 640, 1280),84 AttentionBlock(10, 64, 640, 2, 2048),85 PushBlock(),86 DownSampler(640),87 PushBlock(),88 # CrossAttnDownBlock2D89 ResnetBlock(640, 1280, 1280),90 AttentionBlock(20, 64, 1280, 10, 2048),91 PushBlock(),92 ResnetBlock(1280, 1280, 1280),93 AttentionBlock(20, 64, 1280, 10, 2048),94 PushBlock(),95 # UNetMidBlock2DCrossAttn96 ResnetBlock(1280, 1280, 1280),97 AttentionBlock(20, 64, 1280, 10, 2048),98 ResnetBlock(1280, 1280, 1280),99 PushBlock()100 ])101 102 self.controlnet_blocks = torch.nn.ModuleList([103 torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),104 torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),105 torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),106 torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),107 torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),108 torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),109 torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),110 torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),111 torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),112 torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),113 ])114 115 self.global_pool = global_pool116 117 # 0 -- openpose118 # 1 -- depth119 # 2 -- hed/pidi/scribble/ted120 # 3 -- canny/lineart/anime_lineart/mlsd121 # 4 -- normal122 # 5 -- segment123 # 6 -- tile124 # 7 -- repaint125 self.task_id = {126 "openpose": 0,127 "depth": 1,128 "softedge": 2,129 "canny": 3,130 "lineart": 3,131 "lineart_anime": 3,132 "tile": 6,133 "inpaint": 7134 }135 136 137 def fuse_condition_to_input(self, hidden_states, task_id, conditioning):138 controlnet_cond = self.controlnet_conv_in(conditioning)139 feat_seq = torch.mean(controlnet_cond, dim=(2, 3))140 feat_seq = feat_seq + self.task_embedding[task_id]141 x = torch.stack([feat_seq, torch.mean(hidden_states, dim=(2, 3))], dim=1)142 x = self.controlnet_transformer(x)143 144 alpha = self.spatial_ch_projs(x[:,0]).unsqueeze(-1).unsqueeze(-1)145 controlnet_cond_fuser = controlnet_cond + alpha146 147 hidden_states = hidden_states + controlnet_cond_fuser148 return hidden_states149 150 151 def forward(152 self,153 sample, timestep, encoder_hidden_states,154 conditioning, processor_id, add_time_id, add_text_embeds,155 tiled=False, tile_size=64, tile_stride=32,156 unet:SDXLUNet=None,157 **kwargs158 ):159 task_id = self.task_id[processor_id]160 161 # 1. time162 t_emb = self.time_proj(timestep).to(sample.dtype)163 t_emb = self.time_embedding(t_emb)164 165 time_embeds = self.add_time_proj(add_time_id)166 time_embeds = time_embeds.reshape((add_text_embeds.shape[0], -1))167 add_embeds = torch.concat([add_text_embeds, time_embeds], dim=-1)168 add_embeds = add_embeds.to(sample.dtype)169 if unet is not None and unet.is_kolors:170 add_embeds = unet.add_time_embedding(add_embeds)171 else:172 add_embeds = self.add_time_embedding(add_embeds)173 174 control_type = torch.zeros((sample.shape[0], 8), dtype=sample.dtype, device=sample.device)175 control_type[:, task_id] = 1176 control_embeds = self.control_type_proj(control_type.flatten())177 control_embeds = control_embeds.reshape((sample.shape[0], -1))178 control_embeds = control_embeds.to(sample.dtype)179 control_embeds = self.control_type_embedding(control_embeds)180 time_emb = t_emb + add_embeds + control_embeds181 182 # 2. pre-process183 height, width = sample.shape[2], sample.shape[3]184 hidden_states = self.conv_in(sample)185 hidden_states = self.fuse_condition_to_input(hidden_states, task_id, conditioning)186 text_emb = encoder_hidden_states187 if unet is not None and unet.is_kolors:188 text_emb = unet.text_intermediate_proj(text_emb)189 res_stack = [hidden_states]190 191 # 3. blocks192 for i, block in enumerate(self.blocks):193 if tiled and not isinstance(block, PushBlock):194 _, _, inter_height, _ = hidden_states.shape195 resize_scale = inter_height / height196 hidden_states = TileWorker().tiled_forward(197 lambda x: block(x, time_emb, text_emb, res_stack)[0],198 hidden_states,199 int(tile_size * resize_scale),200 int(tile_stride * resize_scale),201 tile_device=hidden_states.device,202 tile_dtype=hidden_states.dtype203 )204 else:205 hidden_states, _, _, _ = block(hidden_states, time_emb, text_emb, res_stack)206 207 # 4. ControlNet blocks208 controlnet_res_stack = [block(res) for block, res in zip(self.controlnet_blocks, res_stack)]209 210 # pool211 if self.global_pool:212 controlnet_res_stack = [res.mean(dim=(2, 3), keepdim=True) for res in controlnet_res_stack]213 214 return controlnet_res_stack215 216 @staticmethod217 def state_dict_converter():218 return SDXLControlNetUnionStateDictConverter()219 220 221 222class SDXLControlNetUnionStateDictConverter:223 def __init__(self):224 pass225 226 def from_diffusers(self, state_dict):227 # architecture228 block_types = [229 "ResnetBlock", "PushBlock", "ResnetBlock", "PushBlock", "DownSampler", "PushBlock",230 "ResnetBlock", "AttentionBlock", "PushBlock", "ResnetBlock", "AttentionBlock", "PushBlock", "DownSampler", "PushBlock",231 "ResnetBlock", "AttentionBlock", "PushBlock", "ResnetBlock", "AttentionBlock", "PushBlock",232 "ResnetBlock", "AttentionBlock", "ResnetBlock", "PushBlock"233 ]234 235 # controlnet_rename_dict236 controlnet_rename_dict = {237 "controlnet_cond_embedding.conv_in.weight": "controlnet_conv_in.blocks.0.weight",238 "controlnet_cond_embedding.conv_in.bias": "controlnet_conv_in.blocks.0.bias",239 "controlnet_cond_embedding.blocks.0.weight": "controlnet_conv_in.blocks.2.weight",240 "controlnet_cond_embedding.blocks.0.bias": "controlnet_conv_in.blocks.2.bias",241 "controlnet_cond_embedding.blocks.1.weight": "controlnet_conv_in.blocks.4.weight",242 "controlnet_cond_embedding.blocks.1.bias": "controlnet_conv_in.blocks.4.bias",243 "controlnet_cond_embedding.blocks.2.weight": "controlnet_conv_in.blocks.6.weight",244 "controlnet_cond_embedding.blocks.2.bias": "controlnet_conv_in.blocks.6.bias",245 "controlnet_cond_embedding.blocks.3.weight": "controlnet_conv_in.blocks.8.weight",246 "controlnet_cond_embedding.blocks.3.bias": "controlnet_conv_in.blocks.8.bias",247 "controlnet_cond_embedding.blocks.4.weight": "controlnet_conv_in.blocks.10.weight",248 "controlnet_cond_embedding.blocks.4.bias": "controlnet_conv_in.blocks.10.bias",249 "controlnet_cond_embedding.blocks.5.weight": "controlnet_conv_in.blocks.12.weight",250 "controlnet_cond_embedding.blocks.5.bias": "controlnet_conv_in.blocks.12.bias",251 "controlnet_cond_embedding.conv_out.weight": "controlnet_conv_in.blocks.14.weight",252 "controlnet_cond_embedding.conv_out.bias": "controlnet_conv_in.blocks.14.bias",253 "control_add_embedding.linear_1.weight": "control_type_embedding.0.weight",254 "control_add_embedding.linear_1.bias": "control_type_embedding.0.bias",255 "control_add_embedding.linear_2.weight": "control_type_embedding.2.weight",256 "control_add_embedding.linear_2.bias": "control_type_embedding.2.bias",257 }258 259 # Rename each parameter260 name_list = sorted([name for name in state_dict])261 rename_dict = {}262 block_id = {"ResnetBlock": -1, "AttentionBlock": -1, "DownSampler": -1, "UpSampler": -1}263 last_block_type_with_id = {"ResnetBlock": "", "AttentionBlock": "", "DownSampler": "", "UpSampler": ""}264 for name in name_list:265 names = name.split(".")266 if names[0] in ["conv_in", "conv_norm_out", "conv_out", "task_embedding", "spatial_ch_projs"]:267 pass268 elif name in controlnet_rename_dict:269 names = controlnet_rename_dict[name].split(".")270 elif names[0] == "controlnet_down_blocks":271 names[0] = "controlnet_blocks"272 elif names[0] == "controlnet_mid_block":273 names = ["controlnet_blocks", "9", names[-1]]274 elif names[0] in ["time_embedding", "add_embedding"]:275 if names[0] == "add_embedding":276 names[0] = "add_time_embedding"277 names[1] = {"linear_1": "0", "linear_2": "2"}[names[1]]278 elif names[0] == "control_add_embedding":279 names[0] = "control_type_embedding"280 elif names[0] == "transformer_layes":281 names[0] = "controlnet_transformer"282 names.pop(1)283 elif names[0] in ["down_blocks", "mid_block", "up_blocks"]:284 if names[0] == "mid_block":285 names.insert(1, "0")286 block_type = {"resnets": "ResnetBlock", "attentions": "AttentionBlock", "downsamplers": "DownSampler", "upsamplers": "UpSampler"}[names[2]]287 block_type_with_id = ".".join(names[:4])288 if block_type_with_id != last_block_type_with_id[block_type]:289 block_id[block_type] += 1290 last_block_type_with_id[block_type] = block_type_with_id291 while block_id[block_type] < len(block_types) and block_types[block_id[block_type]] != block_type:292 block_id[block_type] += 1293 block_type_with_id = ".".join(names[:4])294 names = ["blocks", str(block_id[block_type])] + names[4:]295 if "ff" in names:296 ff_index = names.index("ff")297 component = ".".join(names[ff_index:ff_index+3])298 component = {"ff.net.0": "act_fn", "ff.net.2": "ff"}[component]299 names = names[:ff_index] + [component] + names[ff_index+3:]300 if "to_out" in names:301 names.pop(names.index("to_out") + 1)302 else:303 print(name, state_dict[name].shape)304 # raise ValueError(f"Unknown parameters: {name}")305 rename_dict[name] = ".".join(names)306 307 # Convert state_dict308 state_dict_ = {}309 for name, param in state_dict.items():310 if name not in rename_dict:311 continue312 if ".proj_in." in name or ".proj_out." in name:313 param = param.squeeze()314 state_dict_[rename_dict[name]] = param315 return state_dict_316 317 def from_civitai(self, state_dict):318 return self.from_diffusers(state_dict)