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
This commit is contained in:
John Pollock
2025-09-20 11:58:29 -05:00
parent 8e4c7fed14
commit 63ff1a4064
5 changed files with 212 additions and 4 deletions
+33 -1
View File
@@ -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:
+116
View File
@@ -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=<TAG>|cpu=<used_GiB>/<total_GiB>|<device>=<used_GiB>/<total_GiB>|..."
- 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)
# ==========================================================================================
+9 -1
View File
@@ -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():
+21 -1
View File
@@ -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
+33 -1
View File
@@ -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}")