Enhance partially_load with weight unpatching and static assignments

Add logic to detect and unpatch weights for modules with comfy_cast_weights, introduce memory and patch counters, and integrate static device assignments from analyze_safetensor_loading to improve distributed safetensor loading efficiency.
This commit is contained in:
John Pollock
2025-08-23 09:09:31 -05:00
parent 240acae8c5
commit 956bd3bfa0
+28 -1
View File
@@ -53,6 +53,7 @@ def create_safetensor_model_hash(model, caller):
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
# Patch ComfyUI's ModelPatcher
if not hasattr(comfy.model_patcher.ModelPatcher, '_distorch_patched'):
original_partially_load = comfy.model_patcher.ModelPatcher.partially_load
@@ -65,11 +66,37 @@ def register_patched_safetensor_modelpatcher():
debug_hash = create_safetensor_model_hash(self, "partial_load")
allocations = safetensor_allocation_store.get(debug_hash)
mem_counter = 0
patch_counter = 0
loading = self._load_list()
load_completely = []
loading.sort(reverse=True)
for x in loading:
n = x[1]
m = x[2]
params = x[3]
module_mem = x[0]
weight_key = "{}.weight".format(n)
bias_key = "{}.bias".format(n)
cast_weight = self.force_cast_weights
if hasattr(m, "comfy_cast_weights"):
logging.info(f"Unpatching weight {weight_key} for Distorch2")
wipe_lowvram_weight(m)
if allocations:
logger.info(f"[MULTIGPU_DISTORCHV2] Using static allocation for model {debug_hash[:8]}")
# Parse allocation string and apply static assignment
device_assignments = analyze_safetensor_loading(self, allocations)
# Apply our static assignments instead of ComfyUI's dynamic ones
for block_name, target_device in device_assignments['block_assignments'].items():
# Find the module by name