diff --git a/distorch.py b/distorch.py index ddbc863..601f44c 100644 --- a/distorch.py +++ b/distorch.py @@ -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 \ No newline at end of file