feat: add memory-safe ControlNet loader

This commit is contained in:
tianzi
2026-08-04 23:12:24 +01:00
parent 34cb470373
commit ceb4fec353
3 changed files with 167 additions and 0 deletions
+2
View File
@@ -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",
+61
View File
@@ -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):
+104
View File
@@ -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()