sd3 empty latent, change entrypoint function name
This commit is contained in:
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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}, )
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user