Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
flux_dit.py575 linesDownload Raw Back to models
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