feat: add CLIP-only checkpoint loader

This commit is contained in:
tianzi
2026-08-05 00:41:27 +01:00
parent 5fd77f478b
commit 71b53a8ca7
3 changed files with 115 additions and 1 deletions
+3 -1
View File
@@ -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",
+48
View File
@@ -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)
+64
View File
@@ -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