refactor(distorch): Implement IS_CHANGED for robust model reloading

This commit refactors the model loading logic to properly integrate with ComfyUI's caching system.

- Implemented the `IS_CHANGED` class method, which creates a hash of the DisTorch-specific settings (e.g., `compute_device`, `virtual_vram_gb`).
- This allows ComfyUI to automatically detect when settings have changed and trigger a model reload, invalidating the cache correctly.
- Removed the previous manual and less reliable logic for unloading and reloading the model from within the `override` function.
- Set the default log level to "Engineering" to provide more detailed output during development.
This commit is contained in:
John Pollock
2025-08-13 16:27:02 -05:00
parent e288152dae
commit fb6e2e6ffa
2 changed files with 39 additions and 35 deletions
+1 -1
View File
@@ -9,7 +9,7 @@ from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
# --- DisTorch V2 Logging Configuration ---
# Set to "E" for Engineering (DEBUG) or "P" for Production (INFO)
LOG_LEVEL = "P"
LOG_LEVEL = "E"
# Configure logger
log_level = logging.DEBUG if LOG_LEVEL == "E" else logging.INFO
+38 -34
View File
@@ -9,6 +9,7 @@ import torch
import logging
import hashlib
import copy
import inspect
from collections import defaultdict
import comfy.model_management as mm
import comfy.model_patcher
@@ -441,6 +442,13 @@ def override_class_with_distorch_safetensor_v2(cls):
FUNCTION = "override"
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,
donor_device="cpu", expert_mode_allocations="", **kwargs):
# Create a hash of our specific settings
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
return hashlib.sha256(settings_str.encode()).hexdigest()
def override(self, *args, compute_device=None, virtual_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
from . import set_current_device
@@ -449,6 +457,34 @@ def override_class_with_distorch_safetensor_v2(cls):
# 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"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
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:
logging.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.")
# The IS_CHANGED mechanism should handle the reload, this is for logging.
else:
logging.info(f"[MultiGPU_DisTorch2] Settings unchanged for model {model_hash[:8]}. Using cached model.")
out = fn(*args, **kwargs)
# Build allocation string - EXACTLY like GGUF
vram_string = ""
@@ -457,49 +493,17 @@ def override_class_with_distorch_safetensor_v2(cls):
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
# --- Force Model Reload on Setting Change ---
# Create a hash of the DisTorch settings
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()[:8]
# Temporarily load the model to get its hash, without applying our patch yet
fn = getattr(super(), cls.FUNCTION)
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:
logging.info(f"[MultiGPU_DisTorch2] Settings changed for model {model_hash[:8]}. Forcing reload.")
mm.unload_model(model_to_check)
# Update the settings store *before* reloading
safetensor_settings_store[model_hash] = settings_hash
# Call the loader again now that the model is unloaded
out = fn(*args, **kwargs)
else:
out = temp_out # Use the already loaded model
else:
out = temp_out # Should not happen, but as a fallback
logging.info(f"[MULTIGPU_DISTORCHV2] Full allocation string: {full_allocation}")
# Store allocation for the model - EXACTLY like GGUF
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 # Ensure it's set
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 # Ensure it's set
safetensor_settings_store[model_hash] = settings_hash
return out