sd3 empty latent, change entrypoint function name

This commit is contained in:
Sida Liu
2024-06-14 06:39:29 +00:00
parent 529a0d42e0
commit a68641e37e
4 changed files with 33 additions and 4 deletions
+3
View File
@@ -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",
}
+26
View File
@@ -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}, )
+2 -2
View File
@@ -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()
+2 -2
View File
@@ -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)