diff --git a/__init__.py b/__init__.py index 836bb46..d53d17e 100644 --- a/__init__.py +++ b/__init__.py @@ -1,12 +1,15 @@ from .nodes.sd3_load_checkpoint import * from .nodes.sd3_load_clips import * +from .nodes.sd3_empty_latent import * NODE_CLASS_MAPPINGS = { "SD3LoadCheckpoint": SD3LoadCheckpoint, "SD3LoadCLIPs": SD3LoadCLIPs, + "SD3EmptyLatent": SD3EmptyLatent, } NODE_DISPLAY_NAME_MAPPINGS = { "SD3LoadCheckpoint": "SD3 Load Checkpoint", "SD3LoadCLIPs": "SD3 Load CLIPs", + "SD3EmptyLatent": "SD3 Empty Latent", } diff --git a/nodes/sd3_empty_latent.py b/nodes/sd3_empty_latent.py new file mode 100644 index 0000000..df101d9 --- /dev/null +++ b/nodes/sd3_empty_latent.py @@ -0,0 +1,26 @@ +import torch +import comfy.model_management + +class SD3EmptyLatent: + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(s): + resolution = [] + # Generate resolutions following this guideline: "Resolution should be around 1 megapixel and width/height must be multiple of 64" + for width in range(512, 1921, 64): + height = int(1024 * 1024 / width / 64) * 64 + resolution.append(f"{width}x{height}") + + return {"required": {"resolution": (resolution, {"default": "1024x1024"}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 16})}} + RETURN_TYPES = ("LATENT",) + FUNCTION = "main" + + CATEGORY = "latent/sd3" + + def main(self, resolution, batch_size=1): + width, height = map(int, resolution.split('x')) + latent = torch.zeros([batch_size, 16, height // 8, width // 8], device=self.device) + return ({"samples":latent}, ) diff --git a/nodes/sd3_load_checkpoint.py b/nodes/sd3_load_checkpoint.py index 229ed70..cb6adb4 100644 --- a/nodes/sd3_load_checkpoint.py +++ b/nodes/sd3_load_checkpoint.py @@ -9,11 +9,11 @@ class SD3LoadCheckpoint (CheckpointLoaderSimple): "shift": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step":0.01}), }} RETURN_TYPES = ("MODEL", "VAE") - FUNCTION = "load_checkpoint" + FUNCTION = "main" CATEGORY = "sd3" - def load_checkpoint(self, ckpt_name, shift): + def main(self, ckpt_name, shift): ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) model, _, vae, _ = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=False, embedding_directory=folder_paths.get_folder_paths("embeddings")) m = model.clone() diff --git a/nodes/sd3_load_clips.py b/nodes/sd3_load_clips.py index 3c3f229..c3a5e6d 100644 --- a/nodes/sd3_load_clips.py +++ b/nodes/sd3_load_clips.py @@ -6,7 +6,7 @@ class SD3LoadCLIPs(TripleCLIPLoader): def INPUT_TYPES(s): return {"required": { "clip_g": (folder_paths.get_filename_list("clip"), ), "clip_l": (folder_paths.get_filename_list("clip"), ), "t5xxl": (folder_paths.get_filename_list("clip"), )}} - FUNCTION = "load_clips" + FUNCTION = "main" - def load_clips(self, clip_g, clip_l, t5xxl): + def main(self, clip_g, clip_l, t5xxl): return self.load_clip(clip_g, clip_l, t5xxl)