feat: add memory-safe ControlNet loader
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user