refactor: migrate DisTorch2 allocation tracking to per-model metadata

- Replace global safetensor_allocation_store/safetensor_settings_store and create_safetensor_model_hash
  with a per-model annotation (_distorch_v2_meta) stored directly on the inner model object.
- Update distorch_2 to remove global stores and hash creation; parse and consume allocation strings
  from inner_model._distorch_v2_meta during model registration and loading.
- Update wrappers, checkpoint_multigpu, device_utils, and __init__ to set and read the new metadata
  instead of writing/reading global stores.
- Simplify detection of DisTorch-managed models (check inner_model._distorch_v2_meta) and adjust
  logging to surface inner model ids and allocation info.
- Clean up related imports and dead code paths.

Files changed: distorch_2.py, wrappers.py, checkpoint_multigpu.py, device_utils.py, model_management_mgpu.py, __init__.py
This commit is contained in:
John Pollock
2025-10-13 19:00:59 -05:00
parent bdd612c619
commit 1951513e92
6 changed files with 84 additions and 157 deletions
-2
View File
@@ -269,8 +269,6 @@ from .wrappers import (
override_class_with_distorch_safetensor_v2_clip_no_device,
)
from .distorch_2 import (
safetensor_allocation_store,
create_safetensor_model_hash,
register_patched_safetensor_modelpatcher,
analyze_safetensor_loading,
calculate_safetensor_vvram_allocation,
+7 -11
View File
@@ -9,7 +9,7 @@ import comfy.clip_vision
from comfy.sd import VAE, CLIP
from .device_utils import get_device_list, soft_empty_cache_multigpu
from .model_management_mgpu import multigpu_memory_log
from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher
from .distorch_2 import register_patched_safetensor_modelpatcher
logger = logging.getLogger("MultiGPU")
@@ -108,12 +108,10 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True,
if distorch_config and 'unet_allocation' in distorch_config:
register_patched_safetensor_modelpatcher()
model_hash = create_safetensor_model_hash(model_patcher, "checkpoint_loader_unet")
safetensor_allocation_store[model_hash] = distorch_config['unet_allocation']
safetensor_settings_store[model_hash] = distorch_config.get('unet_settings','')
model.is_distorch = True
inner_model = model_patcher.model
inner_model._distorch_v2_meta = {"full_allocation": distorch_config['unet_allocation']}
logger.info(f"[CHECKPOINT_META] UNET inner_model id=0x{id(inner_model):x}")
model._distorch_high_precision_loras = distorch_config.get('high_precision_loras', True)
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")
@@ -145,12 +143,10 @@ def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True,
if distorch_config and 'clip_allocation' in distorch_config:
if hasattr(clip, 'patcher'):
register_patched_safetensor_modelpatcher()
clip_hash = create_safetensor_model_hash(clip.patcher, "checkpoint_loader_clip")
safetensor_allocation_store[clip_hash] = distorch_config['clip_allocation']
safetensor_settings_store[clip_hash] = distorch_config.get('clip_settings','')
clip.patcher.model.is_distorch = True
inner_clip = clip.patcher.model
inner_clip._distorch_v2_meta = {"full_allocation": distorch_config['clip_allocation']}
logger.info(f"[CHECKPOINT_META] CLIP inner_model id=0x{id(inner_clip):x}")
clip.patcher.model._distorch_high_precision_loras = distorch_config.get('high_precision_loras', True)
logger.info(f"Stored DisTorch2 config for CLIP (hash {clip_hash[:8]}): {distorch_config['clip_allocation']}")
m, u = clip.load_sd(clip_sd, full_model=True) # This respects the patched text_encoder_device
if len(m) > 0: logger.warning(f"CLIP missing keys: {m}")
+3 -16
View File
@@ -235,30 +235,17 @@ original_soft_empty_cache = mm.soft_empty_cache
def soft_empty_cache_distorch2_patched(force=False):
"""Patched mm.soft_empty_cache managing VRAM across all devices, CPU RAM with adaptive thresholding, and DisTorch store pruning."""
from .model_management_mgpu import multigpu_memory_log, check_cpu_memory_threshold, trigger_executor_cache_reset
from .distorch_2 import safetensor_allocation_store, create_safetensor_model_hash
is_distorch_active = False
# Detect DisTorch2-managed models
# logger.mgpu_mm_log(f"[DETECT_DEBUG] Checking DisTorch2 active status - loaded models: {len(mm.current_loaded_models)}, store entries: {len(safetensor_allocation_store)}")
for i, lm in enumerate(mm.current_loaded_models):
mp = lm.model # weakref call to ModelPatcher
mp = lm.model
if mp is not None:
model_hash = create_safetensor_model_hash(mp, "cache_patch_check")
in_store = model_hash in safetensor_allocation_store
alloc_value = safetensor_allocation_store.get(model_hash, "")
model_name = type(getattr(mp, 'model', mp)).__name__
unload_distorch_model = getattr(getattr(mp, 'model', None), '_mgpu_unload_distorch_model', False)
inner_model = mp.model
#logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: {model_name}, hash={model_hash[:8]}, in_store={in_store}, alloc_value='{alloc_value}', unload_distorch_model={unload_distorch_model}")
if in_store and alloc_value:
if hasattr(inner_model, '_distorch_v2_meta'):
is_distorch_active = True
#logger.mgpu_mm_log(f"[DETECT_DEBUG] DisTorch2 ACTIVE detected on model: {model_name}")
break
#logger.mgpu_mm_log(f"[DETECT_DEBUG] Final DisTorch2 active status: {is_distorch_active}")
# Phase 2: adaptive CPU memory management
check_cpu_memory_threshold()
+39 -59
View File
@@ -20,38 +20,6 @@ from .device_utils import get_device_list, soft_empty_cache_multigpu
from .model_management_mgpu import multigpu_memory_log, force_full_system_cleanup
safetensor_allocation_store = {}
safetensor_settings_store = {}
def create_safetensor_model_hash(model, caller):
"""Create a unique hash for a safetensor model to track allocations"""
if hasattr(model, 'model'):
# For ModelPatcher objects
actual_model = model.model
model_type = type(actual_model).__name__
# Use ComfyUI's model_size if available
if hasattr(model, 'model_size'):
model_size = model.model_size()
else:
model_size = sum(p.numel() * p.element_size() for p in actual_model.parameters())
if hasattr(model, 'model_state_dict'):
first_layers = str(list(model.model_state_dict().keys())[:3])
else:
first_layers = str(list(actual_model.state_dict().keys())[:3])
else:
# Direct model
model_type = type(model).__name__
model_size = sum(p.numel() * p.element_size() for p in model.parameters())
first_layers = str(list(model.state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
# DEBUG STATEMENT - ALWAYS LOG THE HASH
logger.debug(f"[MultiGPU DisTorch V2] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
def register_patched_safetensor_modelpatcher():
"""Register and patch the ModelPatcher for distributed safetensor loading"""
from comfy.model_patcher import wipe_lowvram_weight, move_weight_functions
@@ -128,23 +96,35 @@ def register_patched_safetensor_modelpatcher():
device = loaded_model.device
base_memory = loaded_model.model_memory_required(device)
# Check DisTorch flags
is_distorch = hasattr(loaded_model.model.model, '_mgpu_virtual_vram_gb')
has_eject = hasattr(loaded_model.model.model, '_mgpu_eject_models')
if has_eject:
eject_device = device
logger.mgpu_mm_log("DisTorch eject_models=True, is_distorch=True - MAX memory eviction")
if is_distorch:
# is_distorch=True: use compute device allocation size
virtual_vram_gb = loaded_model.model.model._mgpu_virtual_vram_gb
inner_model = loaded_model.model.model
if hasattr(inner_model, '_distorch_v2_meta'):
meta = inner_model._distorch_v2_meta
allocation_str = meta['full_allocation']
# Parse allocation string: "expert#compute_device;virtual_vram_gb;donors"
parts = allocation_str.split('#')
virtual_vram_gb = 0.0
has_eject = False
if len(parts) > 1:
virtual_vram_str = parts[1]
virtual_info = virtual_vram_str.split(';')
if len(virtual_info) > 1:
virtual_vram_gb = float(virtual_info[1])
if len(virtual_info) > 2 and virtual_info[2]:
has_eject = True
if has_eject:
eject_device = device
logger.mgpu_mm_log("DisTorch eject_models detected - MAX memory eviction")
virtual_vram_bytes = virtual_vram_gb * (1024**3)
adjusted_memory = max(0, base_memory - virtual_vram_bytes)
total_memory_required[device] = total_memory_required.get(device, 0) + adjusted_memory
logger.mgpu_mm_log(f"DisTorch is_distorch=True, model adjusted {(base_memory - virtual_vram_bytes)/(1024**3):.2f}GB for device {device}")
logger.mgpu_mm_log(f"DisTorch model adjusted {(base_memory - virtual_vram_bytes)/(1024**3):.2f}GB for device {device}")
else:
# is_distorch=False: use full model size
# Standard model: use full model size
total_memory_required[device] = total_memory_required.get(device, 0) + base_memory
logger.mgpu_mm_log(f"[LOAD_MODELS_GPU] Standard model {(base_memory)/(1024**3):.2f}GB for device {device}")
@@ -209,23 +189,24 @@ def register_patched_safetensor_modelpatcher():
original_partially_load = comfy.model_patcher.ModelPatcher.partially_load
def new_partially_load(self, device_to, extra_memory=0, full_load=False, force_patch_weights=False, **kwargs):
"""Override to use our static device assignments"""
global safetensor_allocation_store
debug_hash = create_safetensor_model_hash(self, "partial_load")
multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "pre-load")
allocations = safetensor_allocation_store.get(debug_hash)
# Set default precision flag before checking
if not hasattr(self.model, '_distorch_high_precision_loras'):
self.model._distorch_high_precision_loras = True
if not allocations:
"""Override to use direct model annotation for allocation"""
mp_id = id(self)
mp_patches_uuid = self.patches_uuid
inner_model = self.model
inner_model_id = id(inner_model)
if not hasattr(inner_model, "_distorch_v2_meta"):
logger.debug(f"[DISTORCH_SKIP] ModelPatcher=0x{mp_id:x} inner_model=0x{inner_model_id:x} type={type(inner_model).__name__} - no metadata, using standard loading")
result = original_partially_load(self, device_to, extra_memory, force_patch_weights)
multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "post-load")
if hasattr(self, '_distorch_block_assignments'):
del self._distorch_block_assignments
return result
allocations = inner_model._distorch_v2_meta['full_allocation']
if not hasattr(self.model, '_distorch_high_precision_loras'):
self.model._distorch_high_precision_loras = True
if not hasattr(self.model, 'current_weight_patches_uuid'):
self.model.current_weight_patches_uuid = None
@@ -308,7 +289,6 @@ def register_patched_safetensor_modelpatcher():
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
-21
View File
@@ -22,30 +22,9 @@ logger = logging.getLogger("MultiGPU")
# Model Analysis and Store Management (DisTorch V1 & V2)
# ==========================================================================================
# DisTorch V2 SafeTensor stores
safetensor_allocation_store = {}
safetensor_settings_store = {}
# DisTorch V1 GGUF stores (backwards compatibility)
model_allocation_store = {}
def create_safetensor_model_hash(model, caller):
"""Create a unique hash for a safetensor model to track allocations"""
if hasattr(model, 'model'):
actual_model = model.model
model_type = type(actual_model).__name__
model_size = model.model_size() if hasattr(model, 'model_size') else sum(p.numel() * p.element_size() for p in actual_model.parameters())
first_layers = str(list(model.model_state_dict().keys() if hasattr(model, 'model_state_dict') else actual_model.state_dict().keys())[:3])
else:
model_type = type(model).__name__
model_size = sum(p.numel() * p.element_size() for p in model.parameters())
first_layers = str(list(model.state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
logger.debug(f"[MultiGPU DisTorch V2] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
def create_model_hash(model, caller):
"""Create a unique hash for a GGUF model to track allocations (DisTorch V1)"""
model_type = type(model.model).__name__
+35 -48
View File
@@ -17,12 +17,7 @@ logger = logging.getLogger("MultiGPU")
def _create_distorch_safetensor_v2_override(cls, device_param_name, device_setter_func, apply_device_kwarg_workaround, eject_models_default=True):
"""Internal factory function creating DisTorch2 override class with parameterized device selection behavior."""
from .distorch_2 import (
register_patched_safetensor_modelpatcher,
safetensor_allocation_store,
safetensor_settings_store,
create_safetensor_model_hash
)
from .distorch_2 import register_patched_safetensor_modelpatcher
from .model_management_mgpu import force_full_system_cleanup
class NodeOverrideDisTorchSafetensorV2(cls):
@@ -118,13 +113,8 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
model_to_check = out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_store")
settings_str = f"{device_value}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
logger.debug(f"[MultiGPU DisTorch V2] Stored allocation for model {model_hash[:8]}: {full_allocation}")
inner_model = model_to_check.model
inner_model._distorch_v2_meta = {"full_allocation": full_allocation}
logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}")
logger.mgpu_mm_log(f"[MODEL_SETUP] Setting DisTorch model properties: virtual_vram_gb={virtual_vram_gb}")
@@ -139,16 +129,6 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] ModelPatcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}")
# SET VIRTUAL VRAM PROPERTY FOR MEMORY CALCULATION
if inner_model:
inner_model._mgpu_virtual_vram_gb = virtual_vram_gb
logger.mgpu_mm_log(f"[VIRTUAL_VRAM_SET] Set _mgpu_virtual_vram_gb={virtual_vram_gb}GB on inner model (id=0x{inner_model_id:x}) for memory assessment")
# SET EJECT MODELS PROPERTY IF ENABLED
if eject_models and inner_model:
inner_model._mgpu_eject_models = True
logger.mgpu_mm_log(f"[EJECT_FLAG_SET] Set _mgpu_eject_models=True on inner model (id=0x{inner_model_id:x}) - will trigger ejection during load_models_gpu")
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
mp = out[0].patcher
mp_id = id(mp)
@@ -159,11 +139,6 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] ModelPatcher via patcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}")
# SET VIRTUAL VRAM PROPERTY FOR MEMORY CALCULATION
if inner_model:
inner_model._mgpu_virtual_vram_gb = virtual_vram_gb
logger.mgpu_mm_log(f"[VIRTUAL_VRAM_SET] Set _mgpu_virtual_vram_gb={virtual_vram_gb}GB on inner model (id=0x{inner_model_id:x}) for memory assessment")
return out
return NodeOverrideDisTorchSafetensorV2
@@ -211,7 +186,7 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
def override_class_with_distorch_gguf(cls):
"""DisTorch V1 Legacy wrapper - maintains V1 UI but calls V2 backend"""
from . import set_current_device
from .distorch_2 import register_patched_safetensor_modelpatcher, safetensor_allocation_store, create_safetensor_model_hash
from .distorch_2 import register_patched_safetensor_modelpatcher
class NodeOverrideDisTorchGGUFLegacy(cls):
@classmethod
@@ -257,12 +232,15 @@ def override_class_with_distorch_gguf(cls):
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **clean_kwargs)
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "v1_compat")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "v1_compat")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0].patcher
if model_to_check:
inner_model = model_to_check.model
inner_model._distorch_v2_meta = {"full_allocation": full_allocation}
return out
@@ -272,7 +250,7 @@ def override_class_with_distorch_gguf(cls):
def override_class_with_distorch_gguf_v2(cls):
"""DisTorch V2 wrapper for GGUF models"""
from . import set_current_device
from .distorch_2 import register_patched_safetensor_modelpatcher, safetensor_allocation_store, create_safetensor_model_hash
from .distorch_2 import register_patched_safetensor_modelpatcher
class NodeOverrideDisTorchGGUFv2(cls):
@classmethod
@@ -316,12 +294,15 @@ def override_class_with_distorch_gguf_v2(cls):
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **clean_kwargs)
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "v2_gguf")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "v2_gguf")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0].patcher
if model_to_check:
inner_model = model_to_check.model
inner_model._distorch_v2_meta = {"full_allocation": full_allocation}
return out
@@ -331,7 +312,7 @@ def override_class_with_distorch_gguf_v2(cls):
def override_class_with_distorch_clip(cls):
"""DisTorch V1 wrapper for CLIP models - calls V2 backend"""
from . import set_current_text_encoder_device
from .distorch_2 import register_patched_safetensor_modelpatcher, safetensor_allocation_store, create_safetensor_model_hash
from .distorch_2 import register_patched_safetensor_modelpatcher
class NodeOverrideDisTorchClip(cls):
@classmethod
@@ -377,12 +358,15 @@ def override_class_with_distorch_clip(cls):
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **clean_kwargs)
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "v1_clip")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "v1_clip")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0].patcher
if model_to_check:
inner_model = model_to_check.model
inner_model._distorch_v2_meta = {"full_allocation": full_allocation}
return out
@@ -392,7 +376,7 @@ def override_class_with_distorch_clip(cls):
def override_class_with_distorch_clip_no_device(cls):
"""DisTorch V1 wrapper for Triple/Quad CLIP models - calls V2 backend"""
from . import set_current_text_encoder_device
from .distorch_2 import register_patched_safetensor_modelpatcher, safetensor_allocation_store, create_safetensor_model_hash
from .distorch_2 import register_patched_safetensor_modelpatcher
class NodeOverrideDisTorchClipNoDevice(cls):
@classmethod
@@ -438,12 +422,15 @@ def override_class_with_distorch_clip_no_device(cls):
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **clean_kwargs)
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "v1_clip_nodev")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "v1_clip_nodev")
safetensor_allocation_store[model_hash] = full_allocation
model_to_check = out[0].patcher
if model_to_check:
inner_model = model_to_check.model
inner_model._distorch_v2_meta = {"full_allocation": full_allocation}
return out