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:
+105
-61
@@ -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
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user