From 03a81eddd557c4a24728e02c736a13cf7ad3674c Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 21:44:42 +0100 Subject: [PATCH 1/8] 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(( From 1d6c291e752904e08ef96185adeb7034f3391358 Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 21:45:10 +0100 Subject: [PATCH 2/8] Minor optimization to stop a double .to() being needed --- PixArt/models/PixArtMS.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/PixArt/models/PixArtMS.py b/PixArt/models/PixArtMS.py index 79b7614..957811e 100644 --- a/PixArt/models/PixArtMS.py +++ b/PixArt/models/PixArtMS.py @@ -174,7 +174,7 @@ class PixArtMS(PixArt): self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation, base_size=self.base_size ) - ).unsqueeze(0).to(x.device).to(self.dtype) + ).unsqueeze(0).to(device=x.device, dtype=self.dtype) x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 t = self.t_embedder(timestep) # (N, D) @@ -224,7 +224,7 @@ class PixArtMS(PixArt): device=x.device ).repeat(bs, 1) else: - data_info["img_hw"] = img_hw.to(x.dtype).to(x.device) + data_info["img_hw"] = img_hw.to(dtype=x.dtype, device=x.device) if aspect_ratio is None or True: data_info["aspect_ratio"] = torch.tensor( [[x.shape[2]/x.shape[3]]], @@ -232,7 +232,7 @@ class PixArtMS(PixArt): device=x.device ).repeat(bs, 1) else: - data_info["aspect_ratio"] = aspect_ratio.to(x.dtype).to(x.device) + data_info["aspect_ratio"] = aspect_ratio.to(dtype=x.dtype, device=x.device) ## Still accepts the input w/o that dim but returns garbage if len(context.shape) == 3: From e43ecbeb64946ac6eeed0b858509a97b27c2472d Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 23:24:30 +0100 Subject: [PATCH 3/8] Use correct key to get the number of layers, I accidentally used a layer with a 'final' layer giving me an out by one error --- PixArt/diffusers_convert.py | 2 +- PixArt/loader.py | 2 ++ 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 19852a4..1630106 100644 --- a/PixArt/diffusers_convert.py +++ b/PixArt/diffusers_convert.py @@ -17,7 +17,7 @@ conversion_map_ms = [ # for multi_scale_train (MS) ] def get_depth(state_dict): - return sum(key.endswith('.scale_shift_table') for key in state_dict.keys()) + return sum(key.endswith('cross_attn.proj.weight') for key in state_dict.keys()) def get_conversion_map(state_dict): conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) diff --git a/PixArt/loader.py b/PixArt/loader.py index 3e26711..d1581d9 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -76,6 +76,8 @@ def load_pixart(model_path, model_conf): device=model_management.get_torch_device() ) + model_conf.unet_config['depth'] = sum(key.endswith('cross_attn.proj.weight') for key in state_dict.keys()) + if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config) From 5b082f424d848630b0a45765b255eb9f213a1545 Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 23:42:08 +0100 Subject: [PATCH 4/8] Revert "Use correct key to get the number of layers, I accidentally used a layer with a 'final' layer giving me an out by one error" This reverts commit e43ecbeb64946ac6eeed0b858509a97b27c2472d. --- PixArt/diffusers_convert.py | 2 +- PixArt/loader.py | 2 -- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 1630106..19852a4 100644 --- a/PixArt/diffusers_convert.py +++ b/PixArt/diffusers_convert.py @@ -17,7 +17,7 @@ conversion_map_ms = [ # for multi_scale_train (MS) ] def get_depth(state_dict): - return sum(key.endswith('cross_attn.proj.weight') for key in state_dict.keys()) + 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) diff --git a/PixArt/loader.py b/PixArt/loader.py index d1581d9..3e26711 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -76,8 +76,6 @@ def load_pixart(model_path, model_conf): device=model_management.get_torch_device() ) - model_conf.unet_config['depth'] = sum(key.endswith('cross_attn.proj.weight') for key in state_dict.keys()) - if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config) From 91340dedc7d431608f1b6741cbc8e9b632f2e640 Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 23:56:10 +0100 Subject: [PATCH 5/8] Fix loader depth --- PixArt/loader.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/PixArt/loader.py b/PixArt/loader.py index 3e26711..e085d63 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -76,6 +76,8 @@ def load_pixart(model_path, model_conf): device=model_management.get_torch_device() ) + model_conf.unet_config['depth'] = sum(key.endswith('.scale_shift_table') for key in state_dict.keys()) + if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config) From 712f57c915910f6582bffa223f0d0acf40966fd1 Mon Sep 17 00:00:00 2001 From: gchapman Date: Wed, 19 Jun 2024 09:45:33 +0100 Subject: [PATCH 6/8] Use different keys for detection between original model and converted model --- PixArt/diffusers_convert.py | 2 +- PixArt/loader.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 19852a4..45c4672 100644 --- a/PixArt/diffusers_convert.py +++ b/PixArt/diffusers_convert.py @@ -17,7 +17,7 @@ conversion_map_ms = [ # for multi_scale_train (MS) ] def get_depth(state_dict): - return sum(key.endswith('.scale_shift_table') for key in state_dict.keys()) + return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys()) def get_conversion_map(state_dict): conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) diff --git a/PixArt/loader.py b/PixArt/loader.py index e085d63..63294a5 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -76,7 +76,7 @@ def load_pixart(model_path, model_conf): device=model_management.get_torch_device() ) - model_conf.unet_config['depth'] = sum(key.endswith('.scale_shift_table') for key in state_dict.keys()) + model_conf.unet_config['depth'] = sum(key.endswith('mlp.fc1.weight') for key in state_dict.keys()) if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS From 8f42946305742a49a19fdbbd5154b2d93e52faea Mon Sep 17 00:00:00 2001 From: gchapman Date: Sun, 23 Jun 2024 19:44:32 +0100 Subject: [PATCH 7/8] Remove override as the new simple loader works it out for us. --- PixArt/loader.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/PixArt/loader.py b/PixArt/loader.py index 0659bca..930568a 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -81,8 +81,6 @@ def load_pixart(model_path, model_conf=None): device=model_management.get_torch_device() ) - model_conf.unet_config['depth'] = sum(key.endswith('mlp.fc1.weight') for key in state_dict.keys()) - if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config) From 46e4fc8ce32789df03cf9a3258306723b9594320 Mon Sep 17 00:00:00 2001 From: gchapman Date: Sun, 23 Jun 2024 19:48:34 +0100 Subject: [PATCH 8/8] Add new 900M config for the traditional loader, might be better to have depth selectable? --- PixArt/conf.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/PixArt/conf.py b/PixArt/conf.py index 86c1d7d..128146f 100644 --- a/PixArt/conf.py +++ b/PixArt/conf.py @@ -37,6 +37,21 @@ pixart_conf = { }, "sampling_settings" : sampling_settings, }, + "PixArtMS_Sigma_XL_2_900M": { + "target": "PixArtMSSigma", + "unet_config": { + "input_size": 1024 // 8, + "token_num": 300, + "depth": 42, + "num_heads": 16, + "patch_size": 2, + "hidden_size": 1152, + "micro_condition": False, + "pe_interpolation": 2, + "model_max_length": 300, + }, + "sampling_settings": sampling_settings, + }, "PixArtMS_Sigma_XL_2_2K": { "target": "PixArtMSSigma", "unet_config": {