Additonal refinements to DisTorch2 cache/unload to avoid OOM. Needs at least one more clean-up pass.
- Introduce MEMORY_LOG flag and logger.memory method to gate high-volume memory logs - Demote device setter logs from info to debug to reduce noise - Clarify patch announcement (remove text_encoder_initial_device mention) - Update soft_empty_cache patch log to emphasize multi-device allocation/clearing; delegate to original when DisTorch2 is inactive - Rework load_models_gpu preflight for large DisTorch2 models: - more robust ModelPatcher detection (direct or via .patcher) - track allowed devices and incoming model names - improved large-model detection and proactive unload/clearing on donor/offload devices - mitigates OOM during large model (e.g., UNet) swaps - Minor cleanup of verbose comments and wording in logs
This commit is contained in:
+118
-51
@@ -23,7 +23,13 @@ if not logger.handlers:
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(log_level)
|
||||
logger.info(f"[MultiGPU Initialization] Logger initialized with level: {logging.getLevelName(log_level)}")
|
||||
|
||||
MEMORY_LOG = True
|
||||
|
||||
def memory_method(self, msg):
|
||||
if MEMORY_LOG:
|
||||
self.info(msg)
|
||||
logger.memory = memory_method.__get__(logger, type(logger))
|
||||
|
||||
|
||||
# Global device state management
|
||||
@@ -33,12 +39,12 @@ current_text_encoder_device = mm.text_encoder_device()
|
||||
def set_current_device(device):
|
||||
global current_device
|
||||
current_device = device
|
||||
logger.info(f"[MultiGPU Initialization] current_device set to: {device}")
|
||||
logger.debug(f"[MultiGPU Initialization] current_device set to: {device}")
|
||||
|
||||
def set_current_text_encoder_device(device):
|
||||
global current_text_encoder_device
|
||||
current_text_encoder_device = device
|
||||
logger.info(f"[MultiGPU Initialization] current_text_encoder_device set to: {device}")
|
||||
logger.debug(f"[MultiGPU Initialization] current_text_encoder_device set to: {device}")
|
||||
|
||||
def override_class(cls):
|
||||
class NodeOverride(cls):
|
||||
@@ -136,7 +142,7 @@ def text_encoder_device_patched():
|
||||
return device
|
||||
|
||||
|
||||
logger.info(f"[MultiGPU Core Patching] Patching mm.get_torch_device, mm.text_encoder_device, and mm.text_encoder_initial_device")
|
||||
logger.info(f"[MultiGPU Core Patching] Patching mm.get_torch_device and mm.text_encoder_device")
|
||||
logger.debug(f"[MultiGPU DEBUG] Initial current_device: {current_device}")
|
||||
logger.debug(f"[MultiGPU DEBUG] Initial current_text_encoder_device: {current_text_encoder_device}")
|
||||
mm.get_torch_device = get_torch_device_patched
|
||||
@@ -215,12 +221,8 @@ 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")
|
||||
logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for DisTorch2 Multi-Device Allocation/Clearing")
|
||||
|
||||
# Store the original function for fallback behavior
|
||||
original_soft_empty_cache = mm.soft_empty_cache
|
||||
|
||||
def soft_empty_cache_distorch2_patched(force=False):
|
||||
@@ -246,14 +248,8 @@ def soft_empty_cache_distorch2_patched(force=False):
|
||||
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)
|
||||
@@ -274,42 +270,89 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
|
||||
|
||||
# Detect incoming large DisTorch2 request
|
||||
incoming_is_distorch = False
|
||||
incoming_distorch_nonzero = False
|
||||
incoming_is_large = False
|
||||
incoming_patchers = set()
|
||||
incoming_loaded_names = []
|
||||
incoming_allowed_devices = None
|
||||
|
||||
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()
|
||||
# Identify ModelPatcher (prefer direct; fall back to .patcher)
|
||||
if hasattr(lm, "load_device"):
|
||||
patcher = lm
|
||||
elif hasattr(lm, "patcher"):
|
||||
patcher = lm.patcher
|
||||
else:
|
||||
patcher = None
|
||||
|
||||
if size_bytes > LARGE_MODEL_THRESHOLD:
|
||||
model_for_hash = patcher if patcher is not None else getattr(lm, "model", lm)
|
||||
|
||||
if patcher is not None:
|
||||
incoming_patchers.add(patcher)
|
||||
|
||||
# Determine required memory directly from ModelPatcher (no wrapper; no side effects)
|
||||
device_str = str(patcher.load_device)
|
||||
if patcher.current_loaded_device() == patcher.load_device:
|
||||
required_bytes = patcher.model_size() - patcher.loaded_size()
|
||||
else:
|
||||
required_bytes = patcher.model_size()
|
||||
|
||||
if required_bytes > LARGE_MODEL_THRESHOLD:
|
||||
incoming_is_large = True
|
||||
else:
|
||||
device_str = "n/a"
|
||||
required_bytes = 0
|
||||
|
||||
# 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
|
||||
# Check DisTorch2 management via allocation store (unchanged trigger)
|
||||
model_hash = create_safetensor_model_hash(model_for_hash, "load_patch_check")
|
||||
if model_hash in safetensor_allocation_store and safetensor_allocation_store.get(model_hash):
|
||||
incoming_is_distorch = True
|
||||
if required_bytes > 0:
|
||||
incoming_distorch_nonzero = True
|
||||
if incoming_allowed_devices is None:
|
||||
# Derive compute/donor devices from allocation string
|
||||
alloc_str = safetensor_allocation_store.get(model_hash, "")
|
||||
allowed = set()
|
||||
if alloc_str:
|
||||
parts = alloc_str.split("#", 1)
|
||||
if len(parts) == 2 and parts[1]:
|
||||
vram = parts[1]
|
||||
segs = vram.split(";")
|
||||
# compute device
|
||||
if len(segs) >= 1 and segs[0]:
|
||||
allowed.add(segs[0].strip())
|
||||
# donors list (comma-separated)
|
||||
if len(segs) >= 3 and segs[2]:
|
||||
for d in segs[2].split(","):
|
||||
d = d.strip()
|
||||
if d:
|
||||
allowed.add(d)
|
||||
else:
|
||||
# Expert fraction string: "dev,fraction;dev2,fraction2;..."
|
||||
for token in alloc_str.split(";"):
|
||||
if "," in token:
|
||||
dev, frac = token.split(",", 1)
|
||||
fs = frac.strip()
|
||||
numlike = fs.replace(".", "", 1).isdigit()
|
||||
if numlike and float(fs) > 0.0:
|
||||
allowed.add(dev.strip())
|
||||
if not allowed:
|
||||
allowed = {str(patcher.load_device), "cpu"}
|
||||
incoming_allowed_devices = allowed
|
||||
|
||||
# Log informational context
|
||||
incoming_loaded_names.append(f"{type(getattr(mp, 'model', mp)).__name__}:{size_bytes/(1024**3):.2f}GB")
|
||||
# Log informational context with required bytes and device
|
||||
try:
|
||||
model_name = type(getattr(model_for_hash, "model", model_for_hash)).__name__
|
||||
except Exception:
|
||||
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)}")
|
||||
|
||||
# Proactive unload if both conditions are met
|
||||
if incoming_is_distorch and incoming_is_large:
|
||||
if incoming_distorch_nonzero:
|
||||
logger.info("[MultiGPU Core Patching] 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.")
|
||||
|
||||
@@ -329,6 +372,12 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
|
||||
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'):
|
||||
@@ -339,19 +388,17 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
|
||||
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)")
|
||||
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
|
||||
@@ -365,7 +412,27 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
|
||||
# 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)")
|
||||
# Lineage-aligned cache clear when no unloads happened: apply core 25% rule, per DisTorch devices
|
||||
if incoming_allowed_devices is not None and mm.vram_state != mm.VRAMState.HIGH_VRAM:
|
||||
triggered = []
|
||||
for dev_str in incoming_allowed_devices:
|
||||
try:
|
||||
dev_obj = torch.device(dev_str)
|
||||
except Exception:
|
||||
continue
|
||||
free_total, free_torch = mm.get_free_memory(dev_obj, torch_free_too=True)
|
||||
# free_total: system free; free_torch: torch reserved-but-free
|
||||
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()")
|
||||
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")
|
||||
else:
|
||||
logger.info("[MultiGPU Core Patching] 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")
|
||||
|
||||
# Continue with original behavior
|
||||
return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load)
|
||||
|
||||
Reference in New Issue
Block a user