diff --git a/__init__.py b/__init__.py index 11c48eb..24d6fd0 100644 --- a/__init__.py +++ b/__init__.py @@ -34,6 +34,7 @@ NODE_CLASS_MAPPINGS = { "VrchBooleanKeyControlNode": VrchBooleanKeyControlNode, "VrchChannelOSCControlNode": VrchChannelOSCControlNode, "VrchChannelX4OSCControlNode": VrchChannelX4OSCControlNode, + "VrchControlNetLoaderNode": VrchControlNetLoaderNode, "VrchDelayNode": VrchDelayNode, "VrchDelayOSCControlNode": VrchDelayOSCControlNode, "VrchFloatKeyControlNode": VrchFloatKeyControlNode, @@ -108,6 +109,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "VrchBooleanKeyControlNode": "BOOLEAN Key Control @ vrch.ai", "VrchChannelOSCControlNode": "CHANNEL OSC Control @ vrch.ai", "VrchChannelX4OSCControlNode": "CHANNEL x4 OSC Control @ vrch.ai", + "VrchControlNetLoaderNode": "ControlNet Loader (CPU Offload) @ vrch.ai", "VrchDelayNode": "DELAY @ vrch.ai", "VrchDelayOSCControlNode": "DELAY OSC Control @ vrch.ai", "VrchFloatKeyControlNode": "FLOAT Key Control @ vrch.ai", diff --git a/nodes/model_nodes.py b/nodes/model_nodes.py index a446087..b1d5ae4 100644 --- a/nodes/model_nodes.py +++ b/nodes/model_nodes.py @@ -1,5 +1,6 @@ """Model loading and fallback nodes for ComfyUI workflows.""" +import threading from pathlib import Path import folder_paths @@ -24,6 +25,8 @@ _TENSORRT_MODEL_TYPES = { "FluxSchnell": "flux_schnell", } +_CONTROLNET_CPU_LOAD_LOCK = threading.Lock() + def _register_output_engine_root(): get_output_directory = getattr(folder_paths, "get_output_directory", None) @@ -171,6 +174,64 @@ def _set_tensorrt_control_requirement(model, require_controlnet): metadata["control_required"] = required +class VrchControlNetLoaderNode: + """Load ControlNet weights on CPU so a resident TRT Engine is not duplicated.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "control_net_name": ( + folder_paths.get_filename_list("controlnet"), + ) + } + } + + RETURN_TYPES = ("CONTROL_NET",) + FUNCTION = "load_controlnet" + CATEGORY = CATEGORY + + def load_controlnet(self, control_net_name): + controlnet_path = folder_paths.get_full_path_or_raise( + "controlnet", + control_net_name, + ) + + # ComfyUI's --highvram mode normally constructs ControlNet directly on + # CUDA. A residual TensorRT Engine can already own most of a 16 GiB + # device, so construction itself can OOM before model management gets a + # chance to stream or offload weights. Keep this override scoped to the + # synchronous load call and restore it even when checkpoint loading + # fails. ComfyUI executes model-loader nodes on its single prompt worker; + # the lock also prevents overlapping calls through this node. + import torch + import comfy.controlnet + import comfy.model_management + + cpu_device = torch.device("cpu") + with _CONTROLNET_CPU_LOAD_LOCK: + original_offload_device = ( + comfy.model_management.unet_offload_device + ) + comfy.model_management.unet_offload_device = lambda: cpu_device + try: + controlnet = comfy.controlnet.load_controlnet(controlnet_path) + finally: + comfy.model_management.unet_offload_device = ( + original_offload_device + ) + + if controlnet is None: + raise RuntimeError( + "ControlNet checkpoint is invalid and contains no supported model" + ) + print( + "[comfyui-web-viewer] ControlNet CPU-offload load complete: " + f"{control_net_name}" + ) + return (controlnet,) + + class VrchTensorRTAutoLoaderNode: @classmethod def INPUT_TYPES(cls): diff --git a/nodes/tests/model_nodes_test.py b/nodes/tests/model_nodes_test.py index e7c0548..2ba0519 100644 --- a/nodes/tests/model_nodes_test.py +++ b/nodes/tests/model_nodes_test.py @@ -335,5 +335,109 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase): self.assertEqual(options, [model_nodes.NO_ENGINE_OPTION]) +class TestControlNetLoaderNode(unittest.TestCase): + def setUp(self): + self.original_folder_paths = model_nodes.folder_paths + self.original_modules = { + name: sys.modules.get(name) + for name in ( + "torch", + "comfy", + "comfy.controlnet", + "comfy.model_management", + ) + } + 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, error=None): + calls = [] + gpu_device = object() + cpu_device = object() + model_management = types.ModuleType("comfy.model_management") + original_offload = lambda: gpu_device + model_management.unet_offload_device = original_offload + + controlnet_module = types.ModuleType("comfy.controlnet") + + def load_controlnet(path): + calls.append((path, model_management.unet_offload_device())) + if error is not None: + raise error + return types.SimpleNamespace( + control_model_wrapped=types.SimpleNamespace( + offload_device=model_management.unet_offload_device() + ) + ) + + controlnet_module.load_controlnet = load_controlnet + comfy_module = types.ModuleType("comfy") + comfy_module.controlnet = controlnet_module + comfy_module.model_management = model_management + torch_module = types.ModuleType("torch") + torch_module.device = lambda name: cpu_device if name == "cpu" else name + + sys.modules["torch"] = torch_module + sys.modules["comfy"] = comfy_module + sys.modules["comfy.controlnet"] = controlnet_module + sys.modules["comfy.model_management"] = model_management + return calls, model_management, original_offload, cpu_device + + def test_loads_controlnet_with_cpu_offload_and_restores_global(self): + calls, model_management, original_offload, cpu_device = ( + self.install_runtime() + ) + model_nodes.folder_paths = types.SimpleNamespace( + get_full_path_or_raise=lambda folder, name: f"/{folder}/{name}", + ) + + result = model_nodes.VrchControlNetLoaderNode().load_controlnet( + "union.safetensors" + ) + + self.assertEqual( + calls, + [("/controlnet/union.safetensors", cpu_device)], + ) + self.assertIs( + result[0].control_model_wrapped.offload_device, + cpu_device, + ) + self.assertIs( + model_management.unet_offload_device, + original_offload, + ) + + def test_restores_global_after_load_failure(self): + _, model_management, original_offload, _ = self.install_runtime( + RuntimeError("broken checkpoint") + ) + model_nodes.folder_paths = types.SimpleNamespace( + get_full_path_or_raise=lambda _folder, _name: "/broken.safetensors", + ) + + with self.assertRaisesRegex(RuntimeError, "broken checkpoint"): + model_nodes.VrchControlNetLoaderNode().load_controlnet( + "broken.safetensors" + ) + + self.assertIs( + model_management.unet_offload_device, + original_offload, + ) + + if __name__ == "__main__": unittest.main()