Merge pull request #65 from GavChap/get-depth-programatically

Get depth programatically
This commit is contained in:
City
2024-06-23 22:13:59 +02:00
committed by GitHub
3 changed files with 68 additions and 48 deletions
+15
View File
@@ -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": {
+50 -45
View File
@@ -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((
+3 -3
View File
@@ -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: