From 55a0d22b01c93308b88a4a285fc1dddb3444b8ef Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sun, 21 Sep 2025 06:12:20 -0500 Subject: [PATCH] refactor: simplify memory logging in checkpoint loading Replace try-catch wrapped comfyui_memory_load calls with streamlined multigpu_memory_log function. Removes exception handling overhead and uses consistent config hash identifiers for UNet, VAE, and CLIP model loading phases. --- checkpoint_multigpu.py | 40 ++----- device_utils.py | 234 ++++++++++++++++++++++++++--------------- distorch.py | 14 +-- distorch_2.py | 25 +---- 4 files changed, 166 insertions(+), 147 deletions(-) diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index f1365c1..b4eaa64 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, comfyui_memory_load, create_model_identifier +from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher logger = logging.getLogger("MultiGPU") @@ -105,21 +105,14 @@ 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 + multigpu_memory_log(f"unet:{config_hash[:8]}", "pre-load") 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 + multigpu_memory_log(f"unet:{config_hash[:8]}", "post-model") if distorch_config and 'unet_allocation' in distorch_config: register_patched_safetensor_modelpatcher() @@ -131,27 +124,17 @@ 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 + multigpu_memory_log(f"unet:{config_hash[:8]}", "post-weights") 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 + multigpu_memory_log(f"vae:{config_hash[:8]}", "pre-load") 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 + multigpu_memory_log(f"vae:{config_hash[:8]}", "post-load") if output_clip: clip_target_device = device_config.get('clip_device', original_clip_device) @@ -162,10 +145,7 @@ 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 + multigpu_memory_log(f"clip:{config_hash[:8]}", "pre-load") 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) @@ -184,11 +164,7 @@ 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 + multigpu_memory_log(f"clip:{config_hash[:8]}", "post-load") else: logger.warning("No CLIP/text encoder weights in checkpoint.") else: diff --git a/device_utils.py b/device_utils.py index 966f399..7996ac4 100644 --- a/device_utils.py +++ b/device_utils.py @@ -243,19 +243,11 @@ 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 + # Record pre-GC snapshot for general system view + multigpu_memory_log("general", "pre-soft-empty") # 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) @@ -271,103 +263,50 @@ 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 + # Record post-GC snapshot for general system view + multigpu_memory_log("general", "post-soft-empty") + 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 + return float(b) / (1024.0 ** 3) def comfyui_memory_load(tag: str) -> str: @@ -392,28 +331,155 @@ def comfyui_memory_load(tag: str) -> str: # 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)) + 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 + 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}") return "|".join(segments) +# ========================================================================================== +# Delta-capable memory logging (identifier + tag) with timestamped series +# ========================================================================================== + +from datetime import datetime, timezone + +# Stores the last snapshot per identifier: identifier -> (last_tag, snapshot_map) +# snapshot_map: device_str -> (used_bytes, total_bytes) +_MEM_SNAPSHOT_LAST = {} + +# Full chronological series per identifier: identifier -> list[(timestamp, tag, snapshot_map)] +_MEM_SNAPSHOT_SERIES = {} + + +def _capture_memory_snapshot() -> dict[str, tuple[int, int]]: + """ + Capture an absolute memory snapshot for CPU and all non-CPU devices. + Values are returned in bytes (used, total) for each device string key. + """ + snapshot: dict[str, tuple[int, int]] = {} + + # CPU + vm = psutil.virtual_memory() + snapshot["cpu"] = (vm.used, vm.total) + + # Non-CPU devices + devices = [d for d in get_device_list() if d != "cpu"] + for dev_str in devices: + device = torch.device(dev_str) + total = mm.get_total_memory(device) + free_info = mm.get_free_memory(device, torch_free_too=True) + system_free = free_info[0] if isinstance(free_info, tuple) else free_info + used = max(0, (total or 0) - (system_free or 0)) + snapshot[dev_str] = (used, total or 0) + + return snapshot + + +def _format_delta_gib(delta_bytes: int) -> str: + """Format a signed GiB delta with two decimals.""" + gib = _bytes_to_gib(abs(delta_bytes)) + sign = "+" if delta_bytes >= 0 else "-" + return f"{sign}{gib:.2f}" + + +def memory_print_summary(log: logging.Logger = logger): + """ + Print the entire run as absolute actuals with timestamps for each identifier. + One line per recorded snapshot in insertion order. + Format: + YYYY-MM-DDTHH:MM:SS.mmmZ identifier tag | cpu=U/T | cuda:0=U/T | ... + (GiB values, two decimals) + """ + from . import logger as mgpu_logger + + # Stable identifier order for readability + for identifier in sorted(_MEM_SNAPSHOT_SERIES.keys()): + series = _MEM_SNAPSHOT_SERIES[identifier] + if not series: + continue + mgpu_logger.memory(f"=== memory summary: {identifier} ===") + for ts, tag, snap in series: + # Build device list (cpu first, then sorted devices) + parts = [] + # CPU + cpu_used, cpu_total = snap.get("cpu", (0, 0)) + parts.append(f"cpu={_bytes_to_gib(cpu_used):.2f}/{_bytes_to_gib(cpu_total):.2f}") + # Non-CPU (sorted) + devs = sorted([k for k in snap.keys() if k != "cpu"]) + for dev in devs: + used, total = snap[dev] + parts.append(f"{dev}={_bytes_to_gib(used):.2f}/{_bytes_to_gib(total):.2f}") + ts_str = ts.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z" + mgpu_logger.memory(f"{ts_str} {identifier} {tag} | " + " | ".join(parts)) + + +def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger): + """ + Record a timestamped memory snapshot for the given identifier and tag. + - INFO: per-device deltas vs. previous snapshot for the same identifier (GiB, signed, no totals). + - DEBUG: absolute snapshot string via comfyui_memory_load(tag) prefixed by identifier. + - Special identifier: 'print_summary' will dump the entire series as actuals with timestamps. + """ + from . import logger as mgpu_logger + + if identifier == "print_summary": + memory_print_summary(log=log) + return + + # Capture current snapshot and timestamp + ts = datetime.now(timezone.utc) + curr = _capture_memory_snapshot() + + # Append to full series + series = _MEM_SNAPSHOT_SERIES.get(identifier) + if series is None: + series = [] + _MEM_SNAPSHOT_SERIES[identifier] = series + series.append((ts, tag, curr)) + + # Compute and log delta vs last + if identifier in _MEM_SNAPSHOT_LAST: + prev_tag, prev = _MEM_SNAPSHOT_LAST[identifier] + # Union of device keys + keys = set(prev.keys()) | set(curr.keys()) + # Stable order: cpu first, then sorted devices + ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"]) + parts = [] + for k in ordered: + p_used, _p_tot = prev.get(k, (0, curr.get(k, (0, 0))[1])) + c_used, _c_tot = curr.get(k, (0, prev.get(k, (0, 0))[1])) + delta = c_used - p_used + parts.append(f"{k}={_format_delta_gib(delta)}") + mgpu_logger.memory(f"{identifier} {tag} - {prev_tag}: " + " | ".join(parts)) + else: + # Baseline vs zero + keys = set(curr.keys()) + ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"]) + parts = [] + for k in ordered: + c_used, _c_tot = curr.get(k, (0, 0)) + parts.append(f"{k}=+{_bytes_to_gib(c_used):.2f}") + mgpu_logger.memory(f"{identifier} {tag} - : " + " | ".join(parts)) + + # DEBUG absolute + mgpu_logger.memory(f"{identifier}, {comfyui_memory_load(tag)}") + + # Update last snapshot + _MEM_SNAPSHOT_LAST[identifier] = (tag, curr) + + # ========================================================================================== # Model Management Inspection Utilities (End-to-End Tracking) # ========================================================================================== diff --git a/distorch.py b/distorch.py index 897090d..ab81489 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, comfyui_memory_load +from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log # Global store for model allocations model_allocation_store = {} @@ -41,16 +41,10 @@ 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 + multigpu_memory_log(f"gguf:{debug_hash[:8]}", "pre-load") + super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs) + multigpu_memory_log(f"gguf:{debug_hash[:8]}", "post-load") linked = [] module_count = 0 for n, m in self.model.named_modules(): diff --git a/distorch_2.py b/distorch_2.py index 922fccb..9e121b0 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, comfyui_memory_load +from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log safetensor_allocation_store = {} safetensor_settings_store = {} @@ -64,19 +64,13 @@ 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 + multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "pre-load") 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 + multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "post-load") if hasattr(self, '_distorch_block_assignments'): del self._distorch_block_assignments return result @@ -88,15 +82,7 @@ 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) @@ -174,10 +160,7 @@ 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 + multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "post-load") return 0