Additional garbage/cache collection (#101) addressed DisTorch2 Device issue for CLIP hopefully closing (#99,#104)
Add comprehensive memory cache clearing aligned with ComfyUI patterns to improve stability and reduce OOM incidents in multi-device scenarios. **Addresses Memory/Garbage Collection Issues:** - Created `soft_empty_cache_multigpu()` function in device_utils.py - Replicates ComfyUI's cache clearing for all devices (CUDA, MPS, XPU, NPU, MLU) - Includes CUDA IPC collect optimization like ComfyUI - Strategically placed calls before major memory allocations **Addresses CLIP loading issues:** - Fixed DisTorch2 device device varibale management before text encoder operations **`soft_empty_cache_multigpu()` implementation Aligned with ComfyUI's Patterns:** - Called after GC operations - Placed before major memory allocations - Matches ComfyUI's proven memory management strategy - Same device clearing logic for multi-device scenarios
This commit is contained in:
+13
-11
@@ -29,6 +29,7 @@ if not logger.handlers:
|
||||
# Global device state management
|
||||
current_device = mm.get_torch_device()
|
||||
current_text_encoder_device = mm.text_encoder_device()
|
||||
current_text_encoder_initial_device = mm.text_encoder_device()
|
||||
|
||||
def set_current_device(device):
|
||||
global current_device
|
||||
@@ -36,7 +37,7 @@ def set_current_device(device):
|
||||
logger.info(f"[MultiGPU Initialization] current_device set to: {device}")
|
||||
|
||||
def set_current_text_encoder_device(device):
|
||||
global current_text_encoder_device
|
||||
global current_text_encoder_device, current_text_encoder_initial_device
|
||||
current_text_encoder_device = device
|
||||
current_text_encoder_initial_device = device
|
||||
logger.info(f"[MultiGPU Initialization] current_text_encoder_device and current_text_encoder_initial_device set to: {device}")
|
||||
@@ -192,7 +193,8 @@ from .distorch_2 import (
|
||||
register_patched_safetensor_modelpatcher,
|
||||
analyze_safetensor_loading,
|
||||
calculate_safetensor_vvram_allocation,
|
||||
override_class_with_distorch_safetensor_v2
|
||||
override_class_with_distorch_safetensor_v2,
|
||||
override_class_with_distorch_safetensor_v2_clip
|
||||
)
|
||||
|
||||
# Import advanced checkpoint loaders
|
||||
@@ -229,13 +231,13 @@ if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
# DisTorch 2 SafeTensor nodes for FLUX and other safetensor models
|
||||
NODE_CLASS_MAPPINGS["UNETLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
|
||||
NODE_CLASS_MAPPINGS["VAELoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
|
||||
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["DualCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2_clip(GLOBAL_NODE_CLASS_MAPPINGS["DualCLIPLoader"])
|
||||
if "TripleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["TripleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2_clip(GLOBAL_NODE_CLASS_MAPPINGS["TripleCLIPLoader"])
|
||||
if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
|
||||
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2_clip(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
|
||||
NODE_CLASS_MAPPINGS["CLIPVisionLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
|
||||
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
|
||||
NODE_CLASS_MAPPINGS["ControlNetLoaderDisTorch2MultiGPU"] = override_class_with_distorch_safetensor_v2(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
|
||||
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
|
||||
@@ -307,10 +309,10 @@ gguf_nodes = {
|
||||
"QuadrupleCLIPLoaderGGUFDisTorchMultiGPU": override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF),
|
||||
"UnetLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(UnetLoaderGGUF),
|
||||
"UnetLoaderGGUFAdvancedDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(UnetLoaderGGUFAdvanced),
|
||||
"CLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(CLIPLoaderGGUF),
|
||||
"DualCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(DualCLIPLoaderGGUF),
|
||||
"TripleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(TripleCLIPLoaderGGUF),
|
||||
"QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2(QuadrupleCLIPLoaderGGUF),
|
||||
"CLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2_clip(CLIPLoaderGGUF),
|
||||
"DualCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2_clip(DualCLIPLoaderGGUF),
|
||||
"TripleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2_clip(TripleCLIPLoaderGGUF),
|
||||
"QuadrupleCLIPLoaderGGUFDisTorch2MultiGPU": override_class_with_distorch_safetensor_v2_clip(QuadrupleCLIPLoaderGGUF),
|
||||
"UnetLoaderGGUFMultiGPU": override_class(UnetLoaderGGUF),
|
||||
"UnetLoaderGGUFAdvancedMultiGPU": override_class(UnetLoaderGGUFAdvanced),
|
||||
"CLIPLoaderGGUFMultiGPU": override_class_clip(CLIPLoaderGGUF),
|
||||
|
||||
@@ -12,7 +12,7 @@ import comfy.model_management as mm
|
||||
import comfy.model_detection
|
||||
import comfy.clip_vision
|
||||
from comfy.sd import VAE, CLIP
|
||||
from .device_utils import get_device_list
|
||||
from .device_utils import get_device_list, soft_empty_cache_multigpu
|
||||
from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher
|
||||
|
||||
logger = logging.getLogger("MultiGPU")
|
||||
@@ -107,8 +107,9 @@ 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)
|
||||
|
||||
soft_empty_cache_multigpu(logger)
|
||||
model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=unet_compute_device, offload_device=mm.unet_offload_device())
|
||||
|
||||
|
||||
if distorch_config and 'unet_allocation' in distorch_config:
|
||||
register_patched_safetensor_modelpatcher()
|
||||
model_hash = create_safetensor_model_hash(model_patcher, "checkpoint_loader_unet")
|
||||
@@ -136,6 +137,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:
|
||||
soft_empty_cache_multigpu(logger)
|
||||
clip_params = comfy.utils.calculate_parameters(clip_sd)
|
||||
clip = CLIP(clip_target, embedding_directory=embedding_directory, tokenizer_data=clip_sd, parameters=clip_params, model_options=te_model_options)
|
||||
|
||||
|
||||
+58
-18
@@ -45,9 +45,9 @@ def get_device_list():
|
||||
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available():
|
||||
device_count = torch.cuda.device_count()
|
||||
devs += [f"cuda:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU] Found {device_count} CUDA device(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} CUDA device(s)")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] CUDA detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] CUDA detection failed: {e}")
|
||||
|
||||
# XPU devices (Intel GPUs)
|
||||
try:
|
||||
@@ -59,9 +59,9 @@ def get_device_list():
|
||||
if hasattr(torch, "xpu") and hasattr(torch.xpu, "is_available") and torch.xpu.is_available():
|
||||
device_count = torch.xpu.device_count()
|
||||
devs += [f"xpu:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU] Found {device_count} XPU device(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} XPU device(s)")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] XPU detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] XPU detection failed: {e}")
|
||||
|
||||
# NPU devices (Ascend NPUs from Huawei)
|
||||
try:
|
||||
@@ -69,9 +69,9 @@ def get_device_list():
|
||||
if hasattr(torch, "npu") and hasattr(torch.npu, "is_available") and torch.npu.is_available():
|
||||
device_count = torch.npu.device_count()
|
||||
devs += [f"npu:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU] Found {device_count} NPU device(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} NPU device(s)")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] NPU detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] NPU detection failed: {e}")
|
||||
|
||||
# MLU devices (Cambricon MLUs)
|
||||
try:
|
||||
@@ -79,17 +79,17 @@ def get_device_list():
|
||||
if hasattr(torch, "mlu") and hasattr(torch.mlu, "is_available") and torch.mlu.is_available():
|
||||
device_count = torch.mlu.device_count()
|
||||
devs += [f"mlu:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU] Found {device_count} MLU device(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} MLU device(s)")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] MLU detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] MLU detection failed: {e}")
|
||||
|
||||
# MPS device (Apple Metal - single device only)
|
||||
try:
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
devs.append("mps")
|
||||
logger.debug("[MultiGPU] Found MPS device")
|
||||
logger.debug("[MultiGPU_Device_Utils] Found MPS device")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] MPS detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] MPS detection failed: {e}")
|
||||
|
||||
# DirectML devices (Windows DirectML for AMD/Intel/NVIDIA)
|
||||
try:
|
||||
@@ -97,9 +97,9 @@ def get_device_list():
|
||||
adapter_count = torch_directml.device_count()
|
||||
if adapter_count > 0:
|
||||
devs += [f"directml:{i}" for i in range(adapter_count)]
|
||||
logger.debug(f"[MultiGPU] Found {adapter_count} DirectML adapter(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {adapter_count} DirectML adapter(s)")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] DirectML detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] DirectML detection failed: {e}")
|
||||
|
||||
# IXUCA/CoreX devices (special accelerator)
|
||||
try:
|
||||
@@ -108,18 +108,18 @@ def get_device_list():
|
||||
if hasattr(torch.corex, "device_count"):
|
||||
device_count = torch.corex.device_count()
|
||||
devs += [f"corex:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU] Found {device_count} CoreX device(s)")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} CoreX device(s)")
|
||||
else:
|
||||
devs.append("corex:0")
|
||||
logger.debug("[MultiGPU] Found CoreX device")
|
||||
logger.debug("[MultiGPU_Device_Utils] Found CoreX device")
|
||||
except Exception as e:
|
||||
logger.debug(f"[MultiGPU] CoreX detection failed: {e}")
|
||||
logger.debug(f"[MultiGPU_Device_Utils] CoreX detection failed: {e}")
|
||||
|
||||
# Cache the result for future calls
|
||||
_DEVICE_LIST_CACHE = devs
|
||||
|
||||
# Log only once when initially populated
|
||||
logger.info(f"[MultiGPU] Device list initialized: {devs}")
|
||||
logger.info(f"[MultiGPU_Device_Utils] Device list initialized: {devs}")
|
||||
|
||||
return devs
|
||||
|
||||
@@ -218,10 +218,10 @@ def get_device_type(device_string):
|
||||
def parse_device_string(device_string):
|
||||
"""
|
||||
Parse a device string into type and index.
|
||||
|
||||
|
||||
Args:
|
||||
device_string: Device identifier like "cuda:0", "cpu", "xpu:1", etc.
|
||||
|
||||
|
||||
Returns:
|
||||
Tuple of (device_type, device_index) where index is None for non-indexed devices
|
||||
"""
|
||||
@@ -229,3 +229,43 @@ def parse_device_string(device_string):
|
||||
parts = device_string.split(":")
|
||||
return parts[0], int(parts[1])
|
||||
return device_string, None
|
||||
|
||||
|
||||
def soft_empty_cache_multigpu(logger):
|
||||
"""
|
||||
Replicate ComfyUI's cache clearing but for ALL devices in MultiGPU.
|
||||
MultiGPU adaptation of ComfyUI's soft_empty_cache() functionality.
|
||||
"""
|
||||
import gc
|
||||
|
||||
logger.info("[MultiGPU_Device_Utils] Preparing devices for optimized safetensor loading")
|
||||
|
||||
# Python GC (same as all implementations)
|
||||
gc.collect()
|
||||
logger.debug("[MultiGPU_Device_Utils] Performed garbage collection before safetensor loading")
|
||||
|
||||
# Clear cache for ALL devices (not just ComfyUI's single device)
|
||||
all_devices = get_device_list()
|
||||
|
||||
for device_str in all_devices:
|
||||
if device_str.startswith("cuda:"):
|
||||
device_idx = int(device_str.split(":")[1])
|
||||
torch.cuda.set_device(device_idx)
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect() # ComfyUI's CUDA optimization
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Cleared cache + IPC for {device_str}")
|
||||
elif device_str == "mps":
|
||||
torch.mps.empty_cache()
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for MPS")
|
||||
elif device_str.startswith("xpu:"):
|
||||
torch.xpu.empty_cache()
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for Intel XPU")
|
||||
elif device_str.startswith("npu:"):
|
||||
torch.npu.empty_cache()
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for Ascend NPU")
|
||||
elif device_str.startswith("mlu:"):
|
||||
torch.mlu.empty_cache()
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for Cambricon MLU")
|
||||
elif device_str.startswith("corex:"):
|
||||
torch.corex.empty_cache() # Hypothetical based on ComfyUI's ixuca support
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for CoreX")
|
||||
|
||||
+2
-1
@@ -12,7 +12,7 @@ logger = logging.getLogger("MultiGPU")
|
||||
import copy
|
||||
from collections import defaultdict
|
||||
import comfy.model_management as mm
|
||||
from .device_utils import get_device_list
|
||||
from .device_utils import get_device_list, soft_empty_cache_multigpu
|
||||
|
||||
# Global store for model allocations
|
||||
model_allocation_store = {}
|
||||
@@ -62,6 +62,7 @@ def register_patched_ggufmodelpatcher():
|
||||
debug_hash = create_model_hash(self, "patcher")
|
||||
debug_allocations = model_allocation_store.get(debug_hash)
|
||||
if debug_allocations:
|
||||
soft_empty_cache_multigpu(logger)
|
||||
device_assignments = analyze_ggml_loading(self.model, debug_allocations)['device_assignments']
|
||||
for device, layers in device_assignments.items():
|
||||
target_device = torch.device(device)
|
||||
|
||||
+108
-7
@@ -8,6 +8,7 @@ import torch
|
||||
import logging
|
||||
import hashlib
|
||||
import re
|
||||
import gc
|
||||
|
||||
logger = logging.getLogger("MultiGPU")
|
||||
import copy
|
||||
@@ -16,7 +17,7 @@ from collections import defaultdict
|
||||
import comfy.model_management as mm
|
||||
import comfy.model_patcher
|
||||
from . import current_device
|
||||
from .device_utils import get_device_list
|
||||
from .device_utils import get_device_list, soft_empty_cache_multigpu
|
||||
|
||||
safetensor_allocation_store = {}
|
||||
safetensor_settings_store = {}
|
||||
@@ -66,10 +67,11 @@ def register_patched_safetensor_modelpatcher():
|
||||
allocations = safetensor_allocation_store.get(debug_hash)
|
||||
|
||||
if not hasattr(self.model, '_distorch_high_precision_loras') or not allocations:
|
||||
result = original_partially_load(self, device_to, extra_memory, force_patch_weights)
|
||||
result = original_partially_load(self, device_to, extra_memory, force_patch_weights)
|
||||
if hasattr(self, '_distorch_block_assignments'):
|
||||
del self._distorch_block_assignments
|
||||
return result
|
||||
soft_empty_cache_multigpu(logger)
|
||||
|
||||
mem_counter = 0
|
||||
|
||||
@@ -550,14 +552,14 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
|
||||
def override_class_with_distorch_safetensor_v2(cls):
|
||||
"""DisTorch 2.0 wrapper for safetensor models"""
|
||||
from . import current_device
|
||||
|
||||
|
||||
class NodeOverrideDisTorchSafetensorV2(cls):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
||||
devices = get_device_list()
|
||||
compute_device = devices[1] if len(devices) > 1 else devices[0]
|
||||
|
||||
|
||||
inputs["optional"] = inputs.get("optional", {})
|
||||
inputs["optional"]["compute_device"] = (devices, {"default": compute_device})
|
||||
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1})
|
||||
@@ -571,7 +573,7 @@ def override_class_with_distorch_safetensor_v2(cls):
|
||||
TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch2)"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, *args, compute_device=None, virtual_vram_gb=4.0,
|
||||
def IS_CHANGED(s, *args, compute_device=None, virtual_vram_gb=4.0,
|
||||
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
|
||||
# Create a hash of our specific settings
|
||||
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{high_precision_loras}"
|
||||
@@ -627,7 +629,7 @@ def override_class_with_distorch_safetensor_v2(cls):
|
||||
vram_string = compute_device
|
||||
|
||||
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}")
|
||||
|
||||
if hasattr(out[0], 'model'):
|
||||
@@ -635,10 +637,109 @@ def override_class_with_distorch_safetensor_v2(cls):
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash
|
||||
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
|
||||
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
|
||||
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash
|
||||
|
||||
return out
|
||||
|
||||
return NodeOverrideDisTorchSafetensorV2
|
||||
|
||||
|
||||
def override_class_with_distorch_safetensor_v2_clip(cls):
|
||||
"""DisTorch 2.0 wrapper for safetensor CLIP models"""
|
||||
from . import current_device
|
||||
|
||||
class NodeOverrideDisTorchSafetensorV2Clip(cls):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = copy.deepcopy(cls.INPUT_TYPES())
|
||||
devices = get_device_list()
|
||||
default_device = devices[1] if len(devices) > 1 else devices[0]
|
||||
|
||||
inputs["optional"] = inputs.get("optional", {})
|
||||
inputs["optional"]["device"] = (devices, {"default": default_device}) # Changed from compute_device
|
||||
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1})
|
||||
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
|
||||
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
|
||||
inputs["optional"]["high_precision_loras"] = ("BOOLEAN", {"default": True})
|
||||
return inputs
|
||||
|
||||
CATEGORY = "multigpu/distorch_2"
|
||||
FUNCTION = "override"
|
||||
TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch2)"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
|
||||
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
|
||||
# Create a hash of our specific settings
|
||||
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{high_precision_loras}" # Changed from compute_device
|
||||
return hashlib.sha256(settings_str.encode()).hexdigest()
|
||||
|
||||
def override(self, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
|
||||
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
|
||||
|
||||
from . import set_current_text_encoder_device # Use text encoder device setter
|
||||
if device is not None:
|
||||
set_current_text_encoder_device(device)
|
||||
|
||||
kwargs['device'] = 'default' # Hardcode device setting like in standard clip wrapper
|
||||
|
||||
# Register our patched ModelPatcher
|
||||
register_patched_safetensor_modelpatcher()
|
||||
|
||||
# Call original function
|
||||
fn = getattr(super(), cls.FUNCTION)
|
||||
|
||||
# --- Check if we need to unload the model due to settings change ---
|
||||
# This logic is a bit redundant with IS_CHANGED, but provides clear logging
|
||||
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}" # Changed from compute_device
|
||||
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
|
||||
|
||||
# Temporarily load to get hash without applying our patch
|
||||
temp_out = fn(*args, **kwargs)
|
||||
model_to_check = None
|
||||
if hasattr(temp_out[0], 'model'):
|
||||
model_to_check = temp_out[0]
|
||||
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
|
||||
model_to_check = temp_out[0].patcher
|
||||
|
||||
if model_to_check:
|
||||
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
|
||||
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.")
|
||||
else:
|
||||
logger.info(f"[MultiGPU_DisTorch2] Settings unchanged for model {model_hash[:8]}. Using cached model.")
|
||||
|
||||
out = fn(*args, **kwargs)
|
||||
|
||||
# Store high_precision_loras in the model for later retrieval
|
||||
if hasattr(out[0], 'model'):
|
||||
out[0].model._distorch_high_precision_loras = high_precision_loras
|
||||
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
|
||||
out[0].patcher.model._distorch_high_precision_loras = high_precision_loras
|
||||
|
||||
vram_string = ""
|
||||
if virtual_vram_gb > 0:
|
||||
vram_string = f"{device};{virtual_vram_gb};{donor_device}" # Changed from compute_device
|
||||
elif expert_mode_allocations: # Only include device if there's an expert string
|
||||
vram_string = device # Changed from compute_device
|
||||
|
||||
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}")
|
||||
|
||||
if hasattr(out[0], 'model'):
|
||||
model_hash = create_safetensor_model_hash(out[0], "override")
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash
|
||||
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
|
||||
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
|
||||
safetensor_allocation_store[model_hash] = full_allocation
|
||||
safetensor_settings_store[model_hash] = settings_hash
|
||||
|
||||
return out
|
||||
|
||||
return NodeOverrideDisTorchSafetensorV2Clip
|
||||
|
||||
Reference in New Issue
Block a user