Refactor GGUF model patcher to enhance device handling and introduce model hashing for allocation tracking

This commit is contained in:
John Pollock
2025-01-24 21:06:01 -06:00
parent 624c893942
commit b99e152686
+46 -31
View File
@@ -4,55 +4,62 @@ import torch
from collections import defaultdict
import comfy.model_management
import copy
import hashlib
def register_patched_ggufmodelpatcher(node_instance):
model_allocation_store = {}
def create_model_hash(model, caller):
model_type = type(model.model).__name__
model_size = model.model_size()
first_layers = str(list(model.model_state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
return final_hash
def register_patched_ggufmodelpatcher():
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
module = sys.modules[original_loader.__module__]
if not hasattr(module.GGUFModelPatcher, '_patched'):
original_load = module.GGUFModelPatcher.load
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher not yet patched, applying patch")
def new_load(self, *args, force_patch_weights=False, **kwargs):
global model_allocation_store
super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs)
debug_hash = create_model_hash(self, "patcher")
linked = []
module_count = 0
for n, m in self.model.named_modules():
module_count += 1
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None)
logging.info(f"MultiGPU: GGUFDisTorch - Weight Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None)
#logging.info(f"MultiGPU: GGUFDisTorch - Bias Module {n} on device {device}, offload_device is {self.offload_device}")
if device is not None:
linked.append((n, m))
continue
logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules out of {module_count} total modules")
if linked:
logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules, computing reallocation")
device_assignments = analyze_ggml_loading(self.model, node_instance.distorch_allocations)['device_assignments']
for device, layers in device_assignments.items():
logging.info(f"MultiGPU: GGUFDisTorch - Moving {len(layers)} layers to {device}")
target_device = torch.device(device)
#logging.info(f"MultiGPU: GGUFDisTorch - Moving {len(layers)} layers to {device}")
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
logging.info("MultiGPU: GGUFDisTorch - self.mmap_released = True")
if hasattr(self, 'model'):
debug_hash = create_model_hash(self, "patcher")
debug_allocations = model_allocation_store.get(debug_hash)
if debug_allocations:
device_assignments = analyze_ggml_loading(self.model, debug_allocations)['device_assignments']
for device, layers in device_assignments.items():
target_device = torch.device(device)
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
module.GGUFModelPatcher.load = new_load
module.GGUFModelPatcher._patched = True
logging.info("MultiGPU: GGUFDisTorch - Successfully patched GGUF ModelPatcher")
else:
logging.info("MultiGPU: GGUFDisTorch - GGUF ModelPatcher already patched")
def analyze_ggml_loading(model, distorch_allocations):
@@ -191,7 +198,6 @@ def override_class_with_distorch(cls):
class NodeOverrideDisTorch(cls):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.distorch_allocations = {}
self.distorch_compute_device = None
@classmethod
@@ -217,21 +223,30 @@ def override_class_with_distorch(cls):
def override(self, *args, **kwargs):
self.distorch_allocations = {}
self.distorch_compute_device = kwargs.get("compute_device", None)
if self.distorch_compute_device is not None:
global current_device
current_device = self.distorch_compute_device
global current_device, model_allocation_store
register_patched_ggufmodelpatcher(self)
distorch_compute_device = kwargs.get("compute_device", None)
if distorch_compute_device is not None:
current_device = distorch_compute_device
for key, value in list(kwargs.items()):
register_patched_ggufmodelpatcher() # Removed node_instance argument
allocation_params = {}
keys_to_remove = list(kwargs.keys())
for key in keys_to_remove:
if key not in {"unet_name", "clip_name1", "clip_name2", "clip_name2", "type"}:
logging.info(f"MultiGPU: Removing {key} from kwargs")
logging.info(f"MultiGPU: Value: {value}")
self.distorch_allocations[key] = kwargs.pop(key)
value = kwargs.pop(key)
allocation_params[key] = value
fn = getattr(super(), cls.FUNCTION)
return fn(*args, **kwargs)
model = fn(*args, **kwargs)
if hasattr(model[0], 'model'):
model_hash = create_model_hash(model[0], "override")
model_allocation_store[model_hash] = allocation_params.copy()
elif hasattr(model[0], 'patcher') and hasattr(model[0].patcher, 'model'):
model_hash = create_model_hash(model[0].patcher, "override")
model_allocation_store[model_hash] = allocation_params.copy()
return model
return NodeOverrideDisTorch