From 3121b2f70c3cc956443dc5b17eca71144027ca7f Mon Sep 17 00:00:00 2001 From: John Pollock Date: Tue, 23 Sep 2025 04:41:44 -0500 Subject: [PATCH] feat(mgpu): scoped MM logger; parse compute device/VRAM plan - Introduce MGPU_MM_LOG flag and logger.mgpu_mm_log(...) to gate and prefix MultiGPU Model Management logs (disabled by default) - Replace ad-hoc logger.info("[MultiGPU ...]") calls with mgpu_mm_log in DisTorch2 cache-clearing and delegation paths to reduce noise - In load_models_gpu, parse safetensor allocation strings to infer incoming_compute_device and incoming_compute_planned_bytes (supports hash#device;GB and expert fraction syntax); track required bytes - Remove coarse large-model threshold heuristic in favor of allocation- informed planning Why: centralize and quiet verbose MGPU logs by default, and enable smarter, data-driven device selection and memory planning for multi-GPU model loading. --- __init__.py | 166 ++++++++++++++++++++++++++--------------- checkpoint_multigpu.py | 23 +++--- device_utils.py | 45 ++++++----- distorch_2.py | 73 +++++++++--------- 4 files changed, 171 insertions(+), 136 deletions(-) diff --git a/__init__.py b/__init__.py index bcfadca..3d7b3b0 100644 --- a/__init__.py +++ b/__init__.py @@ -24,12 +24,12 @@ if not logger.handlers: logger.addHandler(handler) logger.setLevel(log_level) -MEMORY_LOG = True +MGPU_MM_LOG = False -def memory_method(self, msg): - if MEMORY_LOG: - self.info(msg) -logger.memory = memory_method.__get__(logger, type(logger)) +def mgpu_mm_log_method(self, msg): + if MGPU_MM_LOG: + self.info(f"[MultiGPU Model Management] {msg}") +logger.mgpu_mm_log = mgpu_mm_log_method.__get__(logger, type(logger)) # Global device state management @@ -242,10 +242,10 @@ def soft_empty_cache_distorch2_patched(force=False): break if is_distorch_active: - logger.info("[MultiGPU Core Patching] DisTorch2 active: clearing caches on all devices") + logger.mgpu_mm_log("DisTorch2 active: clearing caches on all devices") soft_empty_cache_multigpu() else: - logger.info("[MultiGPU Core Patching] DisTorch2 not active: delegating to original mm.soft_empty_cache") + logger.mgpu_mm_log("DisTorch2 not active: delegating to original mm.soft_empty_cache") original_soft_empty_cache(force) mm.soft_empty_cache = soft_empty_cache_distorch2_patched @@ -268,13 +268,15 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch logger.error("[MultiGPU Core Patching] CRITICAL: mm.load_models_gpu 'models' is not a list/tuple/set. Bypassing proactive patch.") return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load) - # Detect incoming large DisTorch2 request + # Detect incoming DisTorch2 request incoming_is_distorch = False incoming_distorch_nonzero = False - incoming_is_large = False incoming_patchers = set() incoming_loaded_names = [] incoming_allowed_devices = None + incoming_compute_device = None + incoming_required_bytes = 0 + incoming_compute_planned_bytes = 0 for lm in models: # Identify ModelPatcher (prefer direct; fall back to .patcher) @@ -297,8 +299,6 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch else: required_bytes = patcher.model_size() - if required_bytes > LARGE_MODEL_THRESHOLD: - incoming_is_large = True else: device_str = "n/a" required_bytes = 0 @@ -339,6 +339,44 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch if not allowed: allowed = {str(patcher.load_device), "cpu"} incoming_allowed_devices = allowed + # Determine compute device and planned bytes from allocation string + alloc = safetensor_allocation_store.get(model_hash, "") + if "#" in alloc: + vram = alloc.split("#", 1)[1] + segs = vram.split(";") + if len(segs) >= 2 and segs[0]: + incoming_compute_device = segs[0].strip() + try: + vvram_gb = float(segs[1]) + incoming_compute_planned_bytes = int(vvram_gb * (1024**3)) + except Exception: + incoming_compute_planned_bytes = 0 + else: + # Expert fractions: "dev,fraction;dev2,fraction2;..." + tokens = [t for t in alloc.split(";") if "," in t] + frac_map = {} + for t in tokens: + dev, frac = t.split(",", 1) + try: + frac_val = float(frac.strip()) + except Exception: + continue + frac_map[dev.strip()] = frac_val + if frac_map: + ld = str(patcher.load_device) + # Prefer the explicit load_device if present and > 0 + target_dev = ld if (ld in frac_map and frac_map[ld] > 0.0) else None + if target_dev is None: + # Otherwise pick highest positive fraction + target_dev = max((d for d,v in frac_map.items() if v > 0.0), key=lambda d: frac_map[d], default=None) + if target_dev is not None: + incoming_compute_device = target_dev + total = mm.get_total_memory(torch.device(target_dev)) + incoming_compute_planned_bytes = int(frac_map[target_dev] * (total or 0)) + if incoming_compute_device is None: + incoming_compute_device = str(patcher.load_device) + if incoming_compute_planned_bytes <= 0: + incoming_compute_planned_bytes = required_bytes # Log informational context with required bytes and device try: @@ -347,68 +385,74 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch model_name = "UnknownModel" incoming_loaded_names.append(f"{model_name}:{required_bytes/(1024**3):.2f}GB req on {device_str}") - logger.info(f"[MultiGPU Core Patching] load_models_gpu incoming set: large={incoming_is_large} distorch2={incoming_is_distorch} count={len(incoming_patchers)}") if incoming_loaded_names: - logger.info(f"[MultiGPU Core Patching] Incoming models summary: {', '.join(incoming_loaded_names)}") + logger.mgpu_mm_log(f"Incoming models summary: {', '.join(incoming_loaded_names)}") if incoming_distorch_nonzero: - logger.info("[MultiGPU Core Patching] Non-Zero incoming DisTorch2 model detected. Initiating proactive unload.") + logger.mgpu_mm_log("Non-Zero incoming DisTorch2 model detected. Initiating proactive unload.") if not hasattr(mm, 'current_loaded_models'): raise AttributeError("comfy.model_management is missing 'current_loaded_models'. Proactive unload check failed.") + needed_patchers = incoming_patchers + # Need-based free on compute device only (scale-aware; core-aligned) + dev_str = incoming_compute_device or (next(iter(incoming_allowed_devices)) if incoming_allowed_devices else None) + freed_bytes = 0 to_unload_indices = [] unload_summaries = [] - needed_patchers = incoming_patchers - logger.info("[MultiGPU Core Patching] Incoming large DisTorch2 model detected. Initiating proactive unload of other large models.") - - # Iterate backwards to safely pop from list - for i in range(len(mm.current_loaded_models) - 1, -1, -1): - lm_cur = mm.current_loaded_models[i] - mp_cur = getattr(lm_cur, 'model', None) - if mp_cur is None: - continue # already dead or cleaned up - - # Skip models needed for this load call - if mp_cur in needed_patchers: - continue - - # Only consider models on compute/donor devices for this DisTorch2 load - if incoming_allowed_devices is not None: - cur_dev_str = str(getattr(lm_cur, "device", "")) - if cur_dev_str not in incoming_allowed_devices: - continue - - # Determine size (prefer LoadedModel.model_memory) - size_cur = 0 - if hasattr(lm_cur, 'model_memory'): - try: - size_cur = lm_cur.model_memory() - except Exception: + if dev_str is not None: + dev_obj = torch.device(dev_str) + free_now = mm.get_free_memory(dev_obj) + try: + free_now_val = free_now[0] if isinstance(free_now, tuple) else free_now + except Exception: + free_now_val = free_now + # Use core-aligned immediate needs: planned vs. memory_required vs. minimum_memory_required + effective_needed = max(incoming_compute_planned_bytes or 0, memory_required or 0, minimum_memory_required or 0) + need_bytes = max(0, effective_needed - (free_now_val or 0)) + logger.mgpu_mm_log(f"Need calc on {dev_str}: effective_needed={effective_needed/(1024**3):.2f}GB, free_now={((free_now_val or 0)/(1024**3)):.2f}GB, need_bytes={need_bytes/(1024**3):.2f}GB") + if need_bytes > 0: + logger.mgpu_mm_log(f"Need-based unload on {dev_str}: need ~{need_bytes/(1024**3):.2f}GB") + # Build candidates on this device only, excluding needed patchers + candidates = [] + for idx, lm_cur in enumerate(mm.current_loaded_models): + mp_cur = getattr(lm_cur, 'model', None) + if mp_cur is None or mp_cur in needed_patchers: + continue + if str(getattr(lm_cur, "device", "")) != dev_str: + continue size_cur = 0 - if size_cur <= 0 and hasattr(mp_cur, 'model_size'): - size_cur = mp_cur.model_size() - - model_name = type(getattr(mp_cur, 'model', mp_cur)).__name__ - logger.info(f"[MultiGPU Core Patching] Unloading large model: {model_name} (~{size_cur/(1024**3):.2f}GB)") - # Attempt full unload; unpatch_weights=True to release distributed allocations - success = False - if hasattr(lm_cur, 'model_unload'): - success = lm_cur.model_unload(memory_to_free=None, unpatch_weights=True) - if success: - to_unload_indices.append(i) - unload_summaries.append(f"{model_name}:{size_cur/(1024**3):.2f}GB") - else: - logger.warning(f"[MultiGPU Core Patching] Failed to fully unload model {model_name} (~{size_cur/(1024**3):.2f}GB)") + if hasattr(lm_cur, 'model_memory'): + try: + size_cur = lm_cur.model_memory() + except Exception: + size_cur = 0 + if size_cur <= 0 and hasattr(mp_cur, 'model_size'): + size_cur = mp_cur.model_size() + candidates.append((size_cur, idx, lm_cur, mp_cur)) + # Sort by size descending + candidates.sort(key=lambda x: x[0], reverse=True) + for size_cur, idx, lm_cur, mp_cur in candidates: + model_name = type(getattr(mp_cur, 'model', mp_cur)).__name__ + logger.mgpu_mm_log(f"Unloading model on {dev_str}: {model_name} (~{size_cur/(1024**3):.2f}GB)") + success = False + if hasattr(lm_cur, 'model_unload'): + success = lm_cur.model_unload(memory_to_free=None, unpatch_weights=True) + if success: + to_unload_indices.append(idx) + unload_summaries.append(f"{model_name}:{size_cur/(1024**3):.2f}GB") + freed_bytes += size_cur + if freed_bytes >= need_bytes: + break # Remove from management list and clear caches unloaded_count = 0 - for idx in to_unload_indices: # already in reverse order + for idx in sorted(to_unload_indices, reverse=True): mm.current_loaded_models.pop(idx) unloaded_count += 1 if unloaded_count > 0: - logger.info(f"[MultiGPU Core Patching] Proactively unloaded {unloaded_count} large model(s): {', '.join(unload_summaries)}") - logger.info("[MultiGPU Core Patching] Performing multi-device cache clear after proactive unload") + logger.mgpu_mm_log(f"Proactively unloaded {unloaded_count} large model(s): {', '.join(unload_summaries)}") + logger.mgpu_mm_log("Performing multi-device cache clear after proactive unload") # Force multi-device cache clear via patched soft_empty_cache (which detects DisTorch2) mm.soft_empty_cache(force=True) else: @@ -425,14 +469,14 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch if free_torch > free_total * 0.25: triggered.append(dev_str) if triggered: - logger.info(f"[MultiGPU Core Patching] No unloads; 25% torch-cache rule triggered on: {', '.join(triggered)}. Calling soft_empty_cache()") + logger.mgpu_mm_log(f"No unloads; 25% torch-cache rule triggered on: {', '.join(triggered)}. Calling soft_empty_cache()") mm.soft_empty_cache(force=True) else: - logger.info("[MultiGPU Core Patching] No unloads; 25% torch-cache rule not met on DisTorch devices; skipping cache clear") + logger.mgpu_mm_log("No unloads; 25% torch-cache rule not met on DisTorch devices; skipping cache clear") else: - logger.info("[MultiGPU Core Patching] No unload candidates matched criteria and either HIGH_VRAM or no DisTorch devices; skipping cache clear") + logger.mgpu_mm_log("No unload candidates matched criteria and either HIGH_VRAM or no DisTorch devices; skipping cache clear") elif incoming_is_distorch: - logger.info("[MultiGPU Core Patching] Incoming DisTorch2 model requires 0.00GB; skipping proactive unload") + logger.mgpu_mm_log("Incoming DisTorch2 model requires 0.00GB; skipping proactive unload") # Continue with original behavior return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load) diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index b4eaa64..42b069b 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -30,10 +30,10 @@ def patch_load_state_dict_guess_config(): global original_load_state_dict_guess_config if original_load_state_dict_guess_config is not None: - logger.info("[MultiGPU] load_state_dict_guess_config is already patched.") + logger.debug("[MultiGPU Checkpoint] load_state_dict_guess_config is already patched.") return - logger.info("[MultiGPU] Patching comfy.sd.load_state_dict_guess_config for advanced MultiGPU loading.") + logger.info("[MultiGPU Core Patching] Patching comfy.sd.load_state_dict_guess_config for advanced MultiGPU loading.") original_load_state_dict_guess_config = comfy.sd.load_state_dict_guess_config comfy.sd.load_state_dict_guess_config = patched_load_state_dict_guess_config @@ -51,9 +51,9 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, if not device_config and not distorch_config: return original_load_state_dict_guess_config(sd, output_vae, output_clip, output_clipvision, embedding_directory, output_model, model_options, te_model_options, metadata) - logger.info("--- [MultiGPU] ENTERING Patched Checkpoint Loader ---") - logger.info(f"Received Device Config: {device_config}") - logger.info(f"Received DisTorch2 Config: {distorch_config}") + logger.debug("[MultiGPU Checkpoint] ENTERING Patched Checkpoint Loader") + logger.debug(f"[MultiGPU Checkpoint] Received Device Config: {device_config}") + logger.debug(f"[MultiGPU Checkpoint] Received DisTorch2 Config: {distorch_config}") clip = None clipvision = None @@ -63,7 +63,6 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, original_main_device = current_device original_clip_device = current_text_encoder_device - logger.info(f"Saved original device contexts: UNet/VAE='{original_main_device}', CLIP='{original_clip_device}'") try: diffusion_model_prefix = comfy.model_detection.unet_prefix_from_state_dict(sd) @@ -80,7 +79,7 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, return None return (diffusion_model, None, VAE(sd={}), None) - logger.info(f"[MultiGPU] Detected Model Config: {type(model_config).__name__}, Parameters: {parameters/10**9:.2f}B") + logger.debug(f"[MultiGPU] Detected Model Config: {type(model_config).__name__}, Parameters: {parameters/10**9:.2f}B") unet_weight_dtype = list(model_config.supported_inference_dtypes) if model_config.scaled_fp8 is not None: @@ -109,7 +108,7 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, 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") + logger.mgpu_mm_log("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()) multigpu_memory_log(f"unet:{config_hash[:8]}", "post-model") @@ -121,7 +120,7 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, safetensor_settings_store[model_hash] = distorch_config.get('unet_settings','') model.is_distorch = True model._distorch_high_precision_loras = distorch_config.get('high_precision_loras', True) - logger.info(f"Stored DisTorch2 config for UNet (hash {model_hash[:8]}): {distorch_config['unet_allocation']}") + logger.mgpu_mm_log(f"Stored DisTorch2 config for UNet (hash {model_hash[:8]}): {distorch_config['unet_allocation']}") model.load_model_weights(sd, diffusion_model_prefix) multigpu_memory_log(f"unet:{config_hash[:8]}", "post-weights") @@ -144,7 +143,7 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, if clip_target is not None: 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") + logger.debug("[MultiGPU Checkpoint] Invoking soft_empty_cache_multigpu before CLIP construction") multigpu_memory_log(f"clip:{config_hash[:8]}", "pre-load") soft_empty_cache_multigpu() clip_params = comfy.utils.calculate_parameters(clip_sd) @@ -171,16 +170,12 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, logger.warning("CLIP target not found in model config.") finally: - # --- Restore original device contexts and clean up --- set_current_device(original_main_device) set_current_text_encoder_device(original_clip_device) if config_hash in checkpoint_device_config: del checkpoint_device_config[config_hash] if config_hash in checkpoint_distorch_config: del checkpoint_distorch_config[config_hash] - logger.info(f"Restored original device contexts. UNet/VAE='{original_main_device}', CLIP='{original_clip_device}'") - logger.info("--- [MultiGPU] EXITING Patched Checkpoint Loader ---") - return (model_patcher, clip, vae, clipvision) class CheckpointLoaderAdvancedMultiGPU: diff --git a/device_utils.py b/device_utils.py index 7996ac4..95f14e8 100644 --- a/device_utils.py +++ b/device_utils.py @@ -122,7 +122,7 @@ def get_device_list(): _DEVICE_LIST_CACHE = devs # Log only once when initially populated - logger.info(f"[MultiGPU_Device_Utils] Device list initialized: {devs}") + logger.debug(f"[MultiGPU_Device_Utils] Device list initialized: {devs}") return devs @@ -242,17 +242,17 @@ def soft_empty_cache_multigpu(): """ import gc - logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: starting GC and multi-device cache clear") + logger.mgpu_mm_log("soft_empty_cache_multigpu: starting GC and multi-device cache clear") # Record pre-GC snapshot for general system view multigpu_memory_log("general", "pre-soft-empty") # Python GC (same as all implementations) gc.collect() - logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: garbage collection complete") + logger.mgpu_mm_log("soft_empty_cache_multigpu: garbage collection complete") # Clear cache for ALL devices (not just ComfyUI's single device) all_devices = get_device_list() - logger.info(f"[MultiGPU_Device_Utils] soft_empty_cache_multigpu: devices to clear = {all_devices}") + logger.mgpu_mm_log(f"soft_empty_cache_multigpu: devices to clear = {all_devices}") # Check global availability first to avoid unnecessary iteration if backend is missing is_cuda_available = hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available() @@ -262,42 +262,42 @@ def soft_empty_cache_multigpu(): if is_cuda_available: 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})") + logger.mgpu_mm_log(f"Clearing CUDA cache on {device_str} (idx={device_idx})") 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}") + logger.mgpu_mm_log(f"Cleared CUDA cache (and IPC if available) on {device_str}") elif device_str == "mps": if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"): - logger.info("[MultiGPU_Device_Utils] Clearing MPS cache") + logger.mgpu_mm_log("Clearing MPS cache") torch.mps.empty_cache() - logger.info("[MultiGPU_Device_Utils] Cleared MPS cache") + logger.mgpu_mm_log("Cleared MPS cache") 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}") + logger.mgpu_mm_log(f"Clearing XPU cache on {device_str}") torch.xpu.empty_cache() - logger.info(f"[MultiGPU_Device_Utils] Cleared XPU cache on {device_str}") + logger.mgpu_mm_log(f"Cleared XPU cache on {device_str}") 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}") + logger.mgpu_mm_log(f"Clearing NPU cache on {device_str}") torch.npu.empty_cache() - logger.info(f"[MultiGPU_Device_Utils] Cleared NPU cache on {device_str}") + logger.mgpu_mm_log(f"Cleared NPU cache on {device_str}") 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}") + logger.mgpu_mm_log(f"Clearing MLU cache on {device_str}") torch.mlu.empty_cache() - logger.info(f"[MultiGPU_Device_Utils] Cleared MLU cache on {device_str}") + logger.mgpu_mm_log(f"Cleared MLU cache on {device_str}") 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}") + logger.mgpu_mm_log(f"Clearing CoreX cache on {device_str}") torch.corex.empty_cache() - logger.info(f"[MultiGPU_Device_Utils] Cleared CoreX cache on {device_str}") + logger.mgpu_mm_log(f"Cleared CoreX cache on {device_str}") # Record post-GC snapshot for general system view multigpu_memory_log("general", "post-soft-empty") @@ -402,14 +402,14 @@ def memory_print_summary(log: logging.Logger = logger): 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 + from . import 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} ===") + logger.mgpu_mm_log(f"=== memory summary: {identifier} ===") for ts, tag, snap in series: # Build device list (cpu first, then sorted devices) parts = [] @@ -422,7 +422,7 @@ def memory_print_summary(log: logging.Logger = logger): 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)) + logger.mgpu_mm_log(f"{ts_str} {identifier} {tag} | " + " | ".join(parts)) def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger): @@ -462,7 +462,7 @@ def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger) 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)) + logger.mgpu_mm_log(f"{identifier} {tag} - {prev_tag}: " + " | ".join(parts)) else: # Baseline vs zero keys = set(curr.keys()) @@ -471,10 +471,7 @@ def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger) 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)}") + logger.mgpu_mm_log(f"{identifier} {tag} - : " + " | ".join(parts)) # Update last snapshot _MEM_SNAPSHOT_LAST[identifier] = (tag, curr) diff --git a/distorch_2.py b/distorch_2.py index 9e121b0..c92a439 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -48,7 +48,7 @@ def create_safetensor_model_hash(model, caller): final_hash = hashlib.sha256(identifier.encode()).hexdigest() # DEBUG STATEMENT - ALWAYS LOG THE HASH - logger.debug(f"[MultiGPU_DisTorch2] Created hash for {caller}: {final_hash[:8]}...") + logger.debug(f"[MultiGPU DisTorch V2] Created hash for {caller}: {final_hash[:8]}...") return final_hash @@ -81,7 +81,7 @@ def register_patched_safetensor_modelpatcher(): unpatch_weights = self.model.current_weight_patches_uuid is not None and (self.model.current_weight_patches_uuid != self.patches_uuid or force_patch_weights) if unpatch_weights: - logger.info(f"[MultiGPU_DisTorch2] Patches changed or forced. Unpatching model.") + logger.debug(f"[MultiGPU DisTorch V2] Patches changed or forced. Unpatching model.") self.unpatch_model(self.offload_device, unpatch_weights=True) self.patch_model(load_weights=False) @@ -90,10 +90,10 @@ def register_patched_safetensor_modelpatcher(): is_clip_model = getattr(self, 'is_clip', False) if is_clip_model: - logger.info(f"[MultiGPU_DisTorch2] Using CLIP-specific allocation for model {debug_hash[:8]} (HEAD PRESERVATION ENABLED)") + logger.debug(f"[MultiGPU DisTorch V2] Using CLIP-specific allocation for model {debug_hash[:8]} (HEAD PRESERVATION ENABLED)") device_assignments = analyze_safetensor_loading_clip(self, allocations) else: - logger.debug(f"[MultiGPU_DisTorch2] Using standard allocation for model {debug_hash[:8]} (UNET/VAE - UNTOUCHED)") + logger.debug(f"[MultiGPU DisTorch V2] Using standard allocation for model {debug_hash[:8]} (UNET/VAE - UNTOUCHED)") device_assignments = analyze_safetensor_loading(self, allocations) model_original_dtype = comfy.utils.weight_dtype(self.model.state_dict()) @@ -111,7 +111,7 @@ def register_patched_safetensor_modelpatcher(): pass if current_module_device is not None and str(current_module_device) != str(block_target_device): - logger.debug(f"[MultiGPU_DisTorch2] Moving already patched {module_name} to {block_target_device}") + logger.debug(f"[MultiGPU DisTorch V2] Moving already patched {module_name} to {block_target_device}") module_object.to(block_target_device) mem_counter += module_size @@ -145,11 +145,11 @@ def register_patched_safetensor_modelpatcher(): new_param = torch.nn.Parameter(cast_data.to(torch.float8_e4m3fn)) new_param.requires_grad = param.requires_grad setattr(module_object, param_name, new_param) - logger.debug(f"[MultiGPU_DisTorch2] Cast {module_name}.{param_name} to FP8 for CPU storage") + logger.debug(f"[MultiGPU DisTorch V2] Cast {module_name}.{param_name} to FP8 for CPU storage") # Step 4: Move to ultimate destination based on DisTorch assignment if block_target_device != device_to: - logger.debug(f"[MultiGPU_DisTorch2] Moving {module_name} from {device_to} to {block_target_device}") + logger.debug(f"[MultiGPU DisTorch V2] Moving {module_name} from {device_to} to {block_target_device}") module_object.to(block_target_device) module_object.comfy_cast_weights = True @@ -159,7 +159,8 @@ 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") + logger.info("[MultiGPU DisTorch V2] DisTorch loading completed.") + logger.info(f"[MultiGPU DisTorch V2] Total memory: {mem_counter / (1024 * 1024):.2f}MB") multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "post-load") return 0 @@ -167,7 +168,7 @@ def register_patched_safetensor_modelpatcher(): comfy.model_patcher.ModelPatcher.partially_load = new_partially_load comfy.model_patcher.ModelPatcher._distorch_patched = True - logger.info("[MultiGPU_DisTorch2] Successfully patched ModelPatcher.partially_load") + logger.info("[MultiGPU Core Patching] Successfully patched ModelPatcher.partially_load") def analyze_safetensor_loading(model_patcher, allocations_string): @@ -183,13 +184,10 @@ def analyze_safetensor_loading(model_patcher, allocations_string): distorch_alloc, virtual_vram_str = allocations_string.split('#') compute_device = virtual_vram_str.split(';')[0] - logger.info(f"[MultiGPU_DisTorch2] Compute Device: {compute_device}") + logger.debug(f"[MultiGPU DisTorch V2] Compute Device: {compute_device}") if not distorch_alloc: mode = "fraction" - logger.info("[MultiGPU_DisTorch2] Expert String Examples:") - logger.info(" Direct(byte) Mode - cuda:0,500mb;cuda:1,3.0g;cpu,5gb* -> '*' cpu = over/underflow device, put 0.50gb on cuda0, 3.00gb on cuda1, and 5.00gb (or the rest) on cpu") - logger.info(" Ratio(%) Mode - cuda:0,8%;cuda:1,8%;cpu,4% -> 8:8:4 ratio, put 40% on cuda0, 40% on cuda1, and 20% on cpu") distorch_alloc = calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str) elif any(c in distorch_alloc.lower() for c in ['g', 'm', 'k', 'b']): @@ -205,12 +203,13 @@ def analyze_safetensor_loading(model_patcher, allocations_string): if device not in present_devices: distorch_alloc += f";{device},0.0" - logger.info(f"[MultiGPU_DisTorch2] Final Allocation String: {distorch_alloc}") - eq_line = "=" * 50 dash_line = "-" * 50 fmt_assign = "{:<18}{:>7}{:>14}{:>10}" + logger.info(eq_line) + logger.info(f"[MultiGPU DisTorch V2] Final Allocation String:\n{distorch_alloc}") + for allocation in distorch_alloc.split(';'): if ',' not in allocation: continue @@ -264,8 +263,8 @@ def analyze_safetensor_loading(model_patcher, allocations_string): total_memory = sum(module_size for module_size, _, _, _ in raw_block_list) MIN_BLOCK_THRESHOLD = total_memory * 0.0001 - logger.debug(f"[MultiGPU_DisTorch2] Total model memory: {total_memory} bytes") - logger.debug(f"[MultiGPU_DisTorch2] Tiny block threshold (0.01%): {MIN_BLOCK_THRESHOLD} bytes") + logger.debug(f"[MultiGPU DisTorch V2] Total model memory: {total_memory} bytes") + logger.debug(f"[MultiGPU DisTorch V2] Tiny block threshold (0.01%): {MIN_BLOCK_THRESHOLD} bytes") all_blocks = [] for module_size, module_name, module_object, params in raw_block_list: @@ -278,9 +277,9 @@ def analyze_safetensor_loading(model_patcher, allocations_string): 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] - logger.debug(f"[MultiGPU_DisTorch2] Total blocks: {len(all_blocks)}") - logger.debug(f"[MultiGPU_DisTorch2] Distributable blocks: {len(block_list)}") - logger.debug(f"[MultiGPU_DisTorch2] Tiny blocks (<0.01%): {len(tiny_block_list)}") + logger.debug(f"[MultiGPU DisTorch V2] Total blocks: {len(all_blocks)}") + logger.debug(f"[MultiGPU DisTorch V2] Distributable blocks: {len(block_list)}") + logger.debug(f"[MultiGPU DisTorch V2] Tiny blocks (<0.01%): {len(tiny_block_list)}") logger.info(" DisTorch2 Model Layer Distribution") logger.info(dash_line) @@ -343,7 +342,7 @@ def analyze_safetensor_loading(model_patcher, allocations_string): tiny_mem_percent = (tiny_block_memory / total_memory) * 100 if total_memory > 0 else 0 device_label = f"{compute_device} (<0.01%)" logger.info(fmt_assign.format(device_label, str(len(tiny_block_list)), f"{tiny_mem_mb:.2f}", f"{tiny_mem_percent:.1f}%")) - logger.debug(f"[MultiGPU_DisTorch2] Tiny block memory breakdown: {tiny_block_memory} bytes ({tiny_mem_mb:.2f} MB), which is {tiny_mem_percent:.4f}% of total model memory.") + logger.debug(f"[MultiGPU DisTorch V2] Tiny block memory breakdown: {tiny_block_memory} bytes ({tiny_mem_mb:.2f} MB), which is {tiny_mem_percent:.4f}% of total model memory.") total_assigned_memory = 0 device_memories = {} @@ -415,7 +414,7 @@ def analyze_safetensor_loading_clip(model_patcher, allocations_string): if device not in present_devices: distorch_alloc += f";{device},0.0" - logger.info(f"[MultiGPU_DisTorch2_CLIP] Final CLIP Allocation String: {distorch_alloc}") + logger.info(f"[MultiGPU_DisTorch2_CLIP] Final CLIP Allocation String:\n{distorch_alloc}") eq_line = "=" * 50 dash_line = "-" * 50 @@ -640,16 +639,16 @@ def calculate_fraction_from_byte_expert_string(model_patcher, byte_str): if bytes_to_assign > 0: final_byte_allocations[dev] = bytes_to_assign remaining_model_bytes -= bytes_to_assign - logger.info(f"[MultiGPU_DisTorch2] Assigning {bytes_to_assign / (1024**2):.2f}MB of model to {dev} (requested {requested_bytes / (1024**2):.2f}MB).") + logger.info(f"[MultiGPU DisTorch V2] Assigning {bytes_to_assign / (1024**2):.2f}MB of model to {dev} (requested {requested_bytes / (1024**2):.2f}MB).") if remaining_model_bytes <= 0: - logger.info("[MultiGPU_DisTorch2] All model blocks have been allocated. Subsequent devices in the string will receive no assignment.") + logger.info("[MultiGPU DisTorch V2] All model blocks have been allocated. Subsequent devices in the string will receive no assignment.") break # Assign any leftover model bytes to the wildcard device if remaining_model_bytes > 0: final_byte_allocations[wildcard_device] += remaining_model_bytes - logger.info(f"[MultiGPU_DisTorch2] Assigning remaining {remaining_model_bytes / (1024**2):.2f}MB of model to wildcard device '{wildcard_device}'.") + logger.info(f"[MultiGPU DisTorch V2] Assigning remaining {remaining_model_bytes / (1024**2):.2f}MB of model to wildcard device '{wildcard_device}'.") # Convert the final byte allocations to VRAM fractions allocation_parts = [] @@ -707,7 +706,7 @@ def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str): else: put_part = ", ".join(put_parts[:-1]) + f", and {put_parts[-1]}" - logger.info(f"[MultiGPU_DisTorch2] Ratio(%) Mode - {ratio_str} -> {ratio_string} ratio, put {put_part}") + logger.info(f"[MultiGPU DisTorch V2] Ratio(%) Mode - {ratio_str} -> {ratio_string} ratio, put {put_part}") allocations_string = ";".join(allocation_parts) @@ -775,8 +774,8 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str): # Warning if model too large if model_size_gb > (recipient_vram * 0.9): required_offload_gb = model_size_gb - (recipient_vram * 0.9) - logger.warning(f"[MultiGPU] WARNING: Model size ({model_size_gb:.2f}GB) is larger than 90% of available VRAM on {recipient_device} ({recipient_vram * 0.9:.2f}GB).") - logger.warning(f"[MultiGPU] To prevent an OOM error, set 'virtual_vram_gb' to at least {required_offload_gb:.2f}.") + logger.warning(f"\n\n[MultiGPU DisTorch V2] Model size ({model_size_gb:.2f}GB) is larger than 90% of available VRAM on: {recipient_device} ({recipient_vram * 0.9:.2f}GB).") + logger.warning(f"[MultiGPU DisTorch V2] To prevent an OOM error, set 'virtual_vram_gb' to at least {required_offload_gb:.2f}.\n\n") new_on_recipient = max(0, model_size_gb - virtual_vram_gb) @@ -854,9 +853,9 @@ def override_class_with_distorch_safetensor_v2(cls): last_settings_hash = safetensor_settings_store.get(model_hash) if last_settings_hash != settings_hash: - logger.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") + logger.mgpu_mm_log(f"Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") else: - logger.info(f"[MultiGPU_DisTorch2] Settings unchanged for model {model_hash[:8]}. Using cached model.") + logger.mgpu_mm_log(f"Settings unchanged for model {model_hash[:8]}. Using cached model.") out = fn(*args, **kwargs) @@ -874,7 +873,7 @@ def override_class_with_distorch_safetensor_v2(cls): full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" - logger.info(f"[MultiGPU_DisTorch2] Full allocation string: {full_allocation}") + logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") if hasattr(out[0], 'model'): model_hash = create_safetensor_model_hash(out[0], "override") @@ -953,9 +952,9 @@ def override_class_with_distorch_safetensor_v2_clip(cls): last_settings_hash = safetensor_settings_store.get(model_hash) if last_settings_hash != settings_hash: - logger.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") + logger.mgpu_mm_log(f"Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") else: - logger.info(f"[MultiGPU_DisTorch2] Settings unchanged for model {model_hash[:8]}. Using cached model.") + logger.mgpu_mm_log(f"Settings unchanged for model {model_hash[:8]}. Using cached model.") out = fn(*args, **kwargs) @@ -973,7 +972,7 @@ def override_class_with_distorch_safetensor_v2_clip(cls): full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" - logger.info(f"[MultiGPU_DisTorch2] Full allocation string: {full_allocation}") + logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") if hasattr(out[0], 'model'): model_hash = create_safetensor_model_hash(out[0], "override") @@ -1049,9 +1048,9 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): last_settings_hash = safetensor_settings_store.get(model_hash) if last_settings_hash != settings_hash: - logger.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") + logger.info(f"[MultiGPU DisTorch V2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.") else: - logger.info(f"[MultiGPU_DisTorch2] Settings unchanged for model {model_hash[:8]}. Using cached model.") + logger.info(f"[MultiGPU DisTorch V2] Settings unchanged for model {model_hash[:8]}. Using cached model.") out = fn(*args, **kwargs) @@ -1069,7 +1068,7 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" - logger.info(f"[MultiGPU_DisTorch2] Full allocation string: {full_allocation}") + logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") if hasattr(out[0], 'model'): model_hash = create_safetensor_model_hash(out[0], "override")