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.
This commit is contained in:
+30
-21
@@ -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)
|
||||
|
||||
+4
-4
@@ -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
|
||||
|
||||
+58
-54
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user