modelscope/DiffSynth-Painter
14
1import torch2from .sd3_dit import TimestepEmbeddings, AdaLayerNorm3from einops import rearrange4from .tiler import TileWorker5 6 7 8class RoPEEmbedding(torch.nn.Module):9 def __init__(self, dim, theta, axes_dim):10 super().__init__()11 self.dim = dim12 self.theta = theta13 self.axes_dim = axes_dim14 15 16 def rope(self, pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:17 assert dim % 2 == 0, "The dimension must be even."18 19 scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim20 omega = 1.0 / (theta**scale)21 22 batch_size, seq_length = pos.shape23 out = torch.einsum("...n,d->...nd", pos, omega)24 cos_out = torch.cos(out)25 sin_out = torch.sin(out)26 27 stacked_out = torch.stack([cos_out, -sin_out, sin_out, cos_out], dim=-1)28 out = stacked_out.view(batch_size, -1, dim // 2, 2, 2)29 return out.float()30 31 32 def forward(self, ids):33 n_axes = ids.shape[-1]34 emb = torch.cat([self.rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)], dim=-3)35 return emb.unsqueeze(1)36 37 38 39class RMSNorm(torch.nn.Module):40 def __init__(self, dim, eps):41 super().__init__()42 self.weight = torch.nn.Parameter(torch.ones((dim,)))43 self.eps = eps44 45 def forward(self, hidden_states):46 input_dtype = hidden_states.dtype47 variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)48 hidden_states = hidden_states * torch.rsqrt(variance + self.eps)49 hidden_states = hidden_states.to(input_dtype) * self.weight50 return hidden_states51 52 53 54class FluxJointAttention(torch.nn.Module):55 def __init__(self, dim_a, dim_b, num_heads, head_dim, only_out_a=False):56 super().__init__()57 self.num_heads = num_heads58 self.head_dim = head_dim59 self.only_out_a = only_out_a60 61 self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)62 self.b_to_qkv = torch.nn.Linear(dim_b, dim_b * 3)63 64 self.norm_q_a = RMSNorm(head_dim, eps=1e-6)65 self.norm_k_a = RMSNorm(head_dim, eps=1e-6)66 self.norm_q_b = RMSNorm(head_dim, eps=1e-6)67 self.norm_k_b = RMSNorm(head_dim, eps=1e-6)68 69 self.a_to_out = torch.nn.Linear(dim_a, dim_a)70 if not only_out_a:71 self.b_to_out = torch.nn.Linear(dim_b, dim_b)72 73 74 def apply_rope(self, xq, xk, freqs_cis):75 xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)76 xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)77 xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]78 xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]79 return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)80 81 82 def forward(self, hidden_states_a, hidden_states_b, image_rotary_emb):83 batch_size = hidden_states_a.shape[0]84 85 # Part A86 qkv_a = self.a_to_qkv(hidden_states_a)87 qkv_a = qkv_a.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)88 q_a, k_a, v_a = qkv_a.chunk(3, dim=1)89 q_a, k_a = self.norm_q_a(q_a), self.norm_k_a(k_a)90 91 # Part B92 qkv_b = self.b_to_qkv(hidden_states_b)93 qkv_b = qkv_b.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)94 q_b, k_b, v_b = qkv_b.chunk(3, dim=1)95 q_b, k_b = self.norm_q_b(q_b), self.norm_k_b(k_b)96 97 q = torch.concat([q_b, q_a], dim=2)98 k = torch.concat([k_b, k_a], dim=2)99 v = torch.concat([v_b, v_a], dim=2)100 101 q, k = self.apply_rope(q, k, image_rotary_emb)102 103 hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)104 hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)105 hidden_states = hidden_states.to(q.dtype)106 hidden_states_b, hidden_states_a = hidden_states[:, :hidden_states_b.shape[1]], hidden_states[:, hidden_states_b.shape[1]:]107 hidden_states_a = self.a_to_out(hidden_states_a)108 if self.only_out_a:109 return hidden_states_a110 else:111 hidden_states_b = self.b_to_out(hidden_states_b)112 return hidden_states_a, hidden_states_b113 114 115 116class FluxJointTransformerBlock(torch.nn.Module):117 def __init__(self, dim, num_attention_heads):118 super().__init__()119 self.norm1_a = AdaLayerNorm(dim)120 self.norm1_b = AdaLayerNorm(dim)121 122 self.attn = FluxJointAttention(dim, dim, num_attention_heads, dim // num_attention_heads)123 124 self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)125 self.ff_a = torch.nn.Sequential(126 torch.nn.Linear(dim, dim*4),127 torch.nn.GELU(approximate="tanh"),128 torch.nn.Linear(dim*4, dim)129 )130 131 self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)132 self.ff_b = torch.nn.Sequential(133 torch.nn.Linear(dim, dim*4),134 torch.nn.GELU(approximate="tanh"),135 torch.nn.Linear(dim*4, dim)136 )137 138 139 def forward(self, hidden_states_a, hidden_states_b, temb, image_rotary_emb):140 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)141 norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)142 143 # Attention144 attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b, image_rotary_emb)145 146 # Part A147 hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a148 norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a149 hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)150 151 # Part B152 hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b153 norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b154 hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)155 156 return hidden_states_a, hidden_states_b157 158 159 160class FluxSingleAttention(torch.nn.Module):161 def __init__(self, dim_a, dim_b, num_heads, head_dim):162 super().__init__()163 self.num_heads = num_heads164 self.head_dim = head_dim165 166 self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)167 168 self.norm_q_a = RMSNorm(head_dim, eps=1e-6)169 self.norm_k_a = RMSNorm(head_dim, eps=1e-6)170 171 172 def apply_rope(self, xq, xk, freqs_cis):173 xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)174 xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)175 xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]176 xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]177 return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)178 179 180 def forward(self, hidden_states, image_rotary_emb):181 batch_size = hidden_states.shape[0]182 183 qkv_a = self.a_to_qkv(hidden_states)184 qkv_a = qkv_a.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)185 q_a, k_a, v = qkv_a.chunk(3, dim=1)186 q_a, k_a = self.norm_q_a(q_a), self.norm_k_a(k_a)187 188 q, k = self.apply_rope(q_a, k_a, image_rotary_emb)189 190 hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)191 hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)192 hidden_states = hidden_states.to(q.dtype)193 return hidden_states194 195 196 197class AdaLayerNormSingle(torch.nn.Module):198 def __init__(self, dim):199 super().__init__()200 self.silu = torch.nn.SiLU()201 self.linear = torch.nn.Linear(dim, 3 * dim, bias=True)202 self.norm = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)203 204 205 def forward(self, x, emb):206 emb = self.linear(self.silu(emb))207 shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1)208 x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]209 return x, gate_msa210 211 212 213class FluxSingleTransformerBlock(torch.nn.Module):214 def __init__(self, dim, num_attention_heads):215 super().__init__()216 self.num_heads = num_attention_heads217 self.head_dim = dim // num_attention_heads218 self.dim = dim219 220 self.norm = AdaLayerNormSingle(dim)221 # self.proj_in = torch.nn.Sequential(torch.nn.Linear(dim, dim * 4), torch.nn.GELU(approximate="tanh"))222 # self.attn = FluxSingleAttention(dim, dim, num_attention_heads, dim // num_attention_heads)223 self.linear = torch.nn.Linear(dim, dim * (3 + 4))224 self.norm_q_a = RMSNorm(self.head_dim, eps=1e-6)225 self.norm_k_a = RMSNorm(self.head_dim, eps=1e-6)226 227 self.proj_out = torch.nn.Linear(dim * 5, dim)228 229 230 def apply_rope(self, xq, xk, freqs_cis):231 xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)232 xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)233 xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]234 xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]235 return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)236 237 238 def process_attention(self, hidden_states, image_rotary_emb):239 batch_size = hidden_states.shape[0]240 241 qkv = hidden_states.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)242 q, k, v = qkv.chunk(3, dim=1)243 q, k = self.norm_q_a(q), self.norm_k_a(k)244 245 q, k = self.apply_rope(q, k, image_rotary_emb)246 247 hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)248 hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)249 hidden_states = hidden_states.to(q.dtype)250 return hidden_states251 252 253 def forward(self, hidden_states_a, hidden_states_b, temb, image_rotary_emb):254 residual = hidden_states_a255 norm_hidden_states, gate = self.norm(hidden_states_a, emb=temb)256 hidden_states_a = self.linear(norm_hidden_states)257 attn_output, mlp_hidden_states = hidden_states_a[:, :, :self.dim * 3], hidden_states_a[:, :, self.dim * 3:]258 259 attn_output = self.process_attention(attn_output, image_rotary_emb)260 mlp_hidden_states = torch.nn.functional.gelu(mlp_hidden_states, approximate="tanh")261 262 hidden_states_a = torch.cat([attn_output, mlp_hidden_states], dim=2)263 hidden_states_a = gate.unsqueeze(1) * self.proj_out(hidden_states_a)264 hidden_states_a = residual + hidden_states_a265 266 return hidden_states_a, hidden_states_b267 268 269 270class AdaLayerNormContinuous(torch.nn.Module):271 def __init__(self, dim):272 super().__init__()273 self.silu = torch.nn.SiLU()274 self.linear = torch.nn.Linear(dim, dim * 2, bias=True)275 self.norm = torch.nn.LayerNorm(dim, eps=1e-6, elementwise_affine=False)276 277 def forward(self, x, conditioning):278 emb = self.linear(self.silu(conditioning))279 scale, shift = torch.chunk(emb, 2, dim=1)280 x = self.norm(x) * (1 + scale)[:, None] + shift[:, None]281 return x282 283 284 285class FluxDiT(torch.nn.Module):286 def __init__(self):287 super().__init__()288 self.pos_embedder = RoPEEmbedding(3072, 10000, [16, 56, 56])289 self.time_embedder = TimestepEmbeddings(256, 3072)290 self.guidance_embedder = TimestepEmbeddings(256, 3072)291 self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(768, 3072), torch.nn.SiLU(), torch.nn.Linear(3072, 3072))292 self.context_embedder = torch.nn.Linear(4096, 3072)293 self.x_embedder = torch.nn.Linear(64, 3072)294 295 self.blocks = torch.nn.ModuleList([FluxJointTransformerBlock(3072, 24) for _ in range(19)])296 self.single_blocks = torch.nn.ModuleList([FluxSingleTransformerBlock(3072, 24) for _ in range(38)])297 298 self.norm_out = AdaLayerNormContinuous(3072)299 self.proj_out = torch.nn.Linear(3072, 64)300 301 302 def patchify(self, hidden_states):303 hidden_states = rearrange(hidden_states, "B C (H P) (W Q) -> B (H W) (C P Q)", P=2, Q=2)304 return hidden_states305 306 307 def unpatchify(self, hidden_states, height, width):308 hidden_states = rearrange(hidden_states, "B (H W) (C P Q) -> B C (H P) (W Q)", P=2, Q=2, H=height//2, W=width//2)309 return hidden_states310 311 312 def prepare_image_ids(self, latents):313 batch_size, _, height, width = latents.shape314 latent_image_ids = torch.zeros(height // 2, width // 2, 3)315 latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height // 2)[:, None]316 latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width // 2)[None, :]317 318 latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape319 320 latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1)321 latent_image_ids = latent_image_ids.reshape(322 batch_size, latent_image_id_height * latent_image_id_width, latent_image_id_channels323 )324 latent_image_ids = latent_image_ids.to(device=latents.device, dtype=latents.dtype)325 326 return latent_image_ids327 328 329 def tiled_forward(330 self,331 hidden_states,332 timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids,333 tile_size=128, tile_stride=64,334 **kwargs335 ):336 # Due to the global positional embedding, we cannot implement layer-wise tiled forward.337 hidden_states = TileWorker().tiled_forward(338 lambda x: self.forward(x, timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None),339 hidden_states,340 tile_size,341 tile_stride,342 tile_device=hidden_states.device,343 tile_dtype=hidden_states.dtype344 )345 return hidden_states346 347 348 def forward(349 self,350 hidden_states,351 timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None,352 tiled=False, tile_size=128, tile_stride=64,353 **kwargs354 ):355 if tiled:356 return self.tiled_forward(357 hidden_states,358 timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids,359 tile_size=tile_size, tile_stride=tile_stride,360 **kwargs361 )362 363 if image_ids is None:364 image_ids = self.prepare_image_ids(hidden_states)365 366 conditioning = self.time_embedder(timestep, hidden_states.dtype)\367 + self.guidance_embedder(guidance, hidden_states.dtype)\368 + self.pooled_text_embedder(pooled_prompt_emb)369 prompt_emb = self.context_embedder(prompt_emb)370 image_rotary_emb = self.pos_embedder(torch.cat((text_ids, image_ids), dim=1))371 372 height, width = hidden_states.shape[-2:]373 hidden_states = self.patchify(hidden_states)374 hidden_states = self.x_embedder(hidden_states)375 376 for block in self.blocks:377 hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)378 379 hidden_states = torch.cat([prompt_emb, hidden_states], dim=1)380 for block in self.single_blocks:381 hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)382 hidden_states = hidden_states[:, prompt_emb.shape[1]:]383 384 hidden_states = self.norm_out(hidden_states, conditioning)385 hidden_states = self.proj_out(hidden_states)386 hidden_states = self.unpatchify(hidden_states, height, width)387 388 return hidden_states389 390 391 @staticmethod392 def state_dict_converter():393 return FluxDiTStateDictConverter()394 395 396 397class FluxDiTStateDictConverter:398 def __init__(self):399 pass400 401 def from_diffusers(self, state_dict):402 rename_dict = {403 "context_embedder": "context_embedder",404 "x_embedder": "x_embedder",405 "time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",406 "time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",407 "time_text_embed.guidance_embedder.linear_1": "guidance_embedder.timestep_embedder.0",408 "time_text_embed.guidance_embedder.linear_2": "guidance_embedder.timestep_embedder.2",409 "time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",410 "time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",411 "norm_out.linear": "norm_out.linear",412 "proj_out": "proj_out",413 414 "norm1.linear": "norm1_a.linear",415 "norm1_context.linear": "norm1_b.linear",416 "attn.to_q": "attn.a_to_q",417 "attn.to_k": "attn.a_to_k",418 "attn.to_v": "attn.a_to_v",419 "attn.to_out.0": "attn.a_to_out",420 "attn.add_q_proj": "attn.b_to_q",421 "attn.add_k_proj": "attn.b_to_k",422 "attn.add_v_proj": "attn.b_to_v",423 "attn.to_add_out": "attn.b_to_out",424 "ff.net.0.proj": "ff_a.0",425 "ff.net.2": "ff_a.2",426 "ff_context.net.0.proj": "ff_b.0",427 "ff_context.net.2": "ff_b.2",428 "attn.norm_q": "attn.norm_q_a",429 "attn.norm_k": "attn.norm_k_a",430 "attn.norm_added_q": "attn.norm_q_b",431 "attn.norm_added_k": "attn.norm_k_b",432 }433 rename_dict_single = {434 "attn.to_q": "a_to_q",435 "attn.to_k": "a_to_k",436 "attn.to_v": "a_to_v",437 "attn.norm_q": "norm_q_a",438 "attn.norm_k": "norm_k_a",439 "norm.linear": "norm.linear",440 "proj_mlp": "proj_in_besides_attn",441 "proj_out": "proj_out",442 }443 state_dict_ = {}444 for name, param in state_dict.items():445 if name in rename_dict:446 state_dict_[rename_dict[name]] = param447 elif name.endswith(".weight") or name.endswith(".bias"):448 suffix = ".weight" if name.endswith(".weight") else ".bias"449 prefix = name[:-len(suffix)]450 if prefix in rename_dict:451 state_dict_[rename_dict[prefix] + suffix] = param452 elif prefix.startswith("transformer_blocks."):453 names = prefix.split(".")454 names[0] = "blocks"455 middle = ".".join(names[2:])456 if middle in rename_dict:457 name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])458 state_dict_[name_] = param459 elif prefix.startswith("single_transformer_blocks."):460 names = prefix.split(".")461 names[0] = "single_blocks"462 middle = ".".join(names[2:])463 if middle in rename_dict_single:464 name_ = ".".join(names[:2] + [rename_dict_single[middle]] + [suffix[1:]])465 state_dict_[name_] = param466 else:467 print(name)468 else:469 print(name)470 for name in list(state_dict_.keys()):471 if ".proj_in_besides_attn." in name:472 name_ = name.replace(".proj_in_besides_attn.", ".linear.")473 param = torch.concat([474 state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_q.")],475 state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_k.")],476 state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_v.")],477 state_dict_[name],478 ], dim=0)479 state_dict_[name_] = param480 state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_q."))481 state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_k."))482 state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_v."))483 state_dict_.pop(name)484 for name in list(state_dict_.keys()):485 for component in ["a", "b"]:486 if f".{component}_to_q." in name:487 name_ = name.replace(f".{component}_to_q.", f".{component}_to_qkv.")488 param = torch.concat([489 state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_q.")],490 state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_k.")],491 state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_v.")],492 ], dim=0)493 state_dict_[name_] = param494 state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_q."))495 state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_k."))496 state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_v."))497 return state_dict_498 499 def from_civitai(self, state_dict):500 rename_dict = {501 "time_in.in_layer.bias": "time_embedder.timestep_embedder.0.bias",502 "time_in.in_layer.weight": "time_embedder.timestep_embedder.0.weight",503 "time_in.out_layer.bias": "time_embedder.timestep_embedder.2.bias",504 "time_in.out_layer.weight": "time_embedder.timestep_embedder.2.weight",505 "txt_in.bias": "context_embedder.bias",506 "txt_in.weight": "context_embedder.weight",507 "vector_in.in_layer.bias": "pooled_text_embedder.0.bias",508 "vector_in.in_layer.weight": "pooled_text_embedder.0.weight",509 "vector_in.out_layer.bias": "pooled_text_embedder.2.bias",510 "vector_in.out_layer.weight": "pooled_text_embedder.2.weight",511 "final_layer.linear.bias": "proj_out.bias",512 "final_layer.linear.weight": "proj_out.weight",513 "guidance_in.in_layer.bias": "guidance_embedder.timestep_embedder.0.bias",514 "guidance_in.in_layer.weight": "guidance_embedder.timestep_embedder.0.weight",515 "guidance_in.out_layer.bias": "guidance_embedder.timestep_embedder.2.bias",516 "guidance_in.out_layer.weight": "guidance_embedder.timestep_embedder.2.weight",517 "img_in.bias": "x_embedder.bias",518 "img_in.weight": "x_embedder.weight",519 "final_layer.adaLN_modulation.1.weight": "norm_out.linear.weight",520 "final_layer.adaLN_modulation.1.bias": "norm_out.linear.bias",521 }522 suffix_rename_dict = {523 "img_attn.norm.key_norm.scale": "attn.norm_k_a.weight",524 "img_attn.norm.query_norm.scale": "attn.norm_q_a.weight",525 "img_attn.proj.bias": "attn.a_to_out.bias",526 "img_attn.proj.weight": "attn.a_to_out.weight",527 "img_attn.qkv.bias": "attn.a_to_qkv.bias",528 "img_attn.qkv.weight": "attn.a_to_qkv.weight",529 "img_mlp.0.bias": "ff_a.0.bias",530 "img_mlp.0.weight": "ff_a.0.weight",531 "img_mlp.2.bias": "ff_a.2.bias",532 "img_mlp.2.weight": "ff_a.2.weight",533 "img_mod.lin.bias": "norm1_a.linear.bias",534 "img_mod.lin.weight": "norm1_a.linear.weight",535 "txt_attn.norm.key_norm.scale": "attn.norm_k_b.weight",536 "txt_attn.norm.query_norm.scale": "attn.norm_q_b.weight",537 "txt_attn.proj.bias": "attn.b_to_out.bias",538 "txt_attn.proj.weight": "attn.b_to_out.weight",539 "txt_attn.qkv.bias": "attn.b_to_qkv.bias",540 "txt_attn.qkv.weight": "attn.b_to_qkv.weight",541 "txt_mlp.0.bias": "ff_b.0.bias",542 "txt_mlp.0.weight": "ff_b.0.weight",543 "txt_mlp.2.bias": "ff_b.2.bias",544 "txt_mlp.2.weight": "ff_b.2.weight",545 "txt_mod.lin.bias": "norm1_b.linear.bias",546 "txt_mod.lin.weight": "norm1_b.linear.weight",547 548 "linear1.bias": "linear.bias",549 "linear1.weight": "linear.weight",550 "linear2.bias": "proj_out.bias",551 "linear2.weight": "proj_out.weight",552 "modulation.lin.bias": "norm.linear.bias",553 "modulation.lin.weight": "norm.linear.weight",554 "norm.key_norm.scale": "norm_k_a.weight",555 "norm.query_norm.scale": "norm_q_a.weight",556 }557 state_dict_ = {}558 for name, param in state_dict.items():559 names = name.split(".")560 if name in rename_dict:561 rename = rename_dict[name]562 if name.startswith("final_layer.adaLN_modulation.1."):563 param = torch.concat([param[3072:], param[:3072]], dim=0)564 state_dict_[rename] = param565 elif names[0] == "double_blocks":566 rename = f"blocks.{names[1]}." + suffix_rename_dict[".".join(names[2:])]567 state_dict_[rename] = param568 elif names[0] == "single_blocks":569 if ".".join(names[2:]) in suffix_rename_dict:570 rename = f"single_blocks.{names[1]}." + suffix_rename_dict[".".join(names[2:])]571 state_dict_[rename] = param572 else:573 print(name)574 return state_dict_575 