diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py new file mode 100644 index 0000000..d2bf810 --- /dev/null +++ b/PixArt/diffusers_convert.py @@ -0,0 +1,95 @@ +# For using the diffusers format weights +# Based on the original ComfyUI function + +# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py +import torch + +conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) + # Patch embeddings + ("x_embedder.proj.weight", "pos_embed.proj.weight"), + ("x_embedder.proj.bias", "pos_embed.proj.bias"), + # Caption projection + ("y_embedder.y_embedding", "caption_projection.y_embedding"), + ("y_embedder.y_proj.fc1.weight", "caption_projection.linear_1.weight"), + ("y_embedder.y_proj.fc1.bias", "caption_projection.linear_1.bias"), + ("y_embedder.y_proj.fc2.weight", "caption_projection.linear_2.weight"), + ("y_embedder.y_proj.fc2.bias", "caption_projection.linear_2.bias"), + # AdaLN-single LN + ("t_embedder.mlp.0.weight", "adaln_single.emb.timestep_embedder.linear_1.weight"), + ("t_embedder.mlp.0.bias", "adaln_single.emb.timestep_embedder.linear_1.bias"), + ("t_embedder.mlp.2.weight", "adaln_single.emb.timestep_embedder.linear_2.weight"), + ("t_embedder.mlp.2.bias", "adaln_single.emb.timestep_embedder.linear_2.bias"), + # Shared norm + ("t_block.1.weight", "adaln_single.linear.weight"), + ("t_block.1.bias", "adaln_single.linear.bias"), + # Final block + ("final_layer.linear.weight", "proj_out.weight"), + ("final_layer.linear.bias", "proj_out.bias"), + ("final_layer.scale_shift_table", "scale_shift_table"), +] + +conversion_map_ms = [ # for multi_scale_train (MS) + # Resolution + ("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"), + ("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"), + ("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"), + ("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"), + # Aspect ratio + ("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"), + ("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"), + ("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"), + ("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"), +] + +# Add actual transformer blocks +for depth in range(28): + # Transformer blocks + conversion_map += [ + (f"blocks.{depth}.scale_shift_table", f"transformer_blocks.{depth}.scale_shift_table"), + # Projection + (f"blocks.{depth}.attn.proj.weight", f"transformer_blocks.{depth}.attn1.to_out.0.weight"), + (f"blocks.{depth}.attn.proj.bias", f"transformer_blocks.{depth}.attn1.to_out.0.bias"), + # Feed-forward + (f"blocks.{depth}.mlp.fc1.weight", f"transformer_blocks.{depth}.ff.net.0.proj.weight"), + (f"blocks.{depth}.mlp.fc1.bias", f"transformer_blocks.{depth}.ff.net.0.proj.bias"), + (f"blocks.{depth}.mlp.fc2.weight", f"transformer_blocks.{depth}.ff.net.2.weight"), + (f"blocks.{depth}.mlp.fc2.bias", f"transformer_blocks.{depth}.ff.net.2.bias"), + # Cross-attention (proj) + (f"blocks.{depth}.cross_attn.proj.weight" ,f"transformer_blocks.{depth}.attn2.to_out.0.weight"), + (f"blocks.{depth}.cross_attn.proj.bias" ,f"transformer_blocks.{depth}.attn2.to_out.0.bias"), + ] + +def convert_pixart_state_dict(unet_state_dict): + if "adaln_single.emb.resolution_embedder.linear_1.weight" in unet_state_dict.keys(): + cmap = conversion_map + conversion_map_ms + else: + cmap = conversion_map + + new_state_dict = {k: unet_state_dict.pop(v) for k,v in cmap} + + for depth in range(28): + # Self Attention + q = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_q.weight") + k = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_k.weight") + v = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_v.weight") + new_state_dict[f"blocks.{depth}.attn.qkv.weight"] = torch.cat((q,k,v), dim=0) + qb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_q.bias") + kb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_k.bias") + vb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_v.bias") + new_state_dict[f"blocks.{depth}.attn.qkv.bias"] = torch.cat((qb,kb,vb), dim=0) + + # Cross-attention (linear) + q = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_q.weight") + k = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_k.weight") + v = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_v.weight") + new_state_dict[f"blocks.{depth}.cross_attn.q_linear.weight"] = q + new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.weight"] = torch.cat((k,v), dim=0) + qb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_q.bias") + kb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_k.bias") + vb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_v.bias") + new_state_dict[f"blocks.{depth}.cross_attn.q_linear.bias"] = qb + new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.bias"] = torch.cat((kb,vb), dim=0) + + if len(unet_state_dict.keys()) > 0: + print(f"PixArt: UNET conversion has leftover keys!:\n{unet_state_dict.keys()}") + + return new_state_dict diff --git a/PixArt/loader.py b/PixArt/loader.py index 78005cd..1757a79 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -5,8 +5,7 @@ import comfy.model_base import comfy.utils import torch from comfy import model_management - -from .models import PixArtMS +from .diffusers_convert import convert_pixart_state_dict class EXM_PixArt(comfy.supported_models_base.BASE): unet_config = {} @@ -27,18 +26,18 @@ class EXM_PixArt(comfy.supported_models_base.BASE): def load_pixart(model_path, model_conf): state_dict = comfy.utils.load_torch_file(model_path) state_dict = state_dict.get("model", state_dict) + if "caption_projection.y_embedding" in state_dict: + state_dict = convert_pixart_state_dict(state_dict) # Diffusers parameters = comfy.utils.calculate_parameters(state_dict) unet_dtype = model_management.unet_dtype(model_params=parameters) model_conf = EXM_PixArt(model_conf) # convert to object - model = comfy.model_base.BaseModel( model_conf, model_type=comfy.model_base.ModelType.EPS, device=model_management.get_torch_device() ) - model.pixart_config = model_conf if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config) @@ -48,7 +47,9 @@ def load_pixart(model_path, model_conf): else: raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'") - model.diffusion_model.load_state_dict(state_dict) + m, u = model.diffusion_model.load_state_dict(state_dict, strict=False) + if len(m) > 0: print("Missing UNET keys", m) + if len(u) > 0: print("Leftover UNET keys", u) model.diffusion_model.dtype = unet_dtype model.diffusion_model.eval() model.diffusion_model.to(unet_dtype)