From aa49ad11392683788e164081db328ed758fd9696 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sun, 10 Aug 2025 20:34:48 -0500 Subject: [PATCH] feat: Introduce DisTorch v2 with BlockSwap memory management This commit introduces a major update, "DisTorch v2", which integrates the new `BlockSwap` system for more efficient and dynamic memory management across multiple GPUs. Key changes: - **BlockSwap Integration:** GGUF model loading is completely refactored to use `BlockSwap`, enabling more intelligent VRAM allocation based on tensor analysis. - **Node Renaming:** All SafeTensor loader nodes are renamed from `...DisTorchMultiGPU` to `...DisTorch2MultiGPU` to clearly distinguish the new implementation from the old one. - **Legacy Support:** The previous GGUF loader is preserved as a legacy option for backward compatibility. - **Improved Memory Calculation:** A more accurate memory calculation function (`get_total_memory_v2`) is implemented and used by the new system. --- __init__.py | 51 +++++++++++++---------- block_swap.py | 8 ++-- distorch.py | 112 ++++++++++++++++++++++++++------------------------ 3 files changed, 92 insertions(+), 79 deletions(-) diff --git a/__init__.py b/__init__.py index d54dfa7..51dfdc6 100644 --- a/__init__.py +++ b/__init__.py @@ -179,7 +179,7 @@ from .distorch import ( analyze_ggml_loading, calculate_vvram_allocation_string, override_class_with_distorch_gguf, - override_class_with_distorch_gguf_legacy, + override_class_with_distorch_gguf_v2, override_class_with_distorch_clip, override_class_with_distorch ) @@ -215,22 +215,22 @@ if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS: if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS: NODE_CLASS_MAPPINGS["DiffControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"]) -# DisTorch SafeTensor nodes -NODE_CLASS_MAPPINGS["UNETLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"]) -NODE_CLASS_MAPPINGS["VAELoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"]) -NODE_CLASS_MAPPINGS["CLIPLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"]) -NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"]) +# DisTorch 2 SafeTensor nodes +NODE_CLASS_MAPPINGS["UNETLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"]) +NODE_CLASS_MAPPINGS["VAELoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"]) +NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"]) +NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"]) if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS: - NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"]) + NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"]) if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS: - NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"]) -NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"]) -NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"]) -NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"]) + NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"]) +NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"]) +NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"]) +NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"]) if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS: - NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"]) + NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"]) if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS: - NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"]) + NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"]) # ComfyUI-LTXVideo if check_module_exists("ComfyUI-LTXVideo") or check_module_exists("comfyui-ltxvideo"): @@ -257,21 +257,30 @@ if check_module_exists("ComfyUI-MMAudio") or check_module_exists("comfyui-mmaudi # ComfyUI-GGUF if check_module_exists("ComfyUI-GGUF") or check_module_exists("comfyui-gguf"): - NODE_CLASS_MAPPINGS["UnetLoaderGGUFMultiGPU"] = override_class(UnetLoaderGGUF) + # Legacy DisTorch GGUF nodes NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUF) - NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorchLegacyMultiGPU"] = override_class_with_distorch_gguf_legacy(UnetLoaderGGUF) - NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedMultiGPU"] = override_class(UnetLoaderGGUFAdvanced) NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUFAdvanced) - NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorchLegacyMultiGPU"] = override_class_with_distorch_gguf_legacy(UnetLoaderGGUFAdvanced) - NODE_CLASS_MAPPINGS["CLIPLoaderGGUFMultiGPU"] = override_class_clip(CLIPLoaderGGUF) NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(CLIPLoaderGGUF) - NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFMultiGPU"] = override_class_clip(DualCLIPLoaderGGUF) NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(DualCLIPLoaderGGUF) - NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFMultiGPU"] = override_class_clip(TripleCLIPLoaderGGUF) NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(TripleCLIPLoaderGGUF) - NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFMultiGPU"] = override_class_clip(QuadrupleCLIPLoaderGGUF) NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF) + # DisTorch 2 GGUF nodes + NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_gguf_v2(UnetLoaderGGUF) + NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorch2MultiGPU"] = override_class_with_distorch_gguf_v2(UnetLoaderGGUFAdvanced) + NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(CLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(DualCLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(TripleCLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU"] = override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF) + + # Standard MultiGPU nodes + NODE_CLASS_MAPPINGS["UnetLoaderGGUFMultiGPU"] = override_class(UnetLoaderGGUF) + NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedMultiGPU"] = override_class(UnetLoaderGGUFAdvanced) + NODE_CLASS_MAPPINGS["CLIPLoaderGGUFMultiGPU"] = override_class_clip(CLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFMultiGPU"] = override_class_clip(DualCLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFMultiGPU"] = override_class_clip(TripleCLIPLoaderGGUF) + NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFMultiGPU"] = override_class_clip(QuadrupleCLIPLoaderGGUF) + # PuLID_ComfyUI if check_module_exists("PuLID_ComfyUI") or check_module_exists("pulid_comfyui"): NODE_CLASS_MAPPINGS["PulidModelLoaderMultiGPU"] = override_class(PulidModelLoader) diff --git a/block_swap.py b/block_swap.py index f53884d..55d82c0 100644 --- a/block_swap.py +++ b/block_swap.py @@ -157,10 +157,10 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu", def override_class_with_distorch_safetensor(cls): - """DisTorch wrapper for SafeTensor models, providing block-swap memory optimization.""" + """DisTorch 2.0 wrapper for SafeTensor models, providing block-swap memory optimization.""" from .nodes import get_device_list - class NodeOverrideDisTorchSafeTensor(cls): + class NodeOverrideDisTorchSafeTensorv2(cls): @classmethod def INPUT_TYPES(s): inputs = copy.deepcopy(cls.INPUT_TYPES()) @@ -175,7 +175,7 @@ def override_class_with_distorch_safetensor(cls): inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""}) return inputs - CATEGORY = "multigpu" + CATEGORY = "multigpu/distorch_2" FUNCTION = "override" def override(self, *args, compute_device=None, virtual_ram_gb=4.0, @@ -205,7 +205,7 @@ def override_class_with_distorch_safetensor(cls): return out - return NodeOverrideDisTorchSafeTensor + return NodeOverrideDisTorchSafeTensorv2 # For backwards compatibility, keep the old name pointing to the new safetensor wrapper diff --git a/distorch.py b/distorch.py index 5d499f8..c8ef474 100644 --- a/distorch.py +++ b/distorch.py @@ -290,63 +290,11 @@ def calculate_vvram_allocation_string(model, virtual_vram_str): def override_class_with_distorch_gguf(cls): - """Standardized DisTorch wrapper for GGUF models.""" - from .nodes import get_device_list - from . import current_device - - class NodeOverrideDisTorchGGUF(cls): - @classmethod - def INPUT_TYPES(s): - inputs = copy.deepcopy(cls.INPUT_TYPES()) - devices = get_device_list() - compute_device = devices[1] if len(devices) > 1 else devices[0] - - inputs["optional"] = inputs.get("optional", {}) - inputs["optional"]["compute_device"] = (devices, {"default": compute_device}) - inputs["optional"]["virtual_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1}) - inputs["optional"]["donor_device"] = (devices, {"default": "cpu"}) - inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""}) - return inputs - - CATEGORY = "multigpu" - FUNCTION = "override" - - def override(self, *args, compute_device=None, virtual_ram_gb=4.0, - donor_device="cpu", expert_mode_allocations="", **kwargs): - from . import set_current_device - if compute_device is not None: - set_current_device(compute_device) - - register_patched_ggufmodelpatcher() - fn = getattr(super(), cls.FUNCTION) - out = fn(*args, **kwargs) - - vram_string = "" - if virtual_ram_gb > 0: - vram_string = f"{compute_device};{virtual_ram_gb};{donor_device}" - - full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" - - logging.info(f"[DisTorch GGUF] Full allocation string: {full_allocation}") - - if hasattr(out[0], 'model'): - model_hash = create_model_hash(out[0], "override") - model_allocation_store[model_hash] = full_allocation - elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - model_hash = create_model_hash(out[0].patcher, "override") - model_allocation_store[model_hash] = full_allocation - - return out - - return NodeOverrideDisTorchGGUF - - -def override_class_with_distorch_gguf_legacy(cls): """Legacy DisTorch wrapper for GGUF models for backward compatibility.""" from .nodes import get_device_list from . import current_device - class NodeOverrideDisTorchLegacy(cls): + class NodeOverrideDisTorchGGUFLegacy(cls): @classmethod def INPUT_TYPES(s): inputs = copy.deepcopy(cls.INPUT_TYPES()) @@ -364,6 +312,10 @@ def override_class_with_distorch_gguf_legacy(cls): CATEGORY = "multigpu/legacy" FUNCTION = "override" + if hasattr(cls, 'TITLE'): + TITLE = f"{cls.TITLE} (Legacy)" + else: + TITLE = "Legacy DisTorch Node" def override(self, *args, device=None, expert_mode_allocations=None, use_other_vram=None, virtual_vram_gb=0.0, **kwargs): from . import set_current_device @@ -396,7 +348,59 @@ def override_class_with_distorch_gguf_legacy(cls): return out - return NodeOverrideDisTorchLegacy + return NodeOverrideDisTorchGGUFLegacy + + +def override_class_with_distorch_gguf_v2(cls): + """DisTorch 2.0 wrapper for GGUF models.""" + from .nodes import get_device_list + from . import current_device + + class NodeOverrideDisTorchGGUFv2(cls): + @classmethod + def INPUT_TYPES(s): + inputs = copy.deepcopy(cls.INPUT_TYPES()) + devices = get_device_list() + compute_device = devices[1] if len(devices) > 1 else devices[0] + + inputs["optional"] = inputs.get("optional", {}) + inputs["optional"]["compute_device"] = (devices, {"default": compute_device}) + inputs["optional"]["virtual_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1}) + inputs["optional"]["donor_device"] = (devices, {"default": "cpu"}) + inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""}) + return inputs + + CATEGORY = "multigpu/distorch_2" + FUNCTION = "override" + + def override(self, *args, compute_device=None, virtual_ram_gb=4.0, + donor_device="cpu", expert_mode_allocations="", **kwargs): + from . import set_current_device + if compute_device is not None: + set_current_device(compute_device) + + register_patched_ggufmodelpatcher() + fn = getattr(super(), cls.FUNCTION) + out = fn(*args, **kwargs) + + vram_string = "" + if virtual_ram_gb > 0: + vram_string = f"{compute_device};{virtual_ram_gb};{donor_device}" + + full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" + + logging.info(f"[DisTorch GGUF] Full allocation string: {full_allocation}") + + if hasattr(out[0], 'model'): + model_hash = create_model_hash(out[0], "override") + model_allocation_store[model_hash] = full_allocation + elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): + model_hash = create_model_hash(out[0].patcher, "override") + model_allocation_store[model_hash] = full_allocation + + return out + + return NodeOverrideDisTorchGGUFv2 def override_class_with_distorch_clip(cls):