From 63ff1a4064ea86f82362a1942fdccf1f07efa9bf Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sat, 20 Sep 2025 11:58:29 -0500 Subject: [PATCH] committing so we don't lose verbose logging. - Add comfyui_memory_load and create_model_identifier utilities (device_utils) - Log GPU memory before/after UNet, VAE, and CLIP construction and after UNet weight load - Include model identifiers in logs to correlate memory to specific patchers - Guard logging calls with try/except to avoid impacting load flow - Improves observability of memory usage for multi-GPU checkpoints and aids OOM/debugging --- checkpoint_multigpu.py | 34 +++++++++++- device_utils.py | 116 +++++++++++++++++++++++++++++++++++++++++ distorch.py | 10 +++- distorch_2.py | 22 +++++++- wanvideo.py | 34 +++++++++++- 5 files changed, 212 insertions(+), 4 deletions(-) diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index 663ccef..f1365c1 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -12,7 +12,7 @@ import comfy.model_management as mm import comfy.model_detection import comfy.clip_vision from comfy.sd import VAE, CLIP -from .device_utils import get_device_list, soft_empty_cache_multigpu +from .device_utils import get_device_list, soft_empty_cache_multigpu, comfyui_memory_load, create_model_identifier from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher logger = logging.getLogger("MultiGPU") @@ -105,11 +105,21 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, set_current_device(unet_compute_device) inital_load_device = mm.unet_inital_load_device(parameters, unet_dtype) + try: + logger.info(comfyui_memory_load(f"pre-model-load:unet:{config_hash[:8]}")) + except Exception: + pass + model = model_config.get_model(sd, diffusion_model_prefix, device=inital_load_device) logger.info("[MultiGPU Checkpoint] Invoking soft_empty_cache_multigpu before UNet ModelPatcher setup") soft_empty_cache_multigpu() model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=unet_compute_device, offload_device=mm.unet_offload_device()) + try: + ident = create_model_identifier(model_patcher) + logger.info(comfyui_memory_load(f"post-model-load:unet:{ident}")) + except Exception: + pass if distorch_config and 'unet_allocation' in distorch_config: register_patched_safetensor_modelpatcher() @@ -121,14 +131,27 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, logger.info(f"Stored DisTorch2 config for UNet (hash {model_hash[:8]}): {distorch_config['unet_allocation']}") model.load_model_weights(sd, diffusion_model_prefix) + try: + ident = create_model_identifier(model_patcher) + logger.info(comfyui_memory_load(f"post-weights-load:unet:{ident}")) + except Exception: + pass if output_vae: vae_target_device = torch.device(device_config.get('vae_device', original_main_device)) set_current_device(vae_target_device) # Use main device context for VAE + try: + logger.info(comfyui_memory_load(f"pre-model-load:vae:{config_hash[:8]}")) + except Exception: + pass vae_sd = comfy.utils.state_dict_prefix_replace(sd, {k: "" for k in model_config.vae_key_prefix}, filter_keys=True) vae_sd = model_config.process_vae_state_dict(vae_sd) vae = VAE(sd=vae_sd, metadata=metadata) + try: + logger.info(comfyui_memory_load(f"post-model-load:vae:{config_hash[:8]}")) + except Exception: + pass if output_clip: clip_target_device = device_config.get('clip_device', original_clip_device) @@ -139,6 +162,10 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, clip_sd = model_config.process_clip_state_dict(sd) if len(clip_sd) > 0: logger.info("[MultiGPU Checkpoint] Invoking soft_empty_cache_multigpu before CLIP construction") + try: + logger.info(comfyui_memory_load(f"pre-model-load:clip:{config_hash[:8]}")) + except Exception: + pass soft_empty_cache_multigpu() clip_params = comfy.utils.calculate_parameters(clip_sd) clip = CLIP(clip_target, embedding_directory=embedding_directory, tokenizer_data=clip_sd, parameters=clip_params, model_options=te_model_options) @@ -157,6 +184,11 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, if len(m) > 0: logger.warning(f"CLIP missing keys: {m}") if len(u) > 0: logger.debug(f"CLIP unexpected keys: {u}") logger.info("CLIP Loaded.") + try: + ident = create_model_identifier(clip.patcher) if hasattr(clip, 'patcher') else f"clip:{config_hash[:8]}" + logger.info(comfyui_memory_load(f"post-model-load:clip:{ident}")) + except Exception: + pass else: logger.warning("No CLIP/text encoder weights in checkpoint.") else: diff --git a/device_utils.py b/device_utils.py index 5c96f0c..966f399 100644 --- a/device_utils.py +++ b/device_utils.py @@ -243,9 +243,19 @@ def soft_empty_cache_multigpu(): import gc logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: starting GC and multi-device cache clear") + # Memory snapshot before GC and soft-empty + try: + logger.info(comfyui_memory_load("pre-soft-empty")) + logger.info(comfyui_memory_load("pre-gc")) + except Exception: + pass # Python GC (same as all implementations) gc.collect() + try: + logger.info(comfyui_memory_load("post-gc")) + except Exception: + pass logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: garbage collection complete") # Clear cache for ALL devices (not just ComfyUI's single device) @@ -261,41 +271,147 @@ def soft_empty_cache_multigpu(): device_idx = int(device_str.split(":")[1]) # Use context manager for safe switching and automatic restoration logger.info(f"[MultiGPU_Device_Utils] Clearing CUDA cache on {device_str} (idx={device_idx})") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass with torch.cuda.device(device_idx): torch.cuda.empty_cache() if hasattr(torch.cuda, "ipc_collect"): torch.cuda.ipc_collect() # ComfyUI's CUDA optimization logger.info(f"[MultiGPU_Device_Utils] Cleared CUDA cache (and IPC if available) on {device_str}") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass elif device_str == "mps": if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"): logger.info("[MultiGPU_Device_Utils] Clearing MPS cache") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass torch.mps.empty_cache() logger.info("[MultiGPU_Device_Utils] Cleared MPS cache") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass elif device_str.startswith("xpu:"): if hasattr(torch, "xpu") and hasattr(torch.xpu, "empty_cache"): logger.info(f"[MultiGPU_Device_Utils] Clearing XPU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass torch.xpu.empty_cache() logger.info(f"[MultiGPU_Device_Utils] Cleared XPU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass elif device_str.startswith("npu:"): if hasattr(torch, "npu") and hasattr(torch.npu, "empty_cache"): logger.info(f"[MultiGPU_Device_Utils] Clearing NPU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass torch.npu.empty_cache() logger.info(f"[MultiGPU_Device_Utils] Cleared NPU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass elif device_str.startswith("mlu:"): if hasattr(torch, "mlu") and hasattr(torch.mlu, "empty_cache"): logger.info(f"[MultiGPU_Device_Utils] Clearing MLU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass torch.mlu.empty_cache() logger.info(f"[MultiGPU_Device_Utils] Cleared MLU cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass elif device_str.startswith("corex:"): if hasattr(torch, "corex") and hasattr(torch.corex, "empty_cache"): logger.info(f"[MultiGPU_Device_Utils] Clearing CoreX cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"pre-empty:{device_str}")) + except Exception: + pass torch.corex.empty_cache() logger.info(f"[MultiGPU_Device_Utils] Cleared CoreX cache on {device_str}") + try: + logger.info(comfyui_memory_load(f"post-empty:{device_str}")) + except Exception: + pass + + # Final memory snapshot after completing soft empty across all devices + try: + logger.info(comfyui_memory_load("post-soft-empty")) + except Exception: + pass + + +def _bytes_to_gib(b: int) -> float: + """Convert bytes to GiB as a float.""" + try: + return float(b) / (1024.0 ** 3) + except Exception: + return 0.0 + + +def comfyui_memory_load(tag: str) -> str: + """ + Returns a single-line, pipe-delimited snapshot of system and device memory usage. + + Format: "tag=|cpu=/|=/|..." + - CPU values represent system RAM via psutil. + - Device values represent VRAM via comfy.model_management across all non-CPU devices. + - Device identifiers use the torch device string from get_device_list() (e.g., 'cuda:0', 'xpu:0', 'mps'). + - Values are in GiB with 2 decimals. + """ + # CPU RAM + vm = psutil.virtual_memory() + cpu_used_gib = _bytes_to_gib(vm.used) + cpu_total_gib = _bytes_to_gib(vm.total) + + segments = [f"tag={tag}", f"cpu={cpu_used_gib:.2f}/{cpu_total_gib:.2f}"] + + # Enumerate non-CPU devices + devices = [d for d in get_device_list() if d != "cpu"] + + # Append per-device VRAM used/total + for dev_str in devices: + try: + device = torch.device(dev_str) + total = mm.get_total_memory(device) + free_info = mm.get_free_memory(device, torch_free_too=True) + # free_info may be a tuple (system_free, torch_cache_free) or a single value + if isinstance(free_info, tuple): + system_free = free_info[0] + else: + system_free = free_info + used = max(0, (total or 0) - (system_free or 0)) + + used_gib = _bytes_to_gib(used) + total_gib = _bytes_to_gib(total or 0) + if total_gib > 0: + segments.append(f"{dev_str}={used_gib:.2f}/{total_gib:.2f}") + except Exception: + # Skip devices that error out (backend not initialized, etc.) + continue + + return "|".join(segments) # ========================================================================================== diff --git a/distorch.py b/distorch.py index 3ec7d92..897090d 100644 --- a/distorch.py +++ b/distorch.py @@ -12,7 +12,7 @@ logger = logging.getLogger("MultiGPU") import copy from collections import defaultdict import comfy.model_management as mm -from .device_utils import get_device_list, soft_empty_cache_multigpu +from .device_utils import get_device_list, soft_empty_cache_multigpu, comfyui_memory_load # Global store for model allocations model_allocation_store = {} @@ -41,8 +41,16 @@ def register_patched_ggufmodelpatcher(): def new_load(self, *args, force_patch_weights=False, **kwargs): global model_allocation_store + try: + logger.info(comfyui_memory_load("pre-model-load:gguf")) + except Exception: + pass super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs) debug_hash = create_model_hash(self, "patcher") + try: + logger.info(comfyui_memory_load(f"post-model-load:gguf:{debug_hash[:8]}")) + except Exception: + pass linked = [] module_count = 0 for n, m in self.model.named_modules(): diff --git a/distorch_2.py b/distorch_2.py index d5d05e8..922fccb 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -17,7 +17,7 @@ from collections import defaultdict import comfy.model_management as mm import comfy.model_patcher from . import current_device -from .device_utils import get_device_list, soft_empty_cache_multigpu +from .device_utils import get_device_list, soft_empty_cache_multigpu, comfyui_memory_load safetensor_allocation_store = {} safetensor_settings_store = {} @@ -64,11 +64,19 @@ def register_patched_safetensor_modelpatcher(): global safetensor_allocation_store debug_hash = create_safetensor_model_hash(self, "partial_load") + try: + logger.info(comfyui_memory_load(f"pre-model-load:safetensor:{debug_hash[:8]}")) + except Exception: + pass allocations = safetensor_allocation_store.get(debug_hash) if not hasattr(self.model, '_distorch_high_precision_loras') or not allocations: result = original_partially_load(self, device_to, extra_memory, force_patch_weights) + try: + logger.info(comfyui_memory_load(f"post-model-load:safetensor:{debug_hash[:8]}")) + except Exception: + pass if hasattr(self, '_distorch_block_assignments'): del self._distorch_block_assignments return result @@ -80,7 +88,15 @@ def register_patched_safetensor_modelpatcher(): if unpatch_weights: logger.info(f"[MultiGPU_DisTorch2] Patches changed or forced. Unpatching model.") + try: + logger.info(comfyui_memory_load(f"pre-model-unload:safetensor:{debug_hash[:8]}")) + except Exception: + pass self.unpatch_model(self.offload_device, unpatch_weights=True) + try: + logger.info(comfyui_memory_load(f"post-model-unload:safetensor:{debug_hash[:8]}")) + except Exception: + pass self.patch_model(load_weights=False) @@ -158,6 +174,10 @@ def register_patched_safetensor_modelpatcher(): self.model.current_weight_patches_uuid = self.patches_uuid logger.info(f"[MultiGPU_DisTorch2] DisTorch loading completed. Total memory: {mem_counter / (1024 * 1024):.2f}MB") + try: + logger.info(comfyui_memory_load(f"post-model-load:safetensor:{debug_hash[:8]}")) + except Exception: + pass return 0 diff --git a/wanvideo.py b/wanvideo.py index ad0d897..40a8045 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -4,7 +4,7 @@ import sys import inspect import folder_paths import comfy.model_management as mm -from .device_utils import get_device_list +from .device_utils import get_device_list, comfyui_memory_load class WanVideoModelLoader: @classmethod @@ -89,8 +89,16 @@ class WanVideoModelLoader: logging.debug(f"[MultiGPU] Both WanVideo modules patched successfully") logging.debug(f"[MultiGPU] Calling original WanVideo loader") + try: + logging.info(comfyui_memory_load(f"pre-model-load:wan-model:{model}")) + except Exception: + pass result = original_loader.loadmodel(model, base_precision, load_device, quantization, compile_args, attention_mode, block_swap_args, lora, vram_management_args, extra_model=extra_model, fantasytalking_model=fantasytalking_model, multitalk_model=multitalk_model, fantasyportrait_model=fantasyportrait_model) + try: + logging.info(comfyui_memory_load(f"post-model-load:wan-model:{model}")) + except Exception: + pass if result and len(result) > 0 and hasattr(result[0], 'model'): model_obj = result[0] @@ -156,7 +164,15 @@ class WanVideoVAELoader: setattr(nodes_module, 'device', selected_device) setattr(nodes_module, 'offload_device', selected_device) + try: + logging.info(comfyui_memory_load(f"pre-model-load:wan-vae:{model_name}")) + except Exception: + pass result = original_loader.loadmodel(model_name, precision, compile_args) + try: + logging.info(comfyui_memory_load(f"post-model-load:wan-vae:{model_name}")) + except Exception: + pass # Attach device info to VAE object for downstream nodes if result and len(result) > 0: @@ -219,7 +235,15 @@ class LoadWanVideoT5TextEncoder: if device == "cpu": setattr(nodes_module, 'offload_device', selected_device) + try: + logging.info(comfyui_memory_load(f"pre-model-load:wan-textenc:{model_name}")) + except Exception: + pass result = original_loader.loadmodel(model_name, precision, load_device, quantization) + try: + logging.info(comfyui_memory_load(f"post-model-load:wan-textenc:{model_name}")) + except Exception: + pass logging.info(f"[MultiGPU] WanVideo T5 Text encoder loaded on {selected_device}") @@ -331,7 +355,15 @@ class LoadWanVideoClipTextEncoder: if device == "cpu": setattr(nodes_module, 'offload_device', selected_device) + try: + logging.info(comfyui_memory_load(f"pre-model-load:wan-clip:{model_name}")) + except Exception: + pass result = original_loader.loadmodel(model_name, precision, load_device) + try: + logging.info(comfyui_memory_load(f"post-model-load:wan-clip:{model_name}")) + except Exception: + pass logging.info(f"[MultiGPU] WanVideo CLIP encoder loaded on {selected_device}")