hugging-apps/echo-memory
0
1import torch2from einops import rearrange3from .svd_unet import TemporalTimesteps4from .tiler import TileWorker5 6 7 8class RMSNorm(torch.nn.Module):9 def __init__(self, dim, eps, elementwise_affine=True):10 super().__init__()11 self.eps = eps12 if elementwise_affine:13 self.weight = torch.nn.Parameter(torch.ones((dim,)))14 else:15 self.weight = None16 17 def forward(self, hidden_states):18 input_dtype = hidden_states.dtype19 variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)20 hidden_states = hidden_states * torch.rsqrt(variance + self.eps)21 hidden_states = hidden_states.to(input_dtype)22 if self.weight is not None:23 hidden_states = hidden_states * self.weight24 return hidden_states25 26 27 28class PatchEmbed(torch.nn.Module):29 def __init__(self, patch_size=2, in_channels=16, embed_dim=1536, pos_embed_max_size=192):30 super().__init__()31 self.pos_embed_max_size = pos_embed_max_size32 self.patch_size = patch_size33 34 self.proj = torch.nn.Conv2d(in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size)35 self.pos_embed = torch.nn.Parameter(torch.zeros(1, self.pos_embed_max_size, self.pos_embed_max_size, embed_dim))36 37 def cropped_pos_embed(self, height, width):38 height = height // self.patch_size39 width = width // self.patch_size40 top = (self.pos_embed_max_size - height) // 241 left = (self.pos_embed_max_size - width) // 242 spatial_pos_embed = self.pos_embed[:, top : top + height, left : left + width, :].flatten(1, 2)43 return spatial_pos_embed44 45 def forward(self, latent):46 height, width = latent.shape[-2:]47 latent = self.proj(latent)48 latent = latent.flatten(2).transpose(1, 2)49 pos_embed = self.cropped_pos_embed(height, width)50 return latent + pos_embed51 52 53 54class TimestepEmbeddings(torch.nn.Module):55 def __init__(self, dim_in, dim_out, computation_device=None):56 super().__init__()57 self.time_proj = TemporalTimesteps(num_channels=dim_in, flip_sin_to_cos=True, downscale_freq_shift=0, computation_device=computation_device)58 self.timestep_embedder = torch.nn.Sequential(59 torch.nn.Linear(dim_in, dim_out), torch.nn.SiLU(), torch.nn.Linear(dim_out, dim_out)60 )61 62 def forward(self, timestep, dtype):63 time_emb = self.time_proj(timestep).to(dtype)64 time_emb = self.timestep_embedder(time_emb)65 return time_emb66 67 68 69class AdaLayerNorm(torch.nn.Module):70 def __init__(self, dim, single=False, dual=False):71 super().__init__()72 self.single = single73 self.dual = dual74 self.linear = torch.nn.Linear(dim, dim * [[6, 2][single], 9][dual])75 self.norm = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)76 77 def forward(self, x, emb):78 emb = self.linear(torch.nn.functional.silu(emb))79 if self.single:80 scale, shift = emb.unsqueeze(1).chunk(2, dim=2)81 x = self.norm(x) * (1 + scale) + shift82 return x83 elif self.dual:84 shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp, shift_msa2, scale_msa2, gate_msa2 = emb.unsqueeze(1).chunk(9, dim=2)85 norm_x = self.norm(x)86 x = norm_x * (1 + scale_msa) + shift_msa87 norm_x2 = norm_x * (1 + scale_msa2) + shift_msa288 return x, gate_msa, shift_mlp, scale_mlp, gate_mlp, norm_x2, gate_msa289 else:90 shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.unsqueeze(1).chunk(6, dim=2)91 x = self.norm(x) * (1 + scale_msa) + shift_msa92 return x, gate_msa, shift_mlp, scale_mlp, gate_mlp93 94 95 96class JointAttention(torch.nn.Module):97 def __init__(self, dim_a, dim_b, num_heads, head_dim, only_out_a=False, use_rms_norm=False):98 super().__init__()99 self.num_heads = num_heads100 self.head_dim = head_dim101 self.only_out_a = only_out_a102 103 self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)104 self.b_to_qkv = torch.nn.Linear(dim_b, dim_b * 3)105 106 self.a_to_out = torch.nn.Linear(dim_a, dim_a)107 if not only_out_a:108 self.b_to_out = torch.nn.Linear(dim_b, dim_b)109 110 if use_rms_norm:111 self.norm_q_a = RMSNorm(head_dim, eps=1e-6)112 self.norm_k_a = RMSNorm(head_dim, eps=1e-6)113 self.norm_q_b = RMSNorm(head_dim, eps=1e-6)114 self.norm_k_b = RMSNorm(head_dim, eps=1e-6)115 else:116 self.norm_q_a = None117 self.norm_k_a = None118 self.norm_q_b = None119 self.norm_k_b = None120 121 122 def process_qkv(self, hidden_states, to_qkv, norm_q, norm_k):123 batch_size = hidden_states.shape[0]124 qkv = to_qkv(hidden_states)125 qkv = qkv.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)126 q, k, v = qkv.chunk(3, dim=1)127 if norm_q is not None:128 q = norm_q(q)129 if norm_k is not None:130 k = norm_k(k)131 return q, k, v132 133 134 def forward(self, hidden_states_a, hidden_states_b):135 batch_size = hidden_states_a.shape[0]136 137 qa, ka, va = self.process_qkv(hidden_states_a, self.a_to_qkv, self.norm_q_a, self.norm_k_a)138 qb, kb, vb = self.process_qkv(hidden_states_b, self.b_to_qkv, self.norm_q_b, self.norm_k_b)139 q = torch.concat([qa, qb], dim=2)140 k = torch.concat([ka, kb], dim=2)141 v = torch.concat([va, vb], dim=2)142 143 hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)144 hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)145 hidden_states = hidden_states.to(q.dtype)146 hidden_states_a, hidden_states_b = hidden_states[:, :hidden_states_a.shape[1]], hidden_states[:, hidden_states_a.shape[1]:]147 hidden_states_a = self.a_to_out(hidden_states_a)148 if self.only_out_a:149 return hidden_states_a150 else:151 hidden_states_b = self.b_to_out(hidden_states_b)152 return hidden_states_a, hidden_states_b153 154 155 156class SingleAttention(torch.nn.Module):157 def __init__(self, dim_a, num_heads, head_dim, use_rms_norm=False):158 super().__init__()159 self.num_heads = num_heads160 self.head_dim = head_dim161 162 self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)163 self.a_to_out = torch.nn.Linear(dim_a, dim_a)164 165 if use_rms_norm:166 self.norm_q_a = RMSNorm(head_dim, eps=1e-6)167 self.norm_k_a = RMSNorm(head_dim, eps=1e-6)168 else:169 self.norm_q_a = None170 self.norm_k_a = None171 172 173 def process_qkv(self, hidden_states, to_qkv, norm_q, norm_k):174 batch_size = hidden_states.shape[0]175 qkv = to_qkv(hidden_states)176 qkv = qkv.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)177 q, k, v = qkv.chunk(3, dim=1)178 if norm_q is not None:179 q = norm_q(q)180 if norm_k is not None:181 k = norm_k(k)182 return q, k, v183 184 185 def forward(self, hidden_states_a):186 batch_size = hidden_states_a.shape[0]187 q, k, v = self.process_qkv(hidden_states_a, self.a_to_qkv, self.norm_q_a, self.norm_k_a)188 189 hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)190 hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)191 hidden_states = hidden_states.to(q.dtype)192 hidden_states = self.a_to_out(hidden_states)193 return hidden_states194 195 196 197class DualTransformerBlock(torch.nn.Module):198 def __init__(self, dim, num_attention_heads, use_rms_norm=False):199 super().__init__()200 self.norm1_a = AdaLayerNorm(dim, dual=True)201 self.norm1_b = AdaLayerNorm(dim)202 203 self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)204 self.attn2 = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)205 206 self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)207 self.ff_a = torch.nn.Sequential(208 torch.nn.Linear(dim, dim*4),209 torch.nn.GELU(approximate="tanh"),210 torch.nn.Linear(dim*4, dim)211 )212 213 self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)214 self.ff_b = torch.nn.Sequential(215 torch.nn.Linear(dim, dim*4),216 torch.nn.GELU(approximate="tanh"),217 torch.nn.Linear(dim*4, dim)218 )219 220 221 def forward(self, hidden_states_a, hidden_states_b, temb):222 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a, norm_hidden_states_a_2, gate_msa_a_2 = self.norm1_a(hidden_states_a, emb=temb)223 norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)224 225 # Attention226 attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b)227 228 # Part A229 hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a230 hidden_states_a = hidden_states_a + gate_msa_a_2 * self.attn2(norm_hidden_states_a_2)231 norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a232 hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)233 234 # Part B235 hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b236 norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b237 hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)238 239 return hidden_states_a, hidden_states_b240 241 242 243class JointTransformerBlock(torch.nn.Module):244 def __init__(self, dim, num_attention_heads, use_rms_norm=False, dual=False):245 super().__init__()246 self.norm1_a = AdaLayerNorm(dim, dual=dual)247 self.norm1_b = AdaLayerNorm(dim)248 249 self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)250 if dual:251 self.attn2 = SingleAttention(dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)252 253 self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)254 self.ff_a = torch.nn.Sequential(255 torch.nn.Linear(dim, dim*4),256 torch.nn.GELU(approximate="tanh"),257 torch.nn.Linear(dim*4, dim)258 )259 260 self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)261 self.ff_b = torch.nn.Sequential(262 torch.nn.Linear(dim, dim*4),263 torch.nn.GELU(approximate="tanh"),264 torch.nn.Linear(dim*4, dim)265 )266 267 268 def forward(self, hidden_states_a, hidden_states_b, temb):269 if self.norm1_a.dual:270 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a, norm_hidden_states_a_2, gate_msa_a_2 = self.norm1_a(hidden_states_a, emb=temb)271 else:272 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)273 norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)274 275 # Attention276 attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b)277 278 # Part A279 hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a280 if self.norm1_a.dual:281 hidden_states_a = hidden_states_a + gate_msa_a_2 * self.attn2(norm_hidden_states_a_2)282 norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a283 hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)284 285 # Part B286 hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b287 norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b288 hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)289 290 return hidden_states_a, hidden_states_b291 292 293 294class JointTransformerFinalBlock(torch.nn.Module):295 def __init__(self, dim, num_attention_heads, use_rms_norm=False):296 super().__init__()297 self.norm1_a = AdaLayerNorm(dim)298 self.norm1_b = AdaLayerNorm(dim, single=True)299 300 self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, only_out_a=True, use_rms_norm=use_rms_norm)301 302 self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)303 self.ff_a = torch.nn.Sequential(304 torch.nn.Linear(dim, dim*4),305 torch.nn.GELU(approximate="tanh"),306 torch.nn.Linear(dim*4, dim)307 )308 309 310 def forward(self, hidden_states_a, hidden_states_b, temb):311 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)312 norm_hidden_states_b = self.norm1_b(hidden_states_b, emb=temb)313 314 # Attention315 attn_output_a = self.attn(norm_hidden_states_a, norm_hidden_states_b)316 317 # Part A318 hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a319 norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a320 hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)321 322 return hidden_states_a, hidden_states_b323 324 325 326class SD3DiT(torch.nn.Module):327 def __init__(self, embed_dim=1536, num_layers=24, use_rms_norm=False, num_dual_blocks=0, pos_embed_max_size=192):328 super().__init__()329 self.pos_embedder = PatchEmbed(patch_size=2, in_channels=16, embed_dim=embed_dim, pos_embed_max_size=pos_embed_max_size)330 self.time_embedder = TimestepEmbeddings(256, embed_dim)331 self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(2048, embed_dim), torch.nn.SiLU(), torch.nn.Linear(embed_dim, embed_dim))332 self.context_embedder = torch.nn.Linear(4096, embed_dim)333 self.blocks = torch.nn.ModuleList([JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm, dual=True) for _ in range(num_dual_blocks)]334 + [JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm) for _ in range(num_layers-1-num_dual_blocks)]335 + [JointTransformerFinalBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm)])336 self.norm_out = AdaLayerNorm(embed_dim, single=True)337 self.proj_out = torch.nn.Linear(embed_dim, 64)338 339 def tiled_forward(self, hidden_states, timestep, prompt_emb, pooled_prompt_emb, tile_size=128, tile_stride=64):340 # Due to the global positional embedding, we cannot implement layer-wise tiled forward.341 hidden_states = TileWorker().tiled_forward(342 lambda x: self.forward(x, timestep, prompt_emb, pooled_prompt_emb),343 hidden_states,344 tile_size,345 tile_stride,346 tile_device=hidden_states.device,347 tile_dtype=hidden_states.dtype348 )349 return hidden_states350 351 def forward(self, hidden_states, timestep, prompt_emb, pooled_prompt_emb, tiled=False, tile_size=128, tile_stride=64, use_gradient_checkpointing=False):352 if tiled:353 return self.tiled_forward(hidden_states, timestep, prompt_emb, pooled_prompt_emb, tile_size, tile_stride)354 conditioning = self.time_embedder(timestep, hidden_states.dtype) + self.pooled_text_embedder(pooled_prompt_emb)355 prompt_emb = self.context_embedder(prompt_emb)356 357 height, width = hidden_states.shape[-2:]358 hidden_states = self.pos_embedder(hidden_states)359 360 def create_custom_forward(module):361 def custom_forward(*inputs):362 return module(*inputs)363 return custom_forward364 365 for block in self.blocks:366 if self.training and use_gradient_checkpointing:367 hidden_states, prompt_emb = torch.utils.checkpoint.checkpoint(368 create_custom_forward(block),369 hidden_states, prompt_emb, conditioning,370 use_reentrant=False,371 )372 else:373 hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning)374 375 hidden_states = self.norm_out(hidden_states, conditioning)376 hidden_states = self.proj_out(hidden_states)377 hidden_states = rearrange(hidden_states, "B (H W) (P Q C) -> B C (H P) (W Q)", P=2, Q=2, H=height//2, W=width//2)378 return hidden_states379 380 @staticmethod381 def state_dict_converter():382 return SD3DiTStateDictConverter()383 384 385 386class SD3DiTStateDictConverter:387 def __init__(self):388 pass389 390 def infer_architecture(self, state_dict):391 embed_dim = state_dict["blocks.0.ff_a.0.weight"].shape[1]392 num_layers = 100393 while num_layers > 0 and f"blocks.{num_layers-1}.ff_a.0.bias" not in state_dict:394 num_layers -= 1395 use_rms_norm = "blocks.0.attn.norm_q_a.weight" in state_dict396 num_dual_blocks = 0397 while f"blocks.{num_dual_blocks}.attn2.a_to_out.bias" in state_dict:398 num_dual_blocks += 1399 pos_embed_max_size = state_dict["pos_embedder.pos_embed"].shape[1]400 return {401 "embed_dim": embed_dim,402 "num_layers": num_layers,403 "use_rms_norm": use_rms_norm,404 "num_dual_blocks": num_dual_blocks,405 "pos_embed_max_size": pos_embed_max_size406 }407 408 def from_diffusers(self, state_dict):409 rename_dict = {410 "context_embedder": "context_embedder",411 "pos_embed.pos_embed": "pos_embedder.pos_embed",412 "pos_embed.proj": "pos_embedder.proj",413 "time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",414 "time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",415 "time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",416 "time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",417 "norm_out.linear": "norm_out.linear",418 "proj_out": "proj_out",419 420 "norm1.linear": "norm1_a.linear",421 "norm1_context.linear": "norm1_b.linear",422 "attn.to_q": "attn.a_to_q",423 "attn.to_k": "attn.a_to_k",424 "attn.to_v": "attn.a_to_v",425 "attn.to_out.0": "attn.a_to_out",426 "attn.add_q_proj": "attn.b_to_q",427 "attn.add_k_proj": "attn.b_to_k",428 "attn.add_v_proj": "attn.b_to_v",429 "attn.to_add_out": "attn.b_to_out",430 "ff.net.0.proj": "ff_a.0",431 "ff.net.2": "ff_a.2",432 "ff_context.net.0.proj": "ff_b.0",433 "ff_context.net.2": "ff_b.2",434 435 "attn.norm_q": "attn.norm_q_a",436 "attn.norm_k": "attn.norm_k_a",437 "attn.norm_added_q": "attn.norm_q_b",438 "attn.norm_added_k": "attn.norm_k_b",439 }440 state_dict_ = {}441 for name, param in state_dict.items():442 if name in rename_dict:443 if name == "pos_embed.pos_embed":444 param = param.reshape((1, 192, 192, param.shape[-1]))445 state_dict_[rename_dict[name]] = param446 elif name.endswith(".weight") or name.endswith(".bias"):447 suffix = ".weight" if name.endswith(".weight") else ".bias"448 prefix = name[:-len(suffix)]449 if prefix in rename_dict:450 state_dict_[rename_dict[prefix] + suffix] = param451 elif prefix.startswith("transformer_blocks."):452 names = prefix.split(".")453 names[0] = "blocks"454 middle = ".".join(names[2:])455 if middle in rename_dict:456 name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])457 state_dict_[name_] = param458 merged_keys = [name for name in state_dict_ if ".a_to_q." in name or ".b_to_q." in name]459 for key in merged_keys:460 param = torch.concat([461 state_dict_[key.replace("to_q", "to_q")],462 state_dict_[key.replace("to_q", "to_k")],463 state_dict_[key.replace("to_q", "to_v")],464 ], dim=0)465 name = key.replace("to_q", "to_qkv")466 state_dict_.pop(key.replace("to_q", "to_q"))467 state_dict_.pop(key.replace("to_q", "to_k"))468 state_dict_.pop(key.replace("to_q", "to_v"))469 state_dict_[name] = param470 return state_dict_, self.infer_architecture(state_dict_)471 472 def from_civitai(self, state_dict):473 rename_dict = {474 "model.diffusion_model.context_embedder.bias": "context_embedder.bias",475 "model.diffusion_model.context_embedder.weight": "context_embedder.weight",476 "model.diffusion_model.final_layer.linear.bias": "proj_out.bias",477 "model.diffusion_model.final_layer.linear.weight": "proj_out.weight",478 479 "model.diffusion_model.pos_embed": "pos_embedder.pos_embed",480 "model.diffusion_model.t_embedder.mlp.0.bias": "time_embedder.timestep_embedder.0.bias",481 "model.diffusion_model.t_embedder.mlp.0.weight": "time_embedder.timestep_embedder.0.weight",482 "model.diffusion_model.t_embedder.mlp.2.bias": "time_embedder.timestep_embedder.2.bias",483 "model.diffusion_model.t_embedder.mlp.2.weight": "time_embedder.timestep_embedder.2.weight",484 "model.diffusion_model.x_embedder.proj.bias": "pos_embedder.proj.bias",485 "model.diffusion_model.x_embedder.proj.weight": "pos_embedder.proj.weight",486 "model.diffusion_model.y_embedder.mlp.0.bias": "pooled_text_embedder.0.bias",487 "model.diffusion_model.y_embedder.mlp.0.weight": "pooled_text_embedder.0.weight",488 "model.diffusion_model.y_embedder.mlp.2.bias": "pooled_text_embedder.2.bias",489 "model.diffusion_model.y_embedder.mlp.2.weight": "pooled_text_embedder.2.weight",490 491 "model.diffusion_model.joint_blocks.23.context_block.adaLN_modulation.1.weight": "blocks.23.norm1_b.linear.weight",492 "model.diffusion_model.joint_blocks.23.context_block.adaLN_modulation.1.bias": "blocks.23.norm1_b.linear.bias",493 "model.diffusion_model.final_layer.adaLN_modulation.1.weight": "norm_out.linear.weight",494 "model.diffusion_model.final_layer.adaLN_modulation.1.bias": "norm_out.linear.bias",495 }496 for i in range(40):497 rename_dict.update({498 f"model.diffusion_model.joint_blocks.{i}.context_block.adaLN_modulation.1.bias": f"blocks.{i}.norm1_b.linear.bias",499 f"model.diffusion_model.joint_blocks.{i}.context_block.adaLN_modulation.1.weight": f"blocks.{i}.norm1_b.linear.weight",500 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.proj.bias": f"blocks.{i}.attn.b_to_out.bias",501 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.proj.weight": f"blocks.{i}.attn.b_to_out.weight",502 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.qkv.bias": [f'blocks.{i}.attn.b_to_q.bias', f'blocks.{i}.attn.b_to_k.bias', f'blocks.{i}.attn.b_to_v.bias'],503 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.qkv.weight": [f'blocks.{i}.attn.b_to_q.weight', f'blocks.{i}.attn.b_to_k.weight', f'blocks.{i}.attn.b_to_v.weight'],504 f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc1.bias": f"blocks.{i}.ff_b.0.bias",505 f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc1.weight": f"blocks.{i}.ff_b.0.weight",506 f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc2.bias": f"blocks.{i}.ff_b.2.bias",507 f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc2.weight": f"blocks.{i}.ff_b.2.weight",508 f"model.diffusion_model.joint_blocks.{i}.x_block.adaLN_modulation.1.bias": f"blocks.{i}.norm1_a.linear.bias",509 f"model.diffusion_model.joint_blocks.{i}.x_block.adaLN_modulation.1.weight": f"blocks.{i}.norm1_a.linear.weight",510 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.proj.bias": f"blocks.{i}.attn.a_to_out.bias",511 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.proj.weight": f"blocks.{i}.attn.a_to_out.weight",512 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.qkv.bias": [f'blocks.{i}.attn.a_to_q.bias', f'blocks.{i}.attn.a_to_k.bias', f'blocks.{i}.attn.a_to_v.bias'],513 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.qkv.weight": [f'blocks.{i}.attn.a_to_q.weight', f'blocks.{i}.attn.a_to_k.weight', f'blocks.{i}.attn.a_to_v.weight'],514 f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc1.bias": f"blocks.{i}.ff_a.0.bias",515 f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc1.weight": f"blocks.{i}.ff_a.0.weight",516 f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc2.bias": f"blocks.{i}.ff_a.2.bias",517 f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc2.weight": f"blocks.{i}.ff_a.2.weight",518 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.ln_q.weight": f"blocks.{i}.attn.norm_q_a.weight",519 f"model.diffusion_model.joint_blocks.{i}.x_block.attn.ln_k.weight": f"blocks.{i}.attn.norm_k_a.weight",520 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.ln_q.weight": f"blocks.{i}.attn.norm_q_b.weight",521 f"model.diffusion_model.joint_blocks.{i}.context_block.attn.ln_k.weight": f"blocks.{i}.attn.norm_k_b.weight",522 523 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.ln_q.weight": f"blocks.{i}.attn2.norm_q_a.weight",524 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.ln_k.weight": f"blocks.{i}.attn2.norm_k_a.weight",525 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.qkv.weight": f"blocks.{i}.attn2.a_to_qkv.weight",526 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.qkv.bias": f"blocks.{i}.attn2.a_to_qkv.bias",527 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.proj.weight": f"blocks.{i}.attn2.a_to_out.weight",528 f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.proj.bias": f"blocks.{i}.attn2.a_to_out.bias",529 })530 state_dict_ = {}531 for name in state_dict:532 if name in rename_dict:533 param = state_dict[name]534 if name == "model.diffusion_model.pos_embed":535 pos_embed_max_size = int(param.shape[1] ** 0.5 + 0.4)536 param = param.reshape((1, pos_embed_max_size, pos_embed_max_size, param.shape[-1]))537 if isinstance(rename_dict[name], str):538 state_dict_[rename_dict[name]] = param539 else:540 name_ = rename_dict[name][0].replace(".a_to_q.", ".a_to_qkv.").replace(".b_to_q.", ".b_to_qkv.")541 state_dict_[name_] = param542 extra_kwargs = self.infer_architecture(state_dict_)543 num_layers = extra_kwargs["num_layers"]544 for name in [545 f"blocks.{num_layers-1}.norm1_b.linear.weight", f"blocks.{num_layers-1}.norm1_b.linear.bias", "norm_out.linear.weight", "norm_out.linear.bias",546 ]:547 param = state_dict_[name]548 dim = param.shape[0] // 2549 param = torch.concat([param[dim:], param[:dim]], axis=0)550 state_dict_[name] = param551 return state_dict_, self.infer_architecture(state_dict_)552 