From 2ac057d338ea7729ed6f17eabefdce78099e66e7 Mon Sep 17 00:00:00 2001 From: pollock Date: Fri, 24 Jan 2025 05:37:17 -0500 Subject: [PATCH] Implement model hash creation and allocation storage for MultiGPU support in GGUF model patcher --- distorch.py | 58 +++++++++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 52 insertions(+), 6 deletions(-) diff --git a/distorch.py b/distorch.py index ddbc863..cbbf24b 100644 --- a/distorch.py +++ b/distorch.py @@ -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 \ No newline at end of file