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:
John Pollock
2025-09-21 09:30:03 -05:00
parent 55a0d22b01
commit a0fe72e290
+118 -51
View File
@@ -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)