Implement model hash creation and allocation storage for MultiGPU support in GGUF model patcher
This commit is contained in:
+52
-6
@@ -5,6 +5,30 @@ from collections import defaultdict
|
||||
import comfy.model_management
|
||||
import copy
|
||||
|
||||
model_allocation_store = {}
|
||||
|
||||
def create_model_hash(model, caller):
|
||||
import hashlib
|
||||
|
||||
logging.info(f"\nMultiGPU: Hash creation from {caller}")
|
||||
|
||||
# Model type, size, and first 3 layers as identifier
|
||||
model_type = type(model.model).__name__
|
||||
model_size = model.model_size()
|
||||
first_layers = str(list(model.model_state_dict().keys())[:3])
|
||||
|
||||
logging.info(f"MultiGPU: Type: {model_type}")
|
||||
logging.info(f"MultiGPU: Size: {model_size}")
|
||||
logging.info(f"MultiGPU: Layer sample: {first_layers}")
|
||||
|
||||
identifier = f"{model_type}_{model_size}_{first_layers}"
|
||||
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
|
||||
logging.info(f"MultiGPU: Hash: {final_hash}")
|
||||
|
||||
return final_hash
|
||||
|
||||
# Add to both places as specified
|
||||
|
||||
def register_patched_ggufmodelpatcher(node_instance):
|
||||
from nodes import NODE_CLASS_MAPPINGS
|
||||
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
|
||||
@@ -15,15 +39,17 @@ def register_patched_ggufmodelpatcher(node_instance):
|
||||
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}")
|
||||
#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
|
||||
@@ -33,12 +59,21 @@ def register_patched_ggufmodelpatcher(node_instance):
|
||||
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")
|
||||
#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")
|
||||
#logging.info(f"MultiGPU: GGUFDisTorch - Found {len(linked)} linked modules, computing reallocation")
|
||||
if hasattr(self, 'model'):
|
||||
debug_hash = create_model_hash(self, "patcher")
|
||||
debug_allocations = model_allocation_store.get(debug_hash)
|
||||
logging.info(f"MultiGPU: Hash lookup - Found allocations: {debug_allocations}")
|
||||
logging.info(f"MultiGPU: LOOKUP - Hash {debug_hash}")
|
||||
if debug_allocations:
|
||||
logging.info(f"MultiGPU: FOUND - Hash matches, using allocations: {debug_allocations}")
|
||||
else:
|
||||
logging.info(f"MultiGPU: MISS - Hash not found in store")
|
||||
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}")
|
||||
#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:
|
||||
@@ -217,10 +252,11 @@ def override_class_with_distorch(cls):
|
||||
|
||||
|
||||
def override(self, *args, **kwargs):
|
||||
global current_device, model_allocation_store
|
||||
|
||||
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
|
||||
|
||||
register_patched_ggufmodelpatcher(self)
|
||||
@@ -232,6 +268,16 @@ def override_class_with_distorch(cls):
|
||||
self.distorch_allocations[key] = kwargs.pop(key)
|
||||
|
||||
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] = self.distorch_allocations.copy()
|
||||
logging.info(f"MultiGPU: STORE - Hash {model_hash}, Allocations: {model_allocation_store[model_hash]}")
|
||||
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] = self.distorch_allocations.copy()
|
||||
logging.info(f"MultiGPU: STORE - Hash {model_hash}, Allocations: {model_allocation_store[model_hash]}")
|
||||
return model
|
||||
|
||||
return NodeOverrideDisTorch
|
||||
Reference in New Issue
Block a user