diff --git a/PixArt/loader.py b/PixArt/loader.py index bc0ca00..4fcee7c 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -9,8 +9,8 @@ import comfy.supported_models_base import comfy.supported_models import comfy.latent_formats -from .model.pixart import PixArt -from .model.pixartms import PixArtMS +from .models.pixart import PixArt +from .models.pixartms import PixArtMS from .diffusers_convert import convert_state_dict from ..utils.loader import load_state_dict_from_config from ..text_encoders.pixart.tenc import PixArtTokenizer, PixArtT5XXL diff --git a/PixArt/model/LICENSE b/PixArt/models/LICENSE similarity index 100% rename from PixArt/model/LICENSE rename to PixArt/models/LICENSE diff --git a/PixArt/model/__init__.py b/PixArt/models/__init__.py similarity index 100% rename from PixArt/model/__init__.py rename to PixArt/models/__init__.py diff --git a/PixArt/model/blocks.py b/PixArt/models/blocks.py similarity index 100% rename from PixArt/model/blocks.py rename to PixArt/models/blocks.py diff --git a/PixArt/model/pixart.py b/PixArt/models/pixart.py similarity index 100% rename from PixArt/model/pixart.py rename to PixArt/models/pixart.py diff --git a/PixArt/model/pixartms.py b/PixArt/models/pixartms.py similarity index 100% rename from PixArt/model/pixartms.py rename to PixArt/models/pixartms.py diff --git a/PixArt/model/utils.py b/PixArt/models/utils.py similarity index 100% rename from PixArt/model/utils.py rename to PixArt/models/utils.py diff --git a/Sana/models/__init__.py b/Sana/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/__init__.py b/__init__.py index eeb3dc1..2f9bee0 100644 --- a/__init__.py +++ b/__init__.py @@ -6,11 +6,7 @@ except ImportError: else: NODE_CLASS_MAPPINGS = {} - # All text encoders - from .text_encoders.nodes import NODE_CLASS_MAPPINGS as Tenc_Nodes - NODE_CLASS_MAPPINGS.update(Tenc_Nodes) - - # Generic nodes + # Generic/universal nodes from .nodes import NODE_CLASS_MAPPINGS as Base_Nodes NODE_CLASS_MAPPINGS.update(Base_Nodes) diff --git a/nodes.py b/nodes.py index 31d7f53..c7ba2f5 100644 --- a/nodes.py +++ b/nodes.py @@ -3,6 +3,7 @@ import comfy.utils from .PixArt.loader import load_pixart_state_dict from .Sana.loader import load_sana_state_dict +from .text_encoders.tenc import load_text_encoder, tenc_names loaders = { "PixArt": load_pixart_state_dict, @@ -31,6 +32,39 @@ class EXMUnetLoader: sd = comfy.utils.load_torch_file(unet_path) return (loader_fn(sd),) +class EXMCLIPLoader: + @classmethod + def INPUT_TYPES(s): + files = [] + files += folder_paths.get_filename_list("clip") + # if "clip_gguf" in folder_paths.folder_names_and_paths: + # files += folder_paths.get_filename_list("clip_gguf") + return { + "required": { + "clip_name": (files, ), + "type": (["PixArt", "MiaoBi", "Sana"],), + } + } + + RETURN_TYPES = ("CLIP",) + FUNCTION = "load_clip" + CATEGORY = "ExtraModels" + TITLE = "CLIPLoader (ExtraModels)" + + def load_clip(self, clip_name, type): + clip_path = folder_paths.get_full_path("clip", clip_name) + clip_type = tenc_names.get(type, None) + + clip = load_text_encoder( + ckpt_paths =[clip_path], + embedding_directory = folder_paths.get_folder_paths("embeddings"), + clip_type = clip_type + ) + return (clip,) + +#class EXMResolutionSelect: + NODE_CLASS_MAPPINGS = { "EXMUnetLoader": EXMUnetLoader, + "EXMCLIPLoader": EXMCLIPLoader, } diff --git a/text_encoders/nodes.py b/text_encoders/nodes.py deleted file mode 100644 index 5bda333..0000000 --- a/text_encoders/nodes.py +++ /dev/null @@ -1,37 +0,0 @@ -import folder_paths - -from .tenc import load_text_encoder, tenc_names - -class EXMCLIPLoader: - @classmethod - def INPUT_TYPES(s): - files = [] - files += folder_paths.get_filename_list("clip") - # if "clip_gguf" in folder_paths.folder_names_and_paths: - # files += folder_paths.get_filename_list("clip_gguf") - return { - "required": { - "clip_name": (files, ), - "type": (["PixArt", "MiaoBi", "Sana"],), - } - } - - RETURN_TYPES = ("CLIP",) - FUNCTION = "load_clip" - CATEGORY = "ExtraModels" - TITLE = "CLIPLoader (ExtraModels)" - - def load_clip(self, clip_name, type): - clip_path = folder_paths.get_full_path("clip", clip_name) - clip_type = tenc_names.get(type, None) - - clip = load_text_encoder( - ckpt_paths =[clip_path], - embedding_directory = folder_paths.get_folder_paths("embeddings"), - clip_type = clip_type - ) - return (clip,) - -NODE_CLASS_MAPPINGS = { - "EXMCLIPLoader": EXMCLIPLoader, -}