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
This commit is contained in:
+50
-45
@@ -3,30 +3,6 @@
|
|||||||
# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py
|
# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py
|
||||||
import torch
|
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)
|
conversion_map_ms = [ # for multi_scale_train (MS)
|
||||||
# Resolution
|
# Resolution
|
||||||
("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"),
|
("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"),
|
("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"),
|
||||||
]
|
]
|
||||||
|
|
||||||
# Add actual transformer blocks
|
def get_depth(state_dict):
|
||||||
for depth in range(28):
|
return sum(key.endswith('.scale_shift_table') for key in state_dict.keys())
|
||||||
# Transformer blocks
|
|
||||||
conversion_map += [
|
def get_conversion_map(state_dict):
|
||||||
(f"blocks.{depth}.scale_shift_table", f"transformer_blocks.{depth}.scale_shift_table"),
|
conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers)
|
||||||
# Projection
|
# Patch embeddings
|
||||||
(f"blocks.{depth}.attn.proj.weight", f"transformer_blocks.{depth}.attn1.to_out.0.weight"),
|
("x_embedder.proj.weight", "pos_embed.proj.weight"),
|
||||||
(f"blocks.{depth}.attn.proj.bias", f"transformer_blocks.{depth}.attn1.to_out.0.bias"),
|
("x_embedder.proj.bias", "pos_embed.proj.bias"),
|
||||||
# Feed-forward
|
# Caption projection
|
||||||
(f"blocks.{depth}.mlp.fc1.weight", f"transformer_blocks.{depth}.ff.net.0.proj.weight"),
|
("y_embedder.y_embedding", "caption_projection.y_embedding"),
|
||||||
(f"blocks.{depth}.mlp.fc1.bias", f"transformer_blocks.{depth}.ff.net.0.proj.bias"),
|
("y_embedder.y_proj.fc1.weight", "caption_projection.linear_1.weight"),
|
||||||
(f"blocks.{depth}.mlp.fc2.weight", f"transformer_blocks.{depth}.ff.net.2.weight"),
|
("y_embedder.y_proj.fc1.bias", "caption_projection.linear_1.bias"),
|
||||||
(f"blocks.{depth}.mlp.fc2.bias", f"transformer_blocks.{depth}.ff.net.2.bias"),
|
("y_embedder.y_proj.fc2.weight", "caption_projection.linear_2.weight"),
|
||||||
# Cross-attention (proj)
|
("y_embedder.y_proj.fc2.bias", "caption_projection.linear_2.bias"),
|
||||||
(f"blocks.{depth}.cross_attn.proj.weight" ,f"transformer_blocks.{depth}.attn2.to_out.0.weight"),
|
# AdaLN-single LN
|
||||||
(f"blocks.{depth}.cross_attn.proj.bias" ,f"transformer_blocks.{depth}.attn2.to_out.0.bias"),
|
("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):
|
def find_prefix(state_dict, target_key):
|
||||||
prefix = ""
|
prefix = ""
|
||||||
for k in state_dict.keys():
|
for k in state_dict.keys():
|
||||||
@@ -68,15 +73,15 @@ def find_prefix(state_dict, target_key):
|
|||||||
|
|
||||||
def convert_state_dict(state_dict):
|
def convert_state_dict(state_dict):
|
||||||
if "adaln_single.emb.resolution_embedder.linear_1.weight" in state_dict.keys():
|
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:
|
else:
|
||||||
cmap = conversion_map
|
cmap = get_conversion_map(state_dict)
|
||||||
|
|
||||||
missing = [k for k,v in cmap if v not in 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}
|
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())
|
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"]:
|
for wb in ["weight", "bias"]:
|
||||||
# Self Attention
|
# Self Attention
|
||||||
key = lambda a: f"transformer_blocks.{depth}.attn1.to_{a}.{wb}"
|
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")
|
print(f"Text Encoder not supported for PixArt LoRA, ignoring {len(t5_keys)} keys")
|
||||||
|
|
||||||
cmap = []
|
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:
|
for k, v in cmap_unet:
|
||||||
if v.endswith(".weight"):
|
if v.endswith(".weight"):
|
||||||
cmap.append((rep_ak(k), rep_ap(v)))
|
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())
|
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 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
|
# Self Attention
|
||||||
key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight")
|
key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight")
|
||||||
new_state_dict[fk(f"blocks.{depth}.attn.qkv.weight")] = torch.cat((
|
new_state_dict[fk(f"blocks.{depth}.attn.qkv.weight")] = torch.cat((
|
||||||
|
|||||||
Reference in New Issue
Block a user