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": { diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 7209476..45c4672 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('.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) + # 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(( diff --git a/PixArt/models/PixArtMS.py b/PixArt/models/PixArtMS.py index 478deb3..908589c 100644 --- a/PixArt/models/PixArtMS.py +++ b/PixArt/models/PixArtMS.py @@ -181,7 +181,7 @@ class PixArtMS(PixArt): self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=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) @@ -231,7 +231,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]]], @@ -239,7 +239,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: