diff --git a/__init__.py b/__init__.py index 7b208d0..444e4a3 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/nodes/model_nodes.py b/nodes/model_nodes.py index 999c35b..427897a 100644 --- a/nodes/model_nodes.py +++ b/nodes/model_nodes.py @@ -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) diff --git a/nodes/tests/model_nodes_test.py b/nodes/tests/model_nodes_test.py index 69ee882..45a7f74 100644 --- a/nodes/tests/model_nodes_test.py +++ b/nodes/tests/model_nodes_test.py @@ -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