From 03a81eddd557c4a24728e02c736a13cf7ad3674c Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 21:44:42 +0100 Subject: [PATCH] Get the depth of the model by counting layers to allow for models with more depth. This allows for deeper models to be created. Example: https://huggingface.co/ptx0/pixart-reality-mix which is a 900M model --- PixArt/diffusers_convert.py | 95 +++++++++++++++++++------------------ 1 file changed, 50 insertions(+), 45 deletions(-) diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 7209476..19852a4 100644 --- a/PixArt/diffusers_convert.py +++ b/PixArt/diffusers_convert.py @@ -3,30 +3,6 @@ # 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"), @@ -40,24 +16,53 @@ conversion_map_ms = [ # for multi_scale_train (MS) ("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 get_depth(state_dict): + return sum(key.endswith('.scale_shift_table') for key in state_dict.keys()) + +def get_conversion_map(state_dict): + 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"), ] + # Add actual transformer blocks + for depth in range(get_depth(state_dict)): + # 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"), + ] + return conversion_map + def find_prefix(state_dict, target_key): prefix = "" for k in state_dict.keys(): @@ -68,15 +73,15 @@ def find_prefix(state_dict, target_key): def convert_state_dict(state_dict): if "adaln_single.emb.resolution_embedder.linear_1.weight" in state_dict.keys(): - cmap = conversion_map + conversion_map_ms + cmap = get_conversion_map(state_dict) + conversion_map_ms else: - cmap = conversion_map + cmap = get_conversion_map(state_dict) missing = [k for k,v in cmap if v not in state_dict] new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} matched = list(v for k,v in cmap if v in state_dict.keys()) - for depth in range(28): + for depth in range(get_depth(state_dict)): for wb in ["weight", "bias"]: # Self Attention key = lambda a: f"transformer_blocks.{depth}.attn1.to_{a}.{wb}" @@ -133,7 +138,7 @@ def convert_lora_state_dict(state_dict, peft=True): print(f"Text Encoder not supported for PixArt LoRA, ignoring {len(t5_keys)} keys") cmap = [] - cmap_unet = conversion_map + conversion_map_ms # todo: 512 model + cmap_unet = get_conversion_map(state_dict) + conversion_map_ms # todo: 512 model for k, v in cmap_unet: if v.endswith(".weight"): cmap.append((rep_ak(k), rep_ap(v))) @@ -146,7 +151,7 @@ def convert_lora_state_dict(state_dict, peft=True): matched = list(v for k,v in cmap if v in state_dict.keys()) for fp, fk in ((rep_ap, rep_ak),(rep_bp, rep_bk)): - for depth in range(28): + for depth in range(get_depth(state_dict)): # Self Attention key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") new_state_dict[fk(f"blocks.{depth}.attn.qkv.weight")] = torch.cat((