From af62f98b4ff5ae6806f912d57991cfd6e6f319bc Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Fri, 8 Dec 2023 16:25:01 +0100 Subject: [PATCH] PixArt KSampler support Thanks to @lawrence-cj for his help on this! --- PixArt/conf.py | 44 ++++++++++++++++++++++++++++++-------------- PixArt/loader.py | 22 ++++++++++++++++------ 2 files changed, 46 insertions(+), 20 deletions(-) diff --git a/PixArt/conf.py b/PixArt/conf.py index 922b297..9919c5e 100644 --- a/PixArt/conf.py +++ b/PixArt/conf.py @@ -3,22 +3,38 @@ List of all PixArt model types / settings """ pixart_conf = { "PixArtMS_XL_2": { # models/PixArtMS - "target" : "PixArtMS", - "input_size" : 1024//8, - "lewei_scale" : 2, - "depth" : 28, - "num_heads" : 16, - "patch_size" : 2, - "hidden_size" : 1152, + "target": "PixArtMS", + "unet_config": { + "input_size" : 1024//8, + "lewei_scale" : 2, + "depth" : 28, + "num_heads" : 16, + "patch_size" : 2, + "hidden_size" : 1152, + }, + "sampling_settings": { + "beta_schedule" : "sqrt_linear", + "linear_start" : 0.0001, + "linear_end" : 0.02, + "timesteps" : 1000, + }, }, "PixArt_XL_2": { # models/PixArt - "target" : "PixArt", - "input_size" : 512//8, - "lewei_scale" : 1, - "depth" : 28, - "num_heads" : 16, - "patch_size" : 2, - "hidden_size" : 1152, + "target": "PixArt", + "unet_config": { + "input_size" : 512//8, + "lewei_scale" : 1, + "depth" : 28, + "num_heads" : 16, + "patch_size" : 2, + "hidden_size" : 1152, + }, + "sampling_settings": { + "beta_schedule" : "sqrt_linear", + "linear_start" : 0.0001, + "linear_end" : 0.02, + "timesteps" : 1000, + }, }, } diff --git a/PixArt/loader.py b/PixArt/loader.py index ab42ce7..78005cd 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -13,6 +13,14 @@ class EXM_PixArt(comfy.supported_models_base.BASE): unet_extra_config = {} latent_format = comfy.latent_formats.SD15 + def __init__(self, model_conf): + self.model_target = model_conf.get("target") + self.unet_config = model_conf.get("unet_config", {}) + self.sampling_settings = model_conf.get("sampling_settings", {}) + self.latent_format = self.latent_format() + # UNET is handled by extension + self.unet_config["disable_unet_model_creation"] = True + def model_type(self, state_dict, prefix=""): return comfy.model_base.ModelType.EPS @@ -22,21 +30,23 @@ def load_pixart(model_path, model_conf): parameters = comfy.utils.calculate_parameters(state_dict) unet_dtype = model_management.unet_dtype(model_params=parameters) + model_conf = EXM_PixArt(model_conf) # convert to object + model = comfy.model_base.BaseModel( - EXM_PixArt({"disable_unet_model_creation" : True }), + model_conf, model_type=comfy.model_base.ModelType.EPS, device=model_management.get_torch_device() ) model.pixart_config = model_conf - if model_conf["target"] == "PixArtMS": + if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS - model.diffusion_model = PixArtMS(**model_conf) - elif model_conf["target"] == "PixArt": + model.diffusion_model = PixArtMS(**model_conf.unet_config) + elif model_conf.model_target == "PixArt": from .models.PixArt import PixArt - model.diffusion_model = PixArt(**model_conf) + model.diffusion_model = PixArt(**model_conf.unet_config) else: - raise NotImplementedError + raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'") model.diffusion_model.load_state_dict(state_dict) model.diffusion_model.dtype = unet_dtype