feat: add CLIP-only checkpoint loader
This commit is contained in:
+3
-1
@@ -14,7 +14,7 @@ from .nodes.audio_music2emo_node import *
|
||||
from .nodes.workflow_export_nodes import *
|
||||
from .nodes.model_nodes import *
|
||||
|
||||
__version__ = "1.1.26"
|
||||
__version__ = "1.1.27"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VrchAnyOSCControlNode": VrchAnyOSCControlNode,
|
||||
@@ -34,6 +34,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"VrchBooleanKeyControlNode": VrchBooleanKeyControlNode,
|
||||
"VrchChannelOSCControlNode": VrchChannelOSCControlNode,
|
||||
"VrchChannelX4OSCControlNode": VrchChannelX4OSCControlNode,
|
||||
"VrchCheckpointClipLoaderNode": VrchCheckpointClipLoaderNode,
|
||||
"VrchControlNetLoaderNode": VrchControlNetLoaderNode,
|
||||
"VrchDelayNode": VrchDelayNode,
|
||||
"VrchDelayOSCControlNode": VrchDelayOSCControlNode,
|
||||
@@ -110,6 +111,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VrchBooleanKeyControlNode": "BOOLEAN Key Control @ vrch.ai",
|
||||
"VrchChannelOSCControlNode": "CHANNEL OSC Control @ vrch.ai",
|
||||
"VrchChannelX4OSCControlNode": "CHANNEL x4 OSC Control @ vrch.ai",
|
||||
"VrchCheckpointClipLoaderNode": "Checkpoint CLIP-only Loader @ vrch.ai",
|
||||
"VrchControlNetLoaderNode": "ControlNet Loader (CPU Offload) @ vrch.ai",
|
||||
"VrchDelayNode": "DELAY @ vrch.ai",
|
||||
"VrchDelayOSCControlNode": "DELAY OSC Control @ vrch.ai",
|
||||
|
||||
@@ -28,6 +28,54 @@ _TENSORRT_MODEL_TYPES = {
|
||||
_CONTROLNET_CPU_LOAD_LOCK = threading.Lock()
|
||||
|
||||
|
||||
class VrchCheckpointClipLoaderNode:
|
||||
"""Load only the CLIP component from a checkpoint.
|
||||
|
||||
TensorRT workflows do not use the checkpoint's diffusion model. Loading
|
||||
the complete checkpoint in ``--highvram`` mode can nevertheless place that
|
||||
unused UNet on CUDA and prevent a second TensorRT Engine from fitting.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (
|
||||
folder_paths.get_filename_list("checkpoints"),
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
FUNCTION = "load_clip"
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def load_clip(self, ckpt_name):
|
||||
import comfy.sd
|
||||
|
||||
checkpoint_path = folder_paths.get_full_path_or_raise(
|
||||
"checkpoints",
|
||||
ckpt_name,
|
||||
)
|
||||
result = comfy.sd.load_checkpoint_guess_config(
|
||||
checkpoint_path,
|
||||
output_vae=False,
|
||||
output_clip=True,
|
||||
output_clipvision=False,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
output_model=False,
|
||||
)
|
||||
if result is None or len(result) < 2 or result[1] is None:
|
||||
raise RuntimeError(
|
||||
"Checkpoint does not contain a supported CLIP text encoder"
|
||||
)
|
||||
print(
|
||||
"[comfyui-web-viewer] Checkpoint CLIP-only load complete: "
|
||||
f"{ckpt_name}"
|
||||
)
|
||||
return (result[1],)
|
||||
|
||||
|
||||
def _register_output_engine_root():
|
||||
get_output_directory = getattr(folder_paths, "get_output_directory", None)
|
||||
registry = getattr(folder_paths, "folder_names_and_paths", None)
|
||||
|
||||
@@ -335,6 +335,70 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
self.assertEqual(options, [model_nodes.NO_ENGINE_OPTION])
|
||||
|
||||
|
||||
class TestCheckpointClipLoaderNode(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.original_folder_paths = model_nodes.folder_paths
|
||||
self.original_modules = {
|
||||
name: sys.modules.get(name)
|
||||
for name in ("comfy", "comfy.sd")
|
||||
}
|
||||
self.addCleanup(self.restore_modules)
|
||||
self.addCleanup(
|
||||
setattr,
|
||||
model_nodes,
|
||||
"folder_paths",
|
||||
self.original_folder_paths,
|
||||
)
|
||||
|
||||
def restore_modules(self):
|
||||
for name, module in self.original_modules.items():
|
||||
if module is None:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = module
|
||||
|
||||
def install_runtime(self, clip=object()):
|
||||
calls = []
|
||||
sd_module = types.ModuleType("comfy.sd")
|
||||
|
||||
def load_checkpoint_guess_config(path, **kwargs):
|
||||
calls.append((path, kwargs))
|
||||
return (None, clip, None, None)
|
||||
|
||||
sd_module.load_checkpoint_guess_config = load_checkpoint_guess_config
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
comfy_module.sd = sd_module
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.sd"] = sd_module
|
||||
model_nodes.folder_paths = types.SimpleNamespace(
|
||||
get_filename_list=lambda _name: ["sdxl.safetensors"],
|
||||
get_full_path_or_raise=lambda folder, name: f"/{folder}/{name}",
|
||||
get_folder_paths=lambda folder: [f"/{folder}"],
|
||||
)
|
||||
return calls, clip
|
||||
|
||||
def test_loads_only_clip_without_constructing_checkpoint_unet(self):
|
||||
calls, clip = self.install_runtime()
|
||||
|
||||
result = model_nodes.VrchCheckpointClipLoaderNode().load_clip(
|
||||
"sdxl.safetensors"
|
||||
)
|
||||
|
||||
self.assertIs(result[0], clip)
|
||||
self.assertEqual(calls[0][0], "/checkpoints/sdxl.safetensors")
|
||||
self.assertFalse(calls[0][1]["output_model"])
|
||||
self.assertFalse(calls[0][1]["output_vae"])
|
||||
self.assertTrue(calls[0][1]["output_clip"])
|
||||
|
||||
def test_rejects_checkpoint_without_clip(self):
|
||||
self.install_runtime(clip=None)
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "does not contain"):
|
||||
model_nodes.VrchCheckpointClipLoaderNode().load_clip(
|
||||
"sdxl.safetensors"
|
||||
)
|
||||
|
||||
|
||||
class TestControlNetLoaderNode(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.original_folder_paths = model_nodes.folder_paths
|
||||
|
||||
Reference in New Issue
Block a user