refactor: Introduce DisTorch V2 architecture

This commit introduces a major architectural refactoring, laying the groundwork for DisTorch V2. The changes focus on improving modularity, memory management, and diagnostics.

Key changes include:
- Renaming `distorch_safetensor.py` to `distorch_2.py` to house the new core logic.
- Deleting the legacy `block_swap.py` module.
- Adding `device_memory_audit.py` for more sophisticated analysis of GPU memory usage.
- Implementing a centralized and configurable logging system in `__init__.py` to provide standardized and level-controlled (DEBUG/INFO) output for better debugging.
This commit is contained in:
John Pollock
2025-08-13 13:37:23 -05:00
parent d5dc678c04
commit e288152dae
7 changed files with 656 additions and 652 deletions
+120 -88
View File
@@ -7,6 +7,16 @@ import folder_paths
import comfy.model_management as mm
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
# --- DisTorch V2 Logging Configuration ---
# Set to "E" for Engineering (DEBUG) or "P" for Production (INFO)
LOG_LEVEL = "P"
# Configure logger
log_level = logging.DEBUG if LOG_LEVEL == "E" else logging.INFO
logging.basicConfig(level=log_level, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
# --- End Logging Configuration ---
# Global device state management
current_device = mm.get_torch_device()
current_text_encoder_device = mm.text_encoder_device()
@@ -34,12 +44,12 @@ def get_device_list():
def set_current_device(device):
global current_device
current_device = device
logging.info(f"[MultiGPU] current_device set to: {device}")
logger.info(f"[MultiGPU] current_device set to: {device}")
def set_current_text_encoder_device(device):
global current_text_encoder_device
current_text_encoder_device = device
logging.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
logger.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
def override_class(cls):
class NodeOverride(cls):
@@ -56,16 +66,14 @@ def override_class(cls):
FUNCTION = "override"
def override(self, *args, device=None, **kwargs):
logging.info(f"[MultiGPU override_class] Called with device={device}")
logger.debug(f"[MultiGPU] override_class called for {cls.__name__} with device={device}")
if device is not None:
set_current_device(device)
logging.info(f"[MultiGPU override_class] Setting current_device to {device}")
fn = getattr(super(), cls.FUNCTION)
logging.info(f"[MultiGPU override_class] Calling wrapped function: {cls.__name__}.{cls.FUNCTION}")
out = fn(*args, **kwargs)
logging.info(f"[MultiGPU override_class] Wrapped function completed successfully")
logger.debug(f"[MultiGPU] override_class for {cls.__name__} completed successfully")
return out
@@ -103,7 +111,7 @@ def get_torch_device_patched():
else:
devs = set(get_device_list())
device = torch.device(current_device) if str(current_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU get_torch_device_patched] Returning device: {device} (current_device={current_device})")
logger.debug(f"[MultiGPU] get_torch_device_patched returning device: {device} (current_device={current_device})")
return device
def text_encoder_device_patched():
@@ -113,24 +121,24 @@ def text_encoder_device_patched():
else:
devs = set(get_device_list())
device = torch.device(current_text_encoder_device) if str(current_text_encoder_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU text_encoder_device_patched] Returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
logger.debug(f"[MultiGPU] text_encoder_device_patched returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
return device
# Apply patches
logging.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
logging.info(f"[MultiGPU] Initial current_device: {current_device}")
logging.info(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
logger.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
logger.debug(f"[MultiGPU] Initial current_device: {current_device}")
logger.debug(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
mm.get_torch_device = get_torch_device_patched
mm.text_encoder_device = text_encoder_device_patched
logging.info(f"[MultiGPU] Patches applied successfully")
logger.debug(f"[MultiGPU] Patches applied successfully")
def check_module_exists(module_path):
full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path)
logging.info(f"MultiGPU: Checking for module at {full_path}")
logger.debug(f"[MultiGPU] Checking for module at {full_path}")
if not os.path.exists(full_path):
logging.info(f"MultiGPU: Module {module_path} not found - skipping")
logger.debug(f"[MultiGPU] Module {module_path} not found - skipping")
return False
logging.info(f"MultiGPU: Found {module_path}, creating compatible MultiGPU nodes")
logger.debug(f"[MultiGPU] Found {module_path}, creating compatible MultiGPU nodes")
return True
# Import from nodes.py
@@ -184,16 +192,8 @@ from .distorch import (
override_class_with_distorch
)
# Import from block_swap.py
from .block_swap import (
analyze_safetensor_distorch,
apply_block_swap,
override_class_with_distorch_safetensor,
override_class_with_distorch_bs
)
# Import from distorch_safetensor.py for FLUX support
from .distorch_safetensor import (
# Import from distorch_2.py for DisTorch v2 SafeTensor support
from .distorch_2 import (
safetensor_allocation_store,
create_safetensor_model_hash,
register_patched_safetensor_modelpatcher,
@@ -242,89 +242,121 @@ if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
# DisTorch 2 FLUX-specific nodes (these are the ones users will use for FLUX)
logging.info("[DISTORCH_SAFETENSOR] Registering FLUX DisTorch2 nodes")
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
NODE_CLASS_MAPPINGS["UNETLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
if "DualCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DualCLIPLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["TripleCLIPLoaderFLUXDisTorch2"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
# --- Registration Table ---
logger.info("[MultiGPU] Initiating custom_node Registration. . .")
dash_line = "-" * 47
fmt_reg = "{:<30}{:>5}{:>10}"
logger.info(dash_line)
logger.info(fmt_reg.format("custom_node", "Found", "Nodes"))
logger.info(dash_line)
registration_data = []
def register_and_count(module_names, node_map):
found = False
for name in module_names:
if check_module_exists(name):
found = True
break
count = 0
if found:
initial_len = len(NODE_CLASS_MAPPINGS)
for key, value in node_map.items():
NODE_CLASS_MAPPINGS[key] = value
count = len(NODE_CLASS_MAPPINGS) - initial_len
registration_data.append({"name": module_names[0], "found": "Y" if found else "N", "count": count})
return found
# ComfyUI-LTXVideo
if check_module_exists("ComfyUI-LTXVideo") or check_module_exists("comfyui-ltxvideo"):
NODE_CLASS_MAPPINGS["LTXVLoaderMultiGPU"] = override_class(LTXVLoader)
ltx_nodes = {"LTXVLoaderMultiGPU": override_class(LTXVLoader)}
register_and_count(["ComfyUI-LTXVideo", "comfyui-ltxvideo"], ltx_nodes)
# ComfyUI-Florence2
if check_module_exists("ComfyUI-Florence2") or check_module_exists("comfyui-florence2"):
NODE_CLASS_MAPPINGS["Florence2ModelLoaderMultiGPU"] = override_class(Florence2ModelLoader)
NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2ModelMultiGPU"] = override_class(DownloadAndLoadFlorence2Model)
florence_nodes = {
"Florence2ModelLoaderMultiGPU": override_class(Florence2ModelLoader),
"DownloadAndLoadFlorence2ModelMultiGPU": override_class(DownloadAndLoadFlorence2Model)
}
register_and_count(["ComfyUI-Florence2", "comfyui-florence2"], florence_nodes)
# ComfyUI_bitsandbytes_NF4
if check_module_exists("ComfyUI_bitsandbytes_NF4") or check_module_exists("comfyui_bitsandbytes_nf4"):
NODE_CLASS_MAPPINGS["CheckpointLoaderNF4MultiGPU"] = override_class(CheckpointLoaderNF4)
nf4_nodes = {"CheckpointLoaderNF4MultiGPU": override_class(CheckpointLoaderNF4)}
register_and_count(["ComfyUI_bitsandbytes_NF4", "comfyui_bitsandbytes_nf4"], nf4_nodes)
# x-flux-comfyui
if check_module_exists("x-flux-comfyui") or check_module_exists("x-flux-comfyui"):
NODE_CLASS_MAPPINGS["LoadFluxControlNetMultiGPU"] = override_class(LoadFluxControlNet)
flux_controlnet_nodes = {"LoadFluxControlNetMultiGPU": override_class(LoadFluxControlNet)}
register_and_count(["x-flux-comfyui"], flux_controlnet_nodes)
# ComfyUI-MMAudio
if check_module_exists("ComfyUI-MMAudio") or check_module_exists("comfyui-mmaudio"):
NODE_CLASS_MAPPINGS["MMAudioModelLoaderMultiGPU"] = override_class(MMAudioModelLoader)
NODE_CLASS_MAPPINGS["MMAudioFeatureUtilsLoaderMultiGPU"] = override_class(MMAudioFeatureUtilsLoader)
NODE_CLASS_MAPPINGS["MMAudioSamplerMultiGPU"] = override_class(MMAudioSampler)
mmaudio_nodes = {
"MMAudioModelLoaderMultiGPU": override_class(MMAudioModelLoader),
"MMAudioFeatureUtilsLoaderMultiGPU": override_class(MMAudioFeatureUtilsLoader),
"MMAudioSamplerMultiGPU": override_class(MMAudioSampler)
}
register_and_count(["ComfyUI-MMAudio", "comfyui-mmaudio"], mmaudio_nodes)
# ComfyUI-GGUF
if check_module_exists("ComfyUI-GGUF") or check_module_exists("comfyui-gguf"):
# Legacy DisTorch GGUF nodes
NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUF)
NODE_CLASS_MAPPINGS["UnetLoaderGGUFAdvancedDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUFAdvanced)
NODE_CLASS_MAPPINGS["CLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(CLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(DualCLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(TripleCLIPLoaderGGUF)
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)
gguf_nodes = {
"UnetLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_gguf(UnetLoaderGGUF),
"UnetLoaderGGUFAdvancedDisTorchMultiGPU": override_class_with_distorch_gguf(UnetLoaderGGUFAdvanced),
"CLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(CLIPLoaderGGUF),
"DualCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(DualCLIPLoaderGGUF),
"TripleCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(TripleCLIPLoaderGGUF),
"QuadrupleCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF),
"UnetLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_gguf_v2(UnetLoaderGGUF),
"UnetLoaderGGUFAdvancedDisTorch2MultiGPU": override_class_with_distorch_gguf_v2(UnetLoaderGGUFAdvanced),
"CLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(CLIPLoaderGGUF),
"DualCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(DualCLIPLoaderGGUF),
"TripleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(TripleCLIPLoaderGGUF),
"QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF),
"UnetLoaderGGUFMultiGPU": override_class(UnetLoaderGGUF),
"UnetLoaderGGUFAdvancedMultiGPU": override_class(UnetLoaderGGUFAdvanced),
"CLIPLoaderGGUFMultiGPU": override_class_clip(CLIPLoaderGGUF),
"DualCLIPLoaderGGUFMultiGPU": override_class_clip(DualCLIPLoaderGGUF),
"TripleCLIPLoaderGGUFMultiGPU": override_class_clip(TripleCLIPLoaderGGUF),
"QuadrupleCLIPLoaderGGUFMultiGPU": override_class_clip(QuadrupleCLIPLoaderGGUF)
}
register_and_count(["ComfyUI-GGUF", "comfyui-gguf"], gguf_nodes)
# PuLID_ComfyUI
if check_module_exists("PuLID_ComfyUI") or check_module_exists("pulid_comfyui"):
NODE_CLASS_MAPPINGS["PulidModelLoaderMultiGPU"] = override_class(PulidModelLoader)
NODE_CLASS_MAPPINGS["PulidInsightFaceLoaderMultiGPU"] = override_class(PulidInsightFaceLoader)
NODE_CLASS_MAPPINGS["PulidEvaClipLoaderMultiGPU"] = override_class(PulidEvaClipLoader)
pulid_nodes = {
"PulidModelLoaderMultiGPU": override_class(PulidModelLoader),
"PulidInsightFaceLoaderMultiGPU": override_class(PulidInsightFaceLoader),
"PulidEvaClipLoaderMultiGPU": override_class(PulidEvaClipLoader)
}
register_and_count(["PuLID_ComfyUI", "pulid_comfyui"], pulid_nodes)
# ComfyUI-HunyuanVideoWrapper
if check_module_exists("ComfyUI-HunyuanVideoWrapper") or check_module_exists("comfyui-hunyuanvideowrapper"):
NODE_CLASS_MAPPINGS["HyVideoModelLoaderMultiGPU"] = override_class(HyVideoModelLoader)
NODE_CLASS_MAPPINGS["HyVideoVAELoaderMultiGPU"] = override_class(HyVideoVAELoader)
NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoderMultiGPU"] = override_class(DownloadAndLoadHyVideoTextEncoder)
hunyuan_nodes = {
"HyVideoModelLoaderMultiGPU": override_class(HyVideoModelLoader),
"HyVideoVAELoaderMultiGPU": override_class(HyVideoVAELoader),
"DownloadAndLoadHyVideoTextEncoderMultiGPU": override_class(DownloadAndLoadHyVideoTextEncoder)
}
register_and_count(["ComfyUI-HunyuanVideoWrapper", "comfyui-hunyuanvideowrapper"], hunyuan_nodes)
# ComfyUI-WanVideoWrapper
if check_module_exists("ComfyUI-WanVideoWrapper") or check_module_exists("comfyui-wanvideowrapper"):
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU"] = WanVideoModelLoader
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU_2"] = WanVideoModelLoader_2
NODE_CLASS_MAPPINGS["WanVideoVAELoaderMultiGPU"] = WanVideoVAELoader
NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoderMultiGPU"] = LoadWanVideoT5TextEncoder
NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoderMultiGPU"] = LoadWanVideoClipTextEncoder
NODE_CLASS_MAPPINGS["WanVideoTextEncodeMultiGPU"] = WanVideoTextEncode
NODE_CLASS_MAPPINGS["WanVideoBlockSwapMultiGPU"] = WanVideoBlockSwap
NODE_CLASS_MAPPINGS["WanVideoSamplerMultiGPU"] = WanVideoSampler
wanvideo_nodes = {
"WanVideoModelLoaderMultiGPU": WanVideoModelLoader,
"WanVideoModelLoaderMultiGPU_2": WanVideoModelLoader_2,
"WanVideoVAELoaderMultiGPU": WanVideoVAELoader,
"LoadWanVideoT5TextEncoderMultiGPU": LoadWanVideoT5TextEncoder,
"LoadWanVideoClipTextEncoderMultiGPU": LoadWanVideoClipTextEncoder,
"WanVideoTextEncodeMultiGPU": WanVideoTextEncode,
"WanVideoBlockSwapMultiGPU": WanVideoBlockSwap,
"WanVideoSamplerMultiGPU": WanVideoSampler
}
register_and_count(["ComfyUI-WanVideoWrapper", "comfyui-wanvideowrapper"], wanvideo_nodes)
logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
# Print the registration table
for item in registration_data:
logger.info(fmt_reg.format(item['name'], item['found'], str(item['count'])))
logger.info(dash_line)
logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
# --- Memory Logging Test ---
from .debug_utils import log_memory_usage
logger.debug("ComfyUI Startup Memory Log:")
log_memory_usage("ComfyUI Startup")
-468
View File
@@ -1,468 +0,0 @@
"""
Block Swap Module for SafeTensor Models
Contains all SafeTensor DisTorch code for block-swap memory optimization
"""
import torch
import logging
import copy
from collections import defaultdict
import comfy.model_management as mm
import torch.nn as nn
from .model_sig import get_model_type
class FluxBlockSwapManager:
"""
Manages block-swapping for FLUX models using the original, fast `forward` patching method.
"""
def __init__(self, model_patcher):
self.model_patcher = model_patcher
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[FluxBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device)
original_forward = block.forward
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
b.to(sd, non_blocking=True)
return result
return patched_forward
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[FluxBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
class QwenBlockSwapManager:
"""
Manages block-swapping for Qwen models.
"""
def __init__(self, model_patcher):
self.model_patcher = model_patcher
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[QwenBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device)
original_forward = block.forward
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
b.to(sd, non_blocking=True)
return result
return patched_forward
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[QwenBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
class WanVideoBlockSwapManager:
"""
Manages block-swapping for WanVideo models using a pre-allocated GPU shell block.
"""
def __init__(self, model_patcher, gpu_shell_block):
self.model_patcher = model_patcher
self.gpu_shell_block = gpu_shell_block
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[WanVideoBlockSwapManager] Applying state_dict patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device) # Ensure the source block is on the swap device
original_forward = block.forward
def create_patched_forward(cpu_block, gpu_shell):
def patched_forward(*args, **kwargs):
logging.info(f"[DEBUG WANVIDEO SWAP] Loading state_dict from CPU block into GPU shell.")
gpu_shell.load_state_dict(cpu_block.state_dict())
logging.info(f"[DEBUG WANVIDEO SWAP] Executing forward pass on GPU shell.")
result = gpu_shell.forward(*args, **kwargs)
return result
return patched_forward
block.forward = create_patched_forward(block, self.gpu_shell_block)
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[WanVideoBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap):
"""Provides a detailed analysis of the block swap configuration, mimicking the GGUF DisTorch style."""
eq_line = "=" * 60
dash_line = "-" * 60
logging.info(eq_line)
logging.info(" DisTorch SafeTensor Memory Analysis")
logging.info(eq_line)
fmt_assign = "{:<12}{:>15}{:>15}{:>15}"
logging.info(fmt_assign.format("Device", "Role", "Total Mem (GB)", "Config (GB)"))
logging.info(dash_line)
compute_total_gb = mm.get_total_memory(torch.device(compute_device)) / (1024**3)
swap_total_gb = mm.get_total_memory(torch.device(swap_device)) / (1024**3)
logging.info(fmt_assign.format(compute_device, "Compute", f"{compute_total_gb:.2f}", ""))
logging.info(fmt_assign.format(swap_device, "Swap", f"{swap_total_gb:.2f}", f"Offload: {virtual_vram_gb:.2f}"))
logging.info(dash_line)
block_summary = defaultdict(lambda: {'count': 0, 'memory': 0})
total_memory = 0
for block in all_blocks:
block_type = type(block).__name__
block_memory = sum(p.numel() * p.element_size() for p in block.parameters())
block_summary[block_type]['count'] += 1
block_summary[block_type]['memory'] += block_memory
total_memory += block_memory
logging.info(" DisTorch SafeTensor Block Analysis")
logging.info(dash_line)
fmt_layer = "{:<20}{:>10}{:>15}{:>12}"
logging.info(fmt_layer.format("Block Type", "Count", "Memory (MB)", "% Total"))
logging.info(dash_line)
sorted_blocks = sorted(block_summary.items(), key=lambda x: x[1]['memory'], reverse=True)
for block_type, data in sorted_blocks:
mem_mb = data['memory'] / (1024 * 1024)
mem_percent = (data['memory'] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_layer.format(block_type, str(data['count']), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
logging.info(dash_line)
logging.info(" DisTorch Final Block Assignments")
logging.info(dash_line)
fmt_final = "{:<5} {:<25} {:>15} {:>15}"
logging.info(fmt_final.format("ID", "Block Type", "Size (MB)", "Assignment"))
logging.info(dash_line)
total_swapped_size_mb = 0
swapped_block_ids = {id(b) for b in blocks_to_swap}
for i, block in enumerate(all_blocks):
block_type = type(block).__name__
size_mb = sum(p.numel() * p.element_size() for p in block.parameters()) / (1024**2)
assignment = "SWAP" if id(block) in swapped_block_ids else "COMPUTE"
if assignment == "SWAP":
total_swapped_size_mb += size_mb
logging.info(fmt_final.format(i, block_type, f"{size_mb:.2f}", assignment))
logging.info(dash_line)
logging.info(f"Total Blocks Swapped: {len(blocks_to_swap)} of {len(all_blocks)}")
logging.info(f"Total VRAM Offloaded: {total_swapped_size_mb / 1024:.2f} GB")
logging.info(eq_line)
def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
virtual_vram_gb=4.0, expert_mode_allocations=""):
"""
Identifies the model type and applies the appropriate block swapping strategy.
"""
model_type = get_model_type(model_patcher)
logging.info(f"[BlockSwap] Detected model type: {model_type}")
if model_type == "FLUX":
manager = FluxBlockSwapManager(model_patcher)
model_to_patch = model_patcher.model.diffusion_model
all_blocks = []
if hasattr(model_to_patch, 'double_blocks'):
all_blocks.extend(model_to_patch.double_blocks)
if hasattr(model_to_patch, 'single_blocks'):
all_blocks.extend(model_to_patch.single_blocks)
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for FLUX model.")
return
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
break
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] FLUX block swap setup complete.")
elif model_type == "QWEN":
manager = QwenBlockSwapManager(model_patcher)
model_to_patch = model_patcher.model.diffusion_model
if not hasattr(model_to_patch, 'transformer_blocks'):
logging.error("[BlockSwap] CRITICAL: Could not find 'transformer_blocks' in Qwen model. Please analyze model structure.")
log_unsupported_model_analysis(model_patcher)
return
all_blocks = model_to_patch.transformer_blocks
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for QWEN model.")
return
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
break
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping for QWEN model.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] QWEN block swap setup complete.")
elif model_type == "WANVIDEO":
model_to_patch = model_patcher.model.diffusion_model
if not hasattr(model_to_patch, 'blocks'):
logging.error("[BlockSwap] CRITICAL: Could not find 'blocks' in WanVideo model. Please analyze model structure.")
log_unsupported_model_analysis(model_patcher)
return
all_blocks = model_to_patch.blocks
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for WanVideo model.")
return
# --- Pre-allocation Strategy ---
# 1. Find the largest block to create a shell
largest_block = max(all_blocks, key=lambda b: sum(p.numel() * p.element_size() for p in b.parameters()))
gpu_shell_block = copy.deepcopy(largest_block).to(compute_device)
shell_size_mb = sum(p.numel() * p.element_size() for p in gpu_shell_block.parameters()) / (1024**2)
logging.info(f"[BlockSwap] Created GPU shell block for WanVideo on {compute_device}, size: {shell_size_mb:.2f} MB")
manager = WanVideoBlockSwapManager(model_patcher, gpu_shell_block)
# --- End Pre-allocation ---
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
# We still need to identify which blocks to swap (i.e., which ones will use the shell)
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
# The rest of the blocks will remain on the compute device and not be patched
block.to(compute_device)
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping for WanVideo model.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] WanVideo block swap setup complete using pre-allocation strategy.")
else:
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported. Logging model structure for analysis.")
log_unsupported_model_analysis(model_patcher)
def log_unsupported_model_analysis(model_patcher):
"""
Logs the structure of an unsupported model for development purposes.
This is a diagnostic tool and does not modify the model.
"""
logging.info("========================================================================")
logging.info(" INTERNAL MODEL ANALYZER (UNSUPPORTED MODEL DETECTED)")
logging.info("========================================================================")
if not hasattr(model_patcher, 'model'):
logging.error("[ModelAnalyzer] Model patcher does not contain a 'model' attribute.")
return
model = model_patcher.model
logging.info(f"[ModelAnalyzer] Root Model Type: {type(model).__name__}")
if not hasattr(model, 'diffusion_model'):
logging.warning("[ModelAnalyzer] Model does not have a 'diffusion_model' attribute. Dumping root model attributes.")
_recursive_log_attrs(model, "model")
else:
diffusion_model = model.diffusion_model
logging.info(f"[ModelAnalyzer] Diffusion Model Type: {type(diffusion_model).__name__}")
_recursive_log_attrs(diffusion_model, "diffusion_model")
logging.info("========================================================================")
logging.info(" MODEL ANALYSIS COMPLETE")
logging.info("========================================================================")
def _recursive_log_attrs(module, path, seen_modules=None):
"""Helper function to recursively log model attributes."""
if seen_modules is None:
seen_modules = set()
if id(module) in seen_modules:
return
seen_modules.add(id(module))
logging.info(f"--- Analyzing path: '{path}' (Type: {type(module).__name__}) ---")
# Log named children first
children_found = False
for name, submodule in module.named_children():
children_found = True
new_path = f"{path}.{name}" if path else name
# Heuristic check for potential block lists
if isinstance(submodule, (torch.nn.ModuleList, list)) and submodule and all(isinstance(x, torch.nn.Module) for x in submodule):
logging.info(f" > [POTENTIAL BLOCK LIST] '{new_path}' | Type: {type(submodule).__name__}, Length: {len(submodule)}")
# Also inspect the first block in the list for more detail
if len(submodule) > 0:
_recursive_log_attrs(submodule[0], f"{new_path}[0]", seen_modules)
else:
logging.info(f" - Child: '{new_path}' | Type: {type(submodule).__name__}")
# Recurse into non-list children
_recursive_log_attrs(submodule, new_path, seen_modules)
if not children_found:
logging.info(" No named children found at this level.")
def override_class_with_distorch_safetensor(cls):
"""DisTorch 2.0 wrapper for SafeTensor models, providing block-swap memory optimization."""
from .nodes import get_device_list
class NodeOverrideDisTorchSafeTensorv2(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_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.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_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
from . import set_current_device
logging.info(f"[DisTorch SafeTensor] Override called with: compute_device={compute_device}, donor_device={donor_device}, virtual_vram_gb={virtual_vram_gb}")
if compute_device is not None:
set_current_device(compute_device)
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
model = out[0]
if hasattr(model, 'model'):
logging.info("[DisTorch SafeTensor] Model has 'model' attribute, applying block swap.")
apply_block_swap(
model,
compute_device=compute_device,
swap_device=donor_device,
virtual_vram_gb=virtual_vram_gb,
expert_mode_allocations=expert_mode_allocations
)
else:
logging.warning("[DisTorch SafeTensor] Loaded object does not have a 'model' attribute, skipping block swap.")
return out
return NodeOverrideDisTorchSafeTensorv2
# For backwards compatibility, keep the old name pointing to the new safetensor wrapper
override_class_with_distorch_bs = override_class_with_distorch_safetensor
+4 -4
View File
@@ -11,7 +11,7 @@ def log_memory_usage(label=""):
ram_used_gb = status.get('ram_used', 0) / (1024**3)
ram_total_gb = status.get('ram_total', 0) / (1024**3)
log_message = f"[MEM_DEBUG] {label} | RAM Used: {ram_used_gb:.2f}/{ram_total_gb:.2f} GB"
log_message = f"[MultiGPU] {label} | RAM Used: {ram_used_gb:.2f}/{ram_total_gb:.2f} GB"
if 'gpus' in status:
for i, gpu in enumerate(status['gpus']):
@@ -19,9 +19,9 @@ def log_memory_usage(label=""):
vram_total_gb = gpu.get('vram_total', 0) / (1024**3)
log_message += f" | VRAM cuda:{i}: {vram_used_gb:.2f}/{vram_total_gb:.2f} GB"
logging.info(log_message)
logging.debug(log_message)
except ImportError:
logging.warning("[MEM_DEBUG] Could not import local CHardwareInfo. Cannot log memory usage.")
logging.warning("[MultiGPU] Could not import local CHardwareInfo. Cannot log memory usage.")
except Exception as e:
logging.error(f"[MEM_DEBUG] Error getting memory usage: {e}")
logging.error(f"[MultiGPU] Error getting memory usage: {e}")
+351
View File
@@ -0,0 +1,351 @@
"""
Device Memory Audit Utility for ComfyUI MultiGPU
Provides tools to inspect and monitor GPU/CPU memory usage and model placement
"""
import torch
import gc
import psutil
import logging
from collections import defaultdict
from typing import Dict, List, Tuple, Any
import comfy.model_management as mm
def format_bytes(bytes_value: int) -> str:
"""Convert bytes to human-readable format"""
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if bytes_value < 1024.0:
return f"{bytes_value:.2f} {unit}"
bytes_value /= 1024.0
return f"{bytes_value:.2f} PB"
def get_device_memory_info() -> Dict[str, Dict[str, Any]]:
"""
Get current memory usage for all available devices
Returns:
Dict with device names as keys and memory info as values
"""
memory_info = {}
# Check CUDA devices
if torch.cuda.is_available():
for i in range(torch.cuda.device_count()):
device_name = f"cuda:{i}"
torch.cuda.synchronize(i) # Ensure accurate memory reading
allocated = torch.cuda.memory_allocated(i)
reserved = torch.cuda.memory_reserved(i)
total = torch.cuda.get_device_properties(i).total_memory
free = total - allocated
memory_info[device_name] = {
"allocated": allocated,
"allocated_str": format_bytes(allocated),
"reserved": reserved,
"reserved_str": format_bytes(reserved),
"total": total,
"total_str": format_bytes(total),
"free": free,
"free_str": format_bytes(free),
"usage_percent": (allocated / total * 100) if total > 0 else 0
}
# Check CPU/System memory
vm = psutil.virtual_memory()
memory_info["cpu"] = {
"allocated": vm.used,
"allocated_str": format_bytes(vm.used),
"total": vm.total,
"total_str": format_bytes(vm.total),
"free": vm.available,
"free_str": format_bytes(vm.available),
"usage_percent": vm.percent
}
return memory_info
def audit_torch_tensors() -> Dict[str, List[Dict[str, Any]]]:
"""
Find all torch tensors in memory and group by device
Returns:
Dict with device names as keys and list of tensor info as values
"""
tensors_by_device = defaultdict(list)
# Iterate through all objects in memory
for obj in gc.get_objects():
try:
if torch.is_tensor(obj):
device_str = str(obj.device)
size_bytes = obj.element_size() * obj.numel()
tensor_info = {
"id": id(obj),
"shape": list(obj.shape),
"dtype": str(obj.dtype),
"size_bytes": size_bytes,
"size_str": format_bytes(size_bytes),
"requires_grad": obj.requires_grad,
"is_leaf": obj.is_leaf
}
tensors_by_device[device_str].append(tensor_info)
except:
# Some objects might not be accessible
pass
# Sort tensors by size for each device
for device in tensors_by_device:
tensors_by_device[device].sort(key=lambda x: x["size_bytes"], reverse=True)
return dict(tensors_by_device)
def audit_model_placement(model) -> Dict[str, List[Tuple[str, str, int]]]:
"""
Audit where each layer of a model is placed
Args:
model: PyTorch model to audit
Returns:
Dict with device names as keys and list of (layer_name, layer_type, size) tuples
"""
device_layers = defaultdict(list)
for name, module in model.named_modules():
# Skip container modules
if len(list(module.children())) > 0:
continue
# Find device for this module
device = None
size_bytes = 0
# Check parameters
for param_name, param in module.named_parameters(recurse=False):
if param is not None:
device = str(param.device)
size_bytes += param.element_size() * param.numel()
if device:
layer_type = type(module).__name__
device_layers[device].append((name, layer_type, size_bytes))
return dict(device_layers)
def audit_comfy_models() -> Dict[str, Any]:
"""
Audit ComfyUI's loaded models and their placement
Returns:
Dict with model information
"""
model_info = {}
# Try to get loaded models from ComfyUI's model management
try:
# Access currently loaded models if available
if hasattr(mm, 'current_loaded_models'):
loaded_models = mm.current_loaded_models
for i, model_data in enumerate(loaded_models):
model_name = f"model_{i}"
if hasattr(model_data, 'model'):
model = model_data.model
placement = audit_model_placement(model)
total_size = 0
for device_layers in placement.values():
for _, _, size in device_layers:
total_size += size
model_info[model_name] = {
"type": type(model).__name__,
"device_placement": placement,
"total_size": total_size,
"total_size_str": format_bytes(total_size)
}
except Exception as e:
logging.warning(f"[MultiGPU] Could not audit ComfyUI models: {e}")
return model_info
def print_memory_audit(detailed=False):
"""
Print a formatted memory audit report
Args:
detailed: If True, include tensor-level details
"""
logging.info("\n" + "=" * 60)
logging.info(" DEVICE MEMORY AUDIT REPORT")
logging.info("=" * 60)
# Device memory overview
memory_info = get_device_memory_info()
logging.info("\n📊 MEMORY USAGE BY DEVICE:")
logging.info("-" * 60)
for device, info in memory_info.items():
logging.info(f"\n{device.upper()}:")
logging.info(f" Total: {info['total_str']:>12}")
logging.info(f" Allocated: {info['allocated_str']:>12} ({info['usage_percent']:.1f}%)")
logging.info(f" Free: {info['free_str']:>12}")
if 'reserved_str' in info:
logging.info(f" Reserved: {info['reserved_str']:>12}")
# Tensor audit
if detailed:
tensors = audit_torch_tensors()
logging.info("\n📦 TENSORS BY DEVICE:")
logging.info("-" * 60)
for device, tensor_list in tensors.items():
total_size = sum(t["size_bytes"] for t in tensor_list)
logging.info(f"\n{device}: {len(tensor_list)} tensors, {format_bytes(total_size)} total")
# Show top 5 largest tensors
for i, tensor in enumerate(tensor_list[:5]):
logging.info(f" #{i+1}: Shape {tensor['shape']}, {tensor['size_str']}, {tensor['dtype']}")
# ComfyUI model audit
models = audit_comfy_models()
if models:
logging.info("\n🤖 COMFYUI LOADED MODELS:")
logging.info("-" * 60)
for model_name, info in models.items():
logging.info(f"\n{model_name} ({info['type']}):")
logging.info(f" Total size: {info['total_size_str']}")
for device, layers in info['device_placement'].items():
device_size = sum(size for _, _, size in layers)
logging.info(f" {device}: {len(layers)} layers, {format_bytes(device_size)}")
logging.info("\n" + "=" * 60)
def get_memory_snapshot() -> Dict[str, Any]:
"""
Get a complete memory snapshot for benchmarking
Returns:
Dict containing all memory information
"""
return {
"device_memory": get_device_memory_info(),
"tensors": audit_torch_tensors(),
"models": audit_comfy_models()
}
def compare_snapshots(before: Dict[str, Any], after: Dict[str, Any]) -> Dict[str, Any]:
"""
Compare two memory snapshots to see what changed
Args:
before: Snapshot taken before an operation
after: Snapshot taken after an operation
Returns:
Dict with differences between snapshots
"""
differences = {}
# Compare device memory
memory_diff = {}
for device in after["device_memory"]:
if device in before["device_memory"]:
before_mem = before["device_memory"][device]["allocated"]
after_mem = after["device_memory"][device]["allocated"]
diff = after_mem - before_mem
memory_diff[device] = {
"before": format_bytes(before_mem),
"after": format_bytes(after_mem),
"difference": format_bytes(abs(diff)),
"increased": diff > 0
}
differences["memory_changes"] = memory_diff
# Compare tensor counts
tensor_diff = {}
for device in after["tensors"]:
after_count = len(after["tensors"][device])
before_count = len(before["tensors"].get(device, []))
if after_count != before_count:
tensor_diff[device] = {
"before": before_count,
"after": after_count,
"difference": after_count - before_count
}
differences["tensor_count_changes"] = tensor_diff
return differences
def print_snapshot_comparison(before: Dict[str, Any], after: Dict[str, Any], label: str = "Operation"):
"""
Print a formatted comparison of two snapshots
Args:
before: Snapshot before operation
after: Snapshot after operation
label: Label for the operation being measured
"""
diff = compare_snapshots(before, after)
logging.info(f"\n📈 MEMORY CHANGES AFTER {label}:")
logging.info("-" * 60)
for device, changes in diff["memory_changes"].items():
symbol = "↑" if changes["increased"] else "↓"
logging.info(f"{device}: {changes['before']} → {changes['after']} ({symbol} {changes['difference']})")
if diff["tensor_count_changes"]:
logging.info("\nTensor count changes:")
for device, changes in diff["tensor_count_changes"].items():
logging.info(f"{device}: {changes['before']} → {changes['after']} ({changes['difference']:+d})")
# Example usage functions for benchmarking
def benchmark_memory_points():
"""
Example function showing how to capture the 4 key memory points
"""
logging.info("\n🔍 Starting Memory Benchmark...")
# 1. Pre-load memory
snapshot_pre = get_memory_snapshot()
print_memory_audit(detailed=False)
# User would load model here
logging.info("\n[Load your model now]")
input("Press Enter when model is loaded...")
# 2. Post-load memory
snapshot_post_load = get_memory_snapshot()
print_snapshot_comparison(snapshot_pre, snapshot_post_load, "MODEL LOAD")
# User would run inference here
logging.info("\n[Run inference now]")
input("Press Enter when inference is complete...")
# 3. Active inference memory (captured during)
snapshot_active = get_memory_snapshot()
print_snapshot_comparison(snapshot_post_load, snapshot_active, "INFERENCE")
# 4. Post-generation memory
logging.info("\n[Waiting for cleanup...]")
torch.cuda.empty_cache() # Force cleanup
gc.collect()
snapshot_post_gen = get_memory_snapshot()
print_snapshot_comparison(snapshot_active, snapshot_post_gen, "CLEANUP")
logging.info("\n✅ Benchmark complete!")
if __name__ == "__main__":
# Run a basic audit when script is executed directly
print_memory_audit(detailed=True)
+8 -7
View File
@@ -22,6 +22,7 @@ def create_model_hash(model, caller):
first_layers = str(list(model.model_state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
logging.debug(f"[MultiGPU_DisTorch_HASH] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
@@ -100,7 +101,7 @@ def analyze_ggml_loading(model, allocations_str):
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Device Allocations")
logging.info(" DisTorch Model Device Allocations")
logging.info(eq_line)
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
logging.info(dash_line)
@@ -133,7 +134,7 @@ def analyze_ggml_loading(model, allocations_str):
memory_by_type[layer_type] += layer_memory
total_memory += layer_memory
logging.info(" DisTorch GGML Layer Distribution")
logging.info(" DisTorch Model Layer Distribution")
logging.info(dash_line)
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
@@ -161,7 +162,7 @@ def analyze_ggml_loading(model, allocations_str):
device_assignments[device] = layer_list[start_idx:end_idx]
current_layer += device_layer_count
logging.info(" DisTorch Final Device/Layer Assignments")
logging.info("DisTorch Model Final Device/Layer Assignments")
logging.info(dash_line)
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
@@ -201,7 +202,7 @@ def calculate_vvram_allocation_string(model, virtual_vram_str):
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Virtual VRAM Analysis")
logging.info(" DisTorch Model Virtual VRAM Analysis")
logging.info(eq_line)
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
logging.info(dash_line)
@@ -284,7 +285,7 @@ def calculate_vvram_allocation_string(model, virtual_vram_str):
allocation_string = ";".join(allocation_parts)
fmt_mem = "{:<20}{:>20}"
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
logging.info(fmt_mem.format("\n v1 Expert String", allocation_string))
return allocation_string
@@ -389,7 +390,7 @@ def override_class_with_distorch_gguf_v2(cls):
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}")
logging.info(f"[MultiGPU_DisTorch] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
@@ -450,7 +451,7 @@ def override_class_with_distorch_clip(cls):
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch] Full allocation string: {full_allocation}")
logging.info(f"[MultiGPU_DisTorch] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
+146 -52
View File
@@ -15,6 +15,7 @@ import comfy.model_patcher
# Global store for safetensor model allocations - EXACTLY like GGUF
safetensor_allocation_store = {}
safetensor_settings_store = {}
def create_safetensor_model_hash(model, caller):
@@ -42,7 +43,7 @@ def create_safetensor_model_hash(model, caller):
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
# DEBUG STATEMENT - ALWAYS LOG THE HASH
logging.info(f"[SAFETENSOR_HASH] Created hash for {caller}: {final_hash[:8]}...")
logging.debug(f"[MULTIGPU_DISTORCHV2_HASH] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
@@ -61,7 +62,7 @@ def register_patched_safetensor_modelpatcher():
allocations = safetensor_allocation_store.get(debug_hash)
if allocations:
logging.info(f"[DISTORCH_SAFETENSOR] Using static allocation for model {debug_hash[:8]}")
logging.info(f"[MULTIGPU_DISTORCHV2] Using static allocation for model {debug_hash[:8]}")
# Parse allocation string and apply static assignment
device_assignments = analyze_safetensor_loading(self, allocations)
@@ -78,7 +79,7 @@ def register_patched_safetensor_modelpatcher():
if hasattr(module, 'weight') or hasattr(module, 'comfy_cast_weights'):
# Move to our assigned device
logging.info(f"[DISTORCH_SAFETENSOR] Moving {block_name} to {target_device}")
logging.debug(f"[MULTIGPU_DISTORCHV2] Moving {block_name} to {target_device}")
module.to(target_device)
# Mark for ComfyUI's cast system if not already marked
if hasattr(module, 'comfy_cast_weights'):
@@ -92,7 +93,7 @@ def register_patched_safetensor_modelpatcher():
comfy.model_patcher.ModelPatcher.partially_load = new_partially_load
comfy.model_patcher.ModelPatcher._distorch_patched = True
logging.info("[DISTORCH_SAFETENSOR] Successfully patched ModelPatcher.partially_load")
logging.info("[MULTIGPU_DISTORCHV2] Successfully patched ModelPatcher.partially_load")
def analyze_safetensor_loading(model_patcher, allocations_str):
@@ -134,7 +135,7 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
# IDENTICAL LOGGING TO DISTORCH
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Safetensor Device Allocations")
logging.info(" DisTorch2 Model Device Allocations")
logging.info(eq_line)
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
logging.info(dash_line)
@@ -158,13 +159,29 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
# Get the actual model from the patcher
model = model_patcher.model if hasattr(model_patcher, 'model') else model_patcher
# Analyze all modules with weights - matching GGML pattern
# First pass: calculate total memory to establish threshold
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight") or hasattr(module, "comfy_cast_weights"):
try:
block_memory = mm.module_size(module)
except:
block_memory = 0
if hasattr(module, 'weight') and module.weight is not None:
block_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, 'bias') and module.bias is not None:
block_memory += module.bias.numel() * module.bias.element_size()
total_memory += block_memory
# Set the minimum block size threshold (0.1% of total model memory)
MIN_BLOCK_THRESHOLD = total_memory * 0.001
# Second pass: analyze and collect all blocks, then filter
all_blocks = []
for name, module in model.named_modules():
if hasattr(module, "weight") or hasattr(module, "comfy_cast_weights"):
block_type = type(module).__name__
block_summary[block_type] = block_summary.get(block_type, 0) + 1
# Calculate memory using ComfyUI's module_size or manual calculation
try:
block_memory = mm.module_size(module)
except:
@@ -174,12 +191,17 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
if hasattr(module, 'bias') and module.bias is not None:
block_memory += module.bias.numel() * module.bias.element_size()
# Populate summary dictionaries with ALL blocks for accurate reporting
block_summary[block_type] = block_summary.get(block_type, 0) + 1
memory_by_type[block_type] += block_memory
total_memory += block_memory
block_list.append((name, module, block_type))
all_blocks.append((name, module, block_type, block_memory))
# Filter out tiny blocks from the distribution list
block_list = [b for b in all_blocks if b[3] >= MIN_BLOCK_THRESHOLD]
tiny_block_list = [b for b in all_blocks if b[3] < MIN_BLOCK_THRESHOLD]
# Log layer distribution - IDENTICAL FORMAT TO GGML
logging.info(" DisTorch Safetensor Layer Distribution")
logging.info(" DisTorch2 Model Layer Distribution")
logging.info(dash_line)
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
@@ -192,60 +214,102 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
logging.info(dash_line)
# Distribute blocks across devices - EXACTLY like GGML
nonzero_devices = [d for d, r in DEVICE_RATIOS_DISTORCH.items() if r > 0]
nonzero_total_ratio = sum(DEVICE_RATIOS_DISTORCH[d] for d in nonzero_devices)
# Distribute blocks sequentially from the tail of the model
device_assignments = {device: [] for device in DEVICE_RATIOS_DISTORCH.keys()}
block_assignments = {} # Map block name to device
total_blocks = len(block_list)
current_block = 0
block_assignments = {}
for idx, device in enumerate(nonzero_devices):
ratio = DEVICE_RATIOS_DISTORCH[device]
if idx == len(nonzero_devices) - 1:
# Last device gets remaining blocks
device_block_count = total_blocks - current_block
# Determine the primary compute device (first non-cpu device)
compute_device = "cuda:0" # Fallback
for dev in sorted_devices:
if dev != "cpu":
compute_device = dev
break
# Calculate total memory to be offloaded to donor devices
total_offload_gb = sum(DEVICE_RATIOS_DISTORCH.get(d, 0) for d in sorted_devices if d != compute_device)
total_offload_bytes = total_offload_gb * (1024**3)
offloaded_bytes = 0
# Iterate from the TAIL of the model
for block_name, module, block_type, block_memory in reversed(block_list):
try:
# block_memory is already calculated
pass
except:
block_memory = 0
if hasattr(module, 'weight') and module.weight is not None:
block_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, 'bias') and module.bias is not None:
block_memory += module.bias.numel() * module.bias.element_size()
# Assign to donor device (currently assumes one donor 'cpu') until target is met
if offloaded_bytes < total_offload_bytes:
# For now, simple offload to CPU, will expand for multi-donor
donor_device = "cpu"
for dev in sorted_devices:
if dev != compute_device:
donor_device = dev
break # Use first available donor
block_assignments[block_name] = donor_device
offloaded_bytes += block_memory
else:
device_block_count = int((ratio / nonzero_total_ratio) * total_blocks)
start_idx = current_block
end_idx = current_block + device_block_count
device_blocks = block_list[start_idx:end_idx]
device_assignments[device] = device_blocks
# Track block name to device mapping
for block_name, module, block_type in device_blocks:
block_assignments[block_name] = device
current_block += device_block_count
# Assign remaining blocks to the primary compute device
block_assignments[block_name] = compute_device
# Explicitly assign tiny blocks to the compute device
if tiny_block_list:
for block_name, module, block_type, block_memory in tiny_block_list:
block_assignments[block_name] = compute_device
# Populate device_assignments from the final block_assignments
for block_name, device in block_assignments.items():
# Find the block in the original list to get all its info
for b_name, b_module, b_type, b_mem in all_blocks:
if b_name == block_name:
device_assignments[device].append((b_name, b_module, b_type, b_mem))
break
# Log final assignments - IDENTICAL FORMAT TO GGML
logging.info(" DisTorch Safetensor Final Device/Layer Assignments")
logging.info("DisTorch2 Model Final Device/Layer Assignments")
logging.info(dash_line)
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
# Calculate and log tiny blocks separately
if tiny_block_list:
tiny_block_memory = sum(b[3] for b in tiny_block_list)
tiny_mem_mb = tiny_block_memory / (1024 * 1024)
tiny_mem_percent = (tiny_block_memory / total_memory) * 100 if total_memory > 0 else 0
device_label = f"{compute_device}(<0.1%)"
logging.info(fmt_assign.format(device_label, str(len(tiny_block_list)), f"{tiny_mem_mb:.2f}", f"{tiny_mem_percent:.1f}%"))
# Log distributed blocks
total_assigned_memory = 0
device_memories = {}
for device, blocks in device_assignments.items():
device_memory = 0
for block_name, module, block_type in blocks:
# Use the memory we calculated earlier
if block_summary[block_type] > 0:
mem_per_layer = memory_by_type[block_type] / block_summary[block_type]
device_memory += mem_per_layer
# Exclude tiny blocks from this calculation
dist_blocks = [b for b in blocks if b[3] >= MIN_BLOCK_THRESHOLD]
if not dist_blocks:
continue
device_memory = sum(b[3] for b in dist_blocks)
device_memories[device] = device_memory
total_assigned_memory += device_memory
sorted_assignments = sorted(device_assignments.keys(), key=lambda d: (d == "cpu", d))
sorted_assignments = sorted(device_memories.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_assignments:
blocks = device_assignments[dev]
# Get only the distributed blocks for the count
dist_blocks = [b for b in device_assignments[dev] if b[3] >= MIN_BLOCK_THRESHOLD]
if not dist_blocks:
continue
mem_mb = device_memories[dev] / (1024 * 1024)
mem_percent = (device_memories[dev] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_assign.format(dev, str(len(blocks)), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
logging.info(fmt_assign.format(dev, str(len(dist_blocks)), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
logging.info(dash_line)
@@ -267,7 +331,7 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Safetensor Virtual VRAM Analysis")
logging.info(" DisTorch2 Model Virtual VRAM Analysis")
logging.info(eq_line)
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
logging.info(dash_line)
@@ -349,7 +413,7 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
allocation_string = ";".join(allocation_parts)
fmt_mem = "{:<20}{:>20}"
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
logging.info(fmt_mem.format("\n v2 Expert String", allocation_string))
return allocation_string
@@ -385,10 +449,6 @@ def override_class_with_distorch_safetensor_v2(cls):
# Register our patched ModelPatcher
register_patched_safetensor_modelpatcher()
# Call original function
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
# Build allocation string - EXACTLY like GGUF
vram_string = ""
@@ -397,15 +457,49 @@ def override_class_with_distorch_safetensor_v2(cls):
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch Safetensor] Full allocation string: {full_allocation}")
# --- Force Model Reload on Setting Change ---
# Create a hash of the DisTorch settings
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()[:8]
# Temporarily load the model to get its hash, without applying our patch yet
fn = getattr(super(), cls.FUNCTION)
temp_out = fn(*args, **kwargs)
model_to_check = None
if hasattr(temp_out[0], 'model'):
model_to_check = temp_out[0]
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
model_to_check = temp_out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
last_settings_hash = safetensor_settings_store.get(model_hash)
if last_settings_hash != settings_hash:
logging.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Forcing reload.")
mm.unload_model(model_to_check)
# Update the settings store *before* reloading
safetensor_settings_store[model_hash] = settings_hash
# Call the loader again now that the model is unloaded
out = fn(*args, **kwargs)
else:
out = temp_out # Use the already loaded model
else:
out = temp_out # Should not happen, but as a fallback
logging.info(f"[MULTIGPU_DISTORCHV2] Full allocation string: {full_allocation}")
# Store allocation for the model - EXACTLY like GGUF
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash # Ensure it's set
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash # Ensure it's set
return out
+27 -33
View File
@@ -49,8 +49,7 @@ class WanVideoModelLoader:
def loadmodel(self, model, base_precision, device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
logging.info(f"[MultiGPU WanVideoModelLoader] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU WanVideoModelLoader] User selected device: {device}")
logging.debug(f"[MultiGPU] WanVideoModelLoader: User selected device: {device}")
selected_device = torch.device(device)
@@ -62,7 +61,7 @@ class WanVideoModelLoader:
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU WanVideoModelLoader] Patching WanVideo modules to use {selected_device}")
logging.debug(f"[MultiGPU] Patching WanVideo modules to use {selected_device}")
original_device = getattr(loader_module, 'device', None)
original_offload = getattr(loader_module, 'offload_device', None)
@@ -72,7 +71,7 @@ class WanVideoModelLoader:
setattr(loader_module, 'device', selected_device)
if model_offload_override:
setattr(loader_module, 'offload_device', model_offload_override)
logging.info(f"[MultiGPU WanVideoModelLoader] Using model offload override: {model_offload_override}")
logging.debug(f"[MultiGPU] Using model offload override: {model_offload_override}")
elif device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
@@ -86,9 +85,9 @@ class WanVideoModelLoader:
setattr(nodes_module, 'offload_device', nodes_model_offload_override)
elif device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
logging.info(f"[MultiGPU WanVideoModelLoader] Both modules patched successfully")
logging.debug(f"[MultiGPU] Both WanVideo modules patched successfully")
logging.info(f"[MultiGPU WanVideoModelLoader] Calling original loader")
logging.debug(f"[MultiGPU] Calling original WanVideo loader")
result = original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
@@ -100,13 +99,13 @@ class WanVideoModelLoader:
block_swap_override = getattr(loader_module, '_block_swap_device_override', None)
if block_swap_override:
transformer.offload_device = block_swap_override
logging.info(f"[MultiGPU WanVideoModelLoader] Patched transformer for block swap to use: {block_swap_override}")
logging.debug(f"[MultiGPU] Patched WanVideo transformer for block swap to use: {block_swap_override}")
logging.info(f"[MultiGPU WanVideoModelLoader] Model loaded on {selected_device}")
logging.info(f"[MultiGPU] WanVideo model loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoModelLoader] Could not patch modules, falling back")
logging.error(f"[MultiGPU] Could not patch WanVideo modules, falling back")
return original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
@@ -137,7 +136,7 @@ class WanVideoVAELoader:
DESCRIPTION = "Loads Wan VAE model with explicit device selection"
def loadmodel(self, model_name, device, precision="bf16", compile_args=None):
logging.info(f"[MultiGPU WanVideoVAELoader] User selected device: {device}")
logging.debug(f"[MultiGPU] WanVideoVAELoader: User selected device: {device}")
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]()
@@ -146,7 +145,7 @@ class WanVideoVAELoader:
if loader_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoVAELoader] Patching modules to use {selected_device}")
logging.debug(f"[MultiGPU] Patching WanVideo VAE modules to use {selected_device}")
setattr(loader_module, 'offload_device', selected_device)
setattr(loader_module, 'device', selected_device)
@@ -159,10 +158,10 @@ class WanVideoVAELoader:
result = original_loader.loadmodel(model_name, precision, compile_args)
logging.info(f"[MultiGPU WanVideoVAELoader] VAE loaded on {selected_device}")
logging.info(f"[MultiGPU] WanVideo VAE loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoVAELoader] Could not patch modules")
logging.error(f"[MultiGPU] Could not patch WanVideo VAE modules")
return original_loader.loadmodel(model_name, precision, compile_args)
@@ -193,8 +192,7 @@ class LoadWanVideoT5TextEncoder:
DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/text_encoders'"
def loadmodel(self, model_name, precision, device, quantization="disabled"):
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] User selected device: {device}")
logging.debug(f"[MultiGPU] LoadWanVideoT5TextEncoder: User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
@@ -205,7 +203,7 @@ class LoadWanVideoT5TextEncoder:
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Patching WanVideo modules to use {selected_device}")
logging.debug(f"[MultiGPU] Patching WanVideo T5 modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
@@ -220,11 +218,11 @@ class LoadWanVideoT5TextEncoder:
result = original_loader.loadmodel(model_name, precision, load_device, quantization)
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Text encoder loaded on {selected_device}")
logging.info(f"[MultiGPU] WanVideo T5 Text encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoT5TextEncoder] Could not patch modules, falling back")
logging.error(f"[MultiGPU] Could not patch WanVideo T5 modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device, quantization)
class WanVideoTextEncode:
@@ -255,7 +253,7 @@ class WanVideoTextEncode:
def process(self, positive_prompt, negative_prompt, device, t5=None, force_offload=True,
model_to_offload=None, use_disk_cache=False):
logging.info(f"[MultiGPU WanVideoTextEncode] User selected device: {device}")
logging.debug(f"[MultiGPU] WanVideoTextEncode: User selected device: {device}")
original_device = "gpu" if device != "cpu" else "cpu"
@@ -266,7 +264,7 @@ class WanVideoTextEncode:
if encoder_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoTextEncode] Patching module to use {selected_device}")
logging.debug(f"[MultiGPU] Patching WanVideo TextEncode module to use {selected_device}")
setattr(encoder_module, 'device', selected_device)
model_loading_name = encoder_module.__name__.replace('.nodes', '.nodes_model_loading')
@@ -278,7 +276,7 @@ class WanVideoTextEncode:
force_offload=force_offload, model_to_offload=model_to_offload,
use_disk_cache=use_disk_cache, device=original_device)
logging.info(f"[MultiGPU WanVideoTextEncode] Encoding completed on {selected_device}")
logging.info(f"[MultiGPU] WanVideo TextEncode completed on {selected_device}")
return result
else:
return original_encoder.process(positive_prompt, negative_prompt, t5=t5,
@@ -308,8 +306,7 @@ class LoadWanVideoClipTextEncoder:
DESCRIPTION = "Loads Wan CLIP text encoder model from 'ComfyUI/models/clip_vision'"
def loadmodel(self, model_name, precision, device):
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] User selected device: {device}")
logging.debug(f"[MultiGPU] LoadWanVideoClipTextEncoder: User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
@@ -320,7 +317,7 @@ class LoadWanVideoClipTextEncoder:
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] Patching WanVideo modules to use {selected_device}")
logging.debug(f"[MultiGPU] Patching WanVideo CLIP modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
@@ -335,11 +332,11 @@ class LoadWanVideoClipTextEncoder:
result = original_loader.loadmodel(model_name, precision, load_device)
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] CLIP encoder loaded on {selected_device}")
logging.info(f"[MultiGPU] WanVideo CLIP encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoClipTextEncoder] Could not patch modules, falling back")
logging.error(f"[MultiGPU] Could not patch WanVideo CLIP modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device)
class WanVideoModelLoader_2:
@@ -376,7 +373,7 @@ class WanVideoSampler:
def process(self, model, **kwargs):
model_device = model.load_device
logging.info(f"[MultiGPU WanVideoSampler] Processing on device: {model_device}")
logging.info(f"[MultiGPU] WanVideoSampler: Processing on device: {model_device}")
for module_name in sys.modules.keys():
if 'WanVideoWrapper' in module_name and hasattr(sys.modules[module_name], 'device'):
@@ -419,10 +416,7 @@ class WanVideoBlockSwap:
def setargs(self, blocks_to_swap, swap_device, model_offload_device, offload_img_emb, offload_txt_emb,
use_non_blocking=False, vace_blocks_to_swap=0):
logging.info(f"[MultiGPU WanVideoBlockSwap] ========== CONFIGURATION ==========")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected swap device: {swap_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected model offload device: {model_offload_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] Blocks to swap: {blocks_to_swap}")
logging.debug(f"[MultiGPU] WanVideoBlockSwap: swap_device={swap_device}, model_offload_device={model_offload_device}, blocks_to_swap={blocks_to_swap}")
selected_swap_device = torch.device(swap_device)
selected_offload_device = torch.device(model_offload_device)
@@ -433,7 +427,7 @@ class WanVideoBlockSwap:
setattr(module, 'offload_device', selected_offload_device)
setattr(module, '_block_swap_device_override', selected_swap_device)
setattr(module, '_model_offload_device_override', selected_offload_device)
logging.info(f"[MultiGPU WanVideoBlockSwap] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
logging.debug(f"[MultiGPU] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
if 'WanVideoWrapper' in module_name and module_name.endswith('.nodes'):
module = sys.modules[module_name]
@@ -451,6 +445,6 @@ class WanVideoBlockSwap:
"model_offload_device": model_offload_device,
}
logging.info(f"[MultiGPU WanVideoBlockSwap] Block swap configuration complete")
logging.info(f"[MultiGPU] WanVideoBlockSwap configuration complete")
return (block_swap_args,)