perf: preserve ControlNet across TAESD decode
This commit is contained in:
+54
-3
@@ -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,)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user