diff --git a/README.md b/README.md index 4294549..347d246 100644 --- a/README.md +++ b/README.md @@ -1,2 +1,9 @@ # ComfyUI-Unload-Model - For unloading a model or all models, using the memory management that is already present in ComfyUI. Copied from https://github.com/willblaschko/ComfyUI-Unload-Models but without the unnecessary extra stuff. + +For unloading a model or all models, using the memory management that is already present in ComfyUI. Copied from https://github.com/willblaschko/ComfyUI-Unload-Models but without the unnecessary extra stuff. + +How to install: Clone this repo into the `custom_nodes` folder, then restart ComfyUI. + +How to use: Add in the middle of a workflow to unload a model at that step. Use any value for the `value` field and the model you want to unload for the `model` field, then route the output of the node to wherever you would have routed the input `value`. + +For example, if you want to unload the CLIP models to save VRAM while using Flux, add this node after the `ClipTextEncode` or `ClipTextEncodeFlux` node, using the conditioning for the `value` field, and using the CLIP model for the `model` field, then route the output to wherever you would send the conditioning, e.g. `FluxGuidance` or `BasicGuider`. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2661253 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .unloadModel import NODE_CLASS_MAPPINGS + +__all__ = ['NODE_CLASS_MAPPINGS'] diff --git a/unloadModel.py b/unloadModel.py new file mode 100644 index 0000000..53c2118 --- /dev/null +++ b/unloadModel.py @@ -0,0 +1,86 @@ +import comfy.model_management as model_management +import gc +import torch +import time + +# Note: This doesn't work with reroute for some reason? +class AnyType(str): + def __ne__(self, __value: object) -> bool: + return False + +any = AnyType("*") + +class UnloadModelNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": {"value": (any, )}, # For passthrough + "optional": {"model": (any, )}, + } + + @classmethod + def VALIDATE_INPUTS(s, **kwargs): + return True + + RETURN_TYPES = (any, ) + FUNCTION = "route" + CATEGORY = "Unload Model" + + def route(self, **kwargs): + print("Unload Model:") + loaded_models = model_management.loaded_models() + if kwargs.get("model") in loaded_models: + print(" - Model found in memory, unloading...") + loaded_models.remove(kwargs.get("model")) + model_management.free_memory(1e30, model_management.get_torch_device(), loaded_models) + model_management.soft_empty_cache(True) + try: + print(" - Clearing Cache...") + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + except: + print(" - Unable to clear cache") + #time.sleep(2) # why? + return (list(kwargs.values())) + +class UnloadAllModelsNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": {"value": (any, )}, + } + + @classmethod + def VALIDATE_INPUTS(s, **kwargs): + return True + + RETURN_TYPES = (any, ) + FUNCTION = "route" + CATEGORY = "Unload Model" + + def route(self, **kwargs): + print("Unload Model:") + print(" - Unloading all models...") + model_management.unload_all_models() + model_management.soft_empty_cache(True) + try: + print(" - Clearing Cache...") + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + except: + print(" - Unable to clear cache") + #time.sleep(2) # why? + return (list(kwargs.values())) + + +NODE_CLASS_MAPPINGS = { + "UnloadModel": UnloadModelNode, + "UnloadAllModels": UnloadAllModelsNode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "UnloadModel": "Unload Model", + "UnloadAllModels": "Unload All Models", +}