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:
John Pollock
2025-09-08 23:06:21 -05:00
parent 0adf219f60
commit c63b539f1e
5 changed files with 185 additions and 39 deletions
+13 -11
View File
@@ -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),
+4 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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