Add support for PixArt diffusers weights
This commit is contained in:
@@ -0,0 +1,95 @@
|
|||||||
|
# For using the diffusers format weights
|
||||||
|
# Based on the original ComfyUI function +
|
||||||
|
# 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"),
|
||||||
|
("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"),
|
||||||
|
("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"),
|
||||||
|
("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"),
|
||||||
|
# Aspect ratio
|
||||||
|
("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"),
|
||||||
|
("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"),
|
||||||
|
("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"),
|
||||||
|
("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 convert_pixart_state_dict(unet_state_dict):
|
||||||
|
if "adaln_single.emb.resolution_embedder.linear_1.weight" in unet_state_dict.keys():
|
||||||
|
cmap = conversion_map + conversion_map_ms
|
||||||
|
else:
|
||||||
|
cmap = conversion_map
|
||||||
|
|
||||||
|
new_state_dict = {k: unet_state_dict.pop(v) for k,v in cmap}
|
||||||
|
|
||||||
|
for depth in range(28):
|
||||||
|
# Self Attention
|
||||||
|
q = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_q.weight")
|
||||||
|
k = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_k.weight")
|
||||||
|
v = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_v.weight")
|
||||||
|
new_state_dict[f"blocks.{depth}.attn.qkv.weight"] = torch.cat((q,k,v), dim=0)
|
||||||
|
qb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_q.bias")
|
||||||
|
kb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_k.bias")
|
||||||
|
vb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn1.to_v.bias")
|
||||||
|
new_state_dict[f"blocks.{depth}.attn.qkv.bias"] = torch.cat((qb,kb,vb), dim=0)
|
||||||
|
|
||||||
|
# Cross-attention (linear)
|
||||||
|
q = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_q.weight")
|
||||||
|
k = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_k.weight")
|
||||||
|
v = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_v.weight")
|
||||||
|
new_state_dict[f"blocks.{depth}.cross_attn.q_linear.weight"] = q
|
||||||
|
new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.weight"] = torch.cat((k,v), dim=0)
|
||||||
|
qb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_q.bias")
|
||||||
|
kb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_k.bias")
|
||||||
|
vb = unet_state_dict.pop(f"transformer_blocks.{depth}.attn2.to_v.bias")
|
||||||
|
new_state_dict[f"blocks.{depth}.cross_attn.q_linear.bias"] = qb
|
||||||
|
new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.bias"] = torch.cat((kb,vb), dim=0)
|
||||||
|
|
||||||
|
if len(unet_state_dict.keys()) > 0:
|
||||||
|
print(f"PixArt: UNET conversion has leftover keys!:\n{unet_state_dict.keys()}")
|
||||||
|
|
||||||
|
return new_state_dict
|
||||||
+6
-5
@@ -5,8 +5,7 @@ import comfy.model_base
|
|||||||
import comfy.utils
|
import comfy.utils
|
||||||
import torch
|
import torch
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
|
from .diffusers_convert import convert_pixart_state_dict
|
||||||
from .models import PixArtMS
|
|
||||||
|
|
||||||
class EXM_PixArt(comfy.supported_models_base.BASE):
|
class EXM_PixArt(comfy.supported_models_base.BASE):
|
||||||
unet_config = {}
|
unet_config = {}
|
||||||
@@ -27,18 +26,18 @@ class EXM_PixArt(comfy.supported_models_base.BASE):
|
|||||||
def load_pixart(model_path, model_conf):
|
def load_pixart(model_path, model_conf):
|
||||||
state_dict = comfy.utils.load_torch_file(model_path)
|
state_dict = comfy.utils.load_torch_file(model_path)
|
||||||
state_dict = state_dict.get("model", state_dict)
|
state_dict = state_dict.get("model", state_dict)
|
||||||
|
if "caption_projection.y_embedding" in state_dict:
|
||||||
|
state_dict = convert_pixart_state_dict(state_dict) # Diffusers
|
||||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
parameters = comfy.utils.calculate_parameters(state_dict)
|
||||||
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
||||||
|
|
||||||
model_conf = EXM_PixArt(model_conf) # convert to object
|
model_conf = EXM_PixArt(model_conf) # convert to object
|
||||||
|
|
||||||
model = comfy.model_base.BaseModel(
|
model = comfy.model_base.BaseModel(
|
||||||
model_conf,
|
model_conf,
|
||||||
model_type=comfy.model_base.ModelType.EPS,
|
model_type=comfy.model_base.ModelType.EPS,
|
||||||
device=model_management.get_torch_device()
|
device=model_management.get_torch_device()
|
||||||
)
|
)
|
||||||
|
|
||||||
model.pixart_config = model_conf
|
|
||||||
if model_conf.model_target == "PixArtMS":
|
if model_conf.model_target == "PixArtMS":
|
||||||
from .models.PixArtMS import PixArtMS
|
from .models.PixArtMS import PixArtMS
|
||||||
model.diffusion_model = PixArtMS(**model_conf.unet_config)
|
model.diffusion_model = PixArtMS(**model_conf.unet_config)
|
||||||
@@ -48,7 +47,9 @@ def load_pixart(model_path, model_conf):
|
|||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
|
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
|
||||||
|
|
||||||
model.diffusion_model.load_state_dict(state_dict)
|
m, u = model.diffusion_model.load_state_dict(state_dict, strict=False)
|
||||||
|
if len(m) > 0: print("Missing UNET keys", m)
|
||||||
|
if len(u) > 0: print("Leftover UNET keys", u)
|
||||||
model.diffusion_model.dtype = unet_dtype
|
model.diffusion_model.dtype = unet_dtype
|
||||||
model.diffusion_model.eval()
|
model.diffusion_model.eval()
|
||||||
model.diffusion_model.to(unet_dtype)
|
model.diffusion_model.to(unet_dtype)
|
||||||
|
|||||||
Reference in New Issue
Block a user