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.
This commit is contained in:
John Pollock
2025-09-23 04:41:44 -05:00
parent a0fe72e290
commit 3121b2f70c
4 changed files with 171 additions and 136 deletions
+105 -61
View File
@@ -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)
+9 -14
View File
@@ -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:
+21 -24
View File
@@ -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} - <baseline>: " + " | ".join(parts))
# DEBUG absolute
mgpu_logger.memory(f"{identifier}, {comfyui_memory_load(tag)}")
logger.mgpu_mm_log(f"{identifier} {tag} - <baseline>: " + " | ".join(parts))
# Update last snapshot
_MEM_SNAPSHOT_LAST[identifier] = (tag, curr)
+36 -37
View File
@@ -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")