diff --git a/__init__.py b/__init__.py index 37c837f..5e250d7 100644 --- a/__init__.py +++ b/__init__.py @@ -6,7 +6,7 @@ from pathlib import Path import folder_paths import comfy.model_management as mm from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS -from .device_utils import get_device_list, is_accelerator_available +from .device_utils import get_device_list, is_accelerator_available, soft_empty_cache_multigpu # --- DisTorch V2 Logging Configuration --- # Set to "E" for Engineering (DEBUG) or "P" for Production (INFO) @@ -215,6 +215,170 @@ from .distorch_2 import ( override_class_with_distorch_safetensor_v2_clip_no_device ) +# ========================================================================================== +# Core Patching: soft_empty_cache harmonization for DisTorch2 +# ========================================================================================== +logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for DisTorch2 harmonization") + +# Store the original function for fallback behavior +original_soft_empty_cache = mm.soft_empty_cache + +def soft_empty_cache_distorch2_patched(force=False): + """ + Patched mm.soft_empty_cache. If DisTorch2 models are active, clear cache on ALL devices. + Otherwise, execute original ComfyUI behavior. + """ + is_distorch_active = False + + # Check if any loaded model is managed by DisTorch2 using the allocation store + for lm in mm.current_loaded_models: + mp = lm.model # weakref call to ModelPatcher + if mp is not None: + model_hash = create_safetensor_model_hash(mp, "cache_patch_check") + if model_hash in safetensor_allocation_store and safetensor_allocation_store[model_hash]: + is_distorch_active = True + break + + if is_distorch_active: + logger.info("[MultiGPU Core Patching] 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") + original_soft_empty_cache(force) + +# Apply the patch +mm.soft_empty_cache = soft_empty_cache_distorch2_patched + +# ========================================================================================== +# Core Patching: load_models_gpu Proactive Unloading (NEW FIX for UNet OOM) +# Prevents OOM on offload/donor devices when swapping large DisTorch2 models. +# ========================================================================================== + +LARGE_MODEL_THRESHOLD = 2 * (1024**3) # 2 GB threshold for "large" models + +# Patch only once (handles reloads) +if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch2_proactive_patched"): + logger.info("[MultiGPU Core Patching] Patching mm.load_models_gpu for DisTorch2 proactive unloading") + + original_load_models_gpu = mm.load_models_gpu + + def patched_load_models_gpu(models, memory_required=0, force_patch_weights=False, minimum_memory_required=None, force_full_load=False): + """ + Proactively unload large models that are not needed when loading a large DisTorch2 model. + This frees both compute and donor device memory ahead of ComfyUI's compute-only check. + """ + # Validate models argument loudly + if not isinstance(models, (list, tuple, set)): + 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 + incoming_is_distorch = False + incoming_is_large = False + incoming_patchers = set() + incoming_loaded_names = [] + + for lm in models: + # Expect LoadedModel instances; gather ModelPatcher if alive + mp = getattr(lm, 'model', None) + if mp is not None: + incoming_patchers.add(mp) + # Determine size (prefer LoadedModel.model_memory if available) + size_bytes = 0 + if hasattr(lm, 'model_memory'): + try: + size_bytes = lm.model_memory() + except Exception: + size_bytes = 0 + if size_bytes <= 0 and hasattr(mp, 'model_size'): + size_bytes = mp.model_size() + + if size_bytes > LARGE_MODEL_THRESHOLD: + incoming_is_large = True + + # Check DisTorch2 management via allocation store + model_hash = create_safetensor_model_hash(mp, "load_patch_check") + if model_hash in safetensor_allocation_store and safetensor_allocation_store.get(model_hash): + incoming_is_distorch = True + + # Log informational context + incoming_loaded_names.append(f"{type(getattr(mp, 'model', mp)).__name__}:{size_bytes/(1024**3):.2f}GB") + + 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)}") + + # Proactive unload if both conditions are met + if incoming_is_distorch and incoming_is_large: + if not hasattr(mm, 'current_loaded_models'): + raise AttributeError("comfy.model_management is missing 'current_loaded_models'. Proactive unload check failed.") + + 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 + + # Determine size (prefer LoadedModel.model_memory) + size_cur = 0 + 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() + + # Only unload large models + if size_cur > LARGE_MODEL_THRESHOLD: + 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)") + + # Remove from management list and clear caches + unloaded_count = 0 + for idx in to_unload_indices: # already in reverse order + 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") + # Force multi-device cache clear via patched soft_empty_cache (which detects DisTorch2) + mm.soft_empty_cache(force=True) + else: + logger.info("[MultiGPU Core Patching] No unload candidates matched the criteria (either none large or all required)") + + # Continue with original behavior + return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load) + + # Mark and apply the patch + patched_load_models_gpu._distorch2_proactive_patched = True + mm.load_models_gpu = patched_load_models_gpu +else: + if not hasattr(mm, 'load_models_gpu'): + raise AttributeError("comfy.model_management is missing 'load_models_gpu'. Core patching failed.") + else: + logger.debug("[MultiGPU Core Patching] mm.load_models_gpu already patched; skipping") + # Import advanced checkpoint loaders from .checkpoint_multigpu import ( CheckpointLoaderAdvancedMultiGPU, diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index 681c258..663ccef 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -107,7 +107,8 @@ 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) - soft_empty_cache_multigpu(logger) + 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()) if distorch_config and 'unet_allocation' in distorch_config: @@ -137,7 +138,8 @@ 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: - soft_empty_cache_multigpu(logger) + logger.info("[MultiGPU Checkpoint] Invoking soft_empty_cache_multigpu before CLIP construction") + 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) diff --git a/device_utils.py b/device_utils.py index b2ed027..5c96f0c 100644 --- a/device_utils.py +++ b/device_utils.py @@ -234,44 +234,68 @@ def parse_device_string(device_string): return device_string, None -def soft_empty_cache_multigpu(logger): +def soft_empty_cache_multigpu(): """ Replicate ComfyUI's cache clearing but for ALL devices in MultiGPU. MultiGPU adaptation of ComfyUI's soft_empty_cache() functionality. + Uses context managers to ensure the calling thread's device context is restored. """ import gc - logger.info("[MultiGPU_Device_Utils] Preparing devices for optimized safetensor loading") + logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: starting GC and multi-device cache clear") # Python GC (same as all implementations) gc.collect() - logger.debug("[MultiGPU_Device_Utils] Performed garbage collection before safetensor loading") + logger.info("[MultiGPU_Device_Utils] 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}") + + # 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() for device_str in all_devices: if device_str.startswith("cuda:"): - device_idx = int(device_str.split(":")[1]) - torch.cuda.set_device(device_idx) - torch.cuda.empty_cache() - torch.cuda.ipc_collect() # ComfyUI's CUDA optimization - logger.debug(f"[MultiGPU_Device_Utils] Cleared cache + IPC for {device_str}") + 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})") + 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}") + elif device_str == "mps": - torch.mps.empty_cache() - logger.debug("[MultiGPU_Device_Utils] Cleared cache for MPS") + if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"): + logger.info("[MultiGPU_Device_Utils] Clearing MPS cache") + torch.mps.empty_cache() + logger.info("[MultiGPU_Device_Utils] Cleared MPS cache") + elif device_str.startswith("xpu:"): - torch.xpu.empty_cache() - logger.debug("[MultiGPU_Device_Utils] Cleared cache for Intel XPU") + if hasattr(torch, "xpu") and hasattr(torch.xpu, "empty_cache"): + logger.info(f"[MultiGPU_Device_Utils] Clearing XPU cache on {device_str}") + torch.xpu.empty_cache() + logger.info(f"[MultiGPU_Device_Utils] Cleared XPU cache on {device_str}") + elif device_str.startswith("npu:"): - torch.npu.empty_cache() - logger.debug("[MultiGPU_Device_Utils] Cleared cache for Ascend NPU") + if hasattr(torch, "npu") and hasattr(torch.npu, "empty_cache"): + logger.info(f"[MultiGPU_Device_Utils] Clearing NPU cache on {device_str}") + torch.npu.empty_cache() + logger.info(f"[MultiGPU_Device_Utils] Cleared NPU cache on {device_str}") + elif device_str.startswith("mlu:"): - torch.mlu.empty_cache() - logger.debug("[MultiGPU_Device_Utils] Cleared cache for Cambricon MLU") + if hasattr(torch, "mlu") and hasattr(torch.mlu, "empty_cache"): + logger.info(f"[MultiGPU_Device_Utils] Clearing MLU cache on {device_str}") + torch.mlu.empty_cache() + logger.info(f"[MultiGPU_Device_Utils] Cleared MLU cache on {device_str}") + elif device_str.startswith("corex:"): - torch.corex.empty_cache() # Hypothetical based on ComfyUI's ixuca support - logger.debug("[MultiGPU_Device_Utils] Cleared cache for CoreX") + if hasattr(torch, "corex") and hasattr(torch.corex, "empty_cache"): + logger.info(f"[MultiGPU_Device_Utils] Clearing CoreX cache on {device_str}") + torch.corex.empty_cache() + logger.info(f"[MultiGPU_Device_Utils] Cleared CoreX cache on {device_str}") # ========================================================================================== diff --git a/distorch.py b/distorch.py index 6588fa3..3ec7d92 100644 --- a/distorch.py +++ b/distorch.py @@ -62,7 +62,8 @@ def register_patched_ggufmodelpatcher(): debug_hash = create_model_hash(self, "patcher") debug_allocations = model_allocation_store.get(debug_hash) if debug_allocations: - soft_empty_cache_multigpu(logger) + logger.info("[MultiGPU DisTorch GGUF] Invoking soft_empty_cache_multigpu before GGUF device assignment") + soft_empty_cache_multigpu() device_assignments = analyze_ggml_loading(self.model, debug_allocations)['device_assignments'] for device, layers in device_assignments.items(): target_device = torch.device(device)