From e288152dae9b55be1ea1ad040dccf61e04d8b1b4 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Wed, 13 Aug 2025 13:37:23 -0500 Subject: [PATCH] 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. --- __init__.py | 208 ++++++----- block_swap.py | 468 ------------------------ debug_utils.py | 8 +- device_memory_audit.py | 351 ++++++++++++++++++ distorch.py | 15 +- distorch_safetensor.py => distorch_2.py | 198 +++++++--- wanvideo.py | 60 ++- 7 files changed, 656 insertions(+), 652 deletions(-) delete mode 100644 block_swap.py create mode 100644 device_memory_audit.py rename distorch_safetensor.py => distorch_2.py (67%) diff --git a/__init__.py b/__init__.py index cb78598..5ded178 100644 --- a/__init__.py +++ b/__init__.py @@ -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") diff --git a/block_swap.py b/block_swap.py deleted file mode 100644 index 9c50270..0000000 --- a/block_swap.py +++ /dev/null @@ -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 diff --git a/debug_utils.py b/debug_utils.py index 08babc5..97695d8 100644 --- a/debug_utils.py +++ b/debug_utils.py @@ -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}") diff --git a/device_memory_audit.py b/device_memory_audit.py new file mode 100644 index 0000000..025991e --- /dev/null +++ b/device_memory_audit.py @@ -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) diff --git a/distorch.py b/distorch.py index bfa270e..ab8e23e 100644 --- a/distorch.py +++ b/distorch.py @@ -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") diff --git a/distorch_safetensor.py b/distorch_2.py similarity index 67% rename from distorch_safetensor.py rename to distorch_2.py index 890eee3..b2fd39e 100644 --- a/distorch_safetensor.py +++ b/distorch_2.py @@ -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 diff --git a/wanvideo.py b/wanvideo.py index 83b80fd..a49c2b0 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -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,)