From bcdecad7a7d0f980b44e5ccf56c75b41c43f9e84 Mon Sep 17 00:00:00 2001 From: tianzi Date: Tue, 4 Aug 2026 23:45:08 +0100 Subject: [PATCH] perf: preserve ControlNet across TAESD decode --- nodes/model_nodes.py | 57 +++++++++++++++++++++++++++++++-- nodes/tests/model_nodes_test.py | 38 ++++++++++++++++++++++ 2 files changed, 92 insertions(+), 3 deletions(-) diff --git a/nodes/model_nodes.py b/nodes/model_nodes.py index 6f16f3b..887cf0f 100644 --- a/nodes/model_nodes.py +++ b/nodes/model_nodes.py @@ -249,14 +249,24 @@ class VrchTAESDMemoryProfileNode: "step": 64, }, ), - } + }, + "optional": { + "control_net": ("CONTROL_NET",), + "clip": ("CLIP",), + }, } RETURN_TYPES = ("VAE",) FUNCTION = "apply_profile" CATEGORY = CATEGORY - def apply_profile(self, vae, memory_mib=256): + def apply_profile( + self, + vae, + memory_mib=256, + control_net=None, + clip=None, + ): first_stage_model = getattr(vae, "first_stage_model", None) if type(first_stage_model).__name__ != "TAESD": raise RuntimeError( @@ -274,9 +284,50 @@ class VrchTAESDMemoryProfileNode: "kind": "taesd", "memory_mib": int(memory_mib), } + resident_models = [] + if control_net is not None: + control_patcher = getattr( + control_net, + "control_model_wrapped", + None, + ) + if control_patcher is None: + raise RuntimeError( + "ControlNet has no managed model patcher" + ) + resident_models.append(control_patcher) + if clip is not None: + clip_patcher = getattr(clip, "patcher", None) + if clip_patcher is None: + raise RuntimeError("CLIP has no managed model patcher") + resident_models.append(clip_patcher) + + vae_patcher = getattr(vae, "patcher", None) + if resident_models and vae_patcher is None: + raise RuntimeError("TAESD VAE has no managed model patcher") + if resident_models: + base_models = getattr( + vae_patcher, + "_vrch_base_model_patches_models", + vae_patcher.model_patches_models, + ) + vae_patcher._vrch_base_model_patches_models = base_models + vae_patcher._vrch_resident_models = tuple(resident_models) + + def model_patches_models(): + models = list(base_models()) + for model in vae_patcher._vrch_resident_models: + if model not in models: + models.append(model) + return models + + vae_patcher.model_patches_models = model_patches_models + vae.vrch_memory_profile["resident_models"] = len( + resident_models + ) print( "[comfyui-web-viewer] TAESD memory profile active: " - f"{int(memory_mib)} MiB" + f"{int(memory_mib)} MiB; resident_models={len(resident_models)}" ) return (vae,) diff --git a/nodes/tests/model_nodes_test.py b/nodes/tests/model_nodes_test.py index c772fe8..35108e1 100644 --- a/nodes/tests/model_nodes_test.py +++ b/nodes/tests/model_nodes_test.py @@ -469,6 +469,44 @@ class TestTAESDMemoryProfileNode(unittest.TestCase): with self.assertRaisesRegex(RuntimeError, "requires a TAESD VAE"): model_nodes.VrchTAESDMemoryProfileNode().apply_profile(vae, 256) + def test_keeps_controlnet_and_clip_resident_during_vae_loads(self): + class TAESD: + pass + + base_patcher = object() + control_patcher = object() + clip_patcher = object() + vae_patcher = types.SimpleNamespace( + model_patches_models=lambda: [base_patcher], + ) + vae = types.SimpleNamespace( + first_stage_model=TAESD(), + patcher=vae_patcher, + ) + control_net = types.SimpleNamespace( + control_model_wrapped=control_patcher, + ) + clip = types.SimpleNamespace(patcher=clip_patcher) + + model_nodes.VrchTAESDMemoryProfileNode().apply_profile( + vae, + 64, + control_net=control_net, + clip=clip, + ) + model_nodes.VrchTAESDMemoryProfileNode().apply_profile( + vae, + 64, + control_net=control_net, + clip=clip, + ) + + self.assertEqual( + vae_patcher.model_patches_models(), + [base_patcher, control_patcher, clip_patcher], + ) + self.assertEqual(vae.vrch_memory_profile["resident_models"], 2) + if __name__ == "__main__": unittest.main()