From 3a4c6d50c8d795fa00aff9f534d3cb5abcbdfae5 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Fri, 7 Feb 2025 15:05:08 -0600 Subject: [PATCH] Virtual VRAM "automatic" mode for DisTorch, WIP but working --- __init__.py | 136 +++++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 125 insertions(+), 11 deletions(-) diff --git a/__init__.py b/__init__.py index 8559761..bd2b549 100644 --- a/__init__.py +++ b/__init__.py @@ -99,8 +99,21 @@ def register_patched_ggufmodelpatcher(): def analyze_ggml_loading(model, allocations_str): DEVICE_RATIOS_DISTORCH = {} device_table = {} + distorch_alloc = allocations_str + virtual_vram_gb = 0.0 - for allocation in allocations_str.split(';'): + if '#' in allocations_str: + distorch_alloc, virtual_vram_str = allocations_str.split('#') + virtual_vram_gb = float(virtual_vram_str.split(';')[1]) + distorch_alloc = calculate_vvram_allocation_string(model, virtual_vram_str) + + eq_line = "=" * 47 + dash_line = "-" * 47 + fmt_assign = "{:<12}{:>10}{:>14}{:>10}" + + logging.info(dash_line) + + for allocation in distorch_alloc.split(';'): dev_name, fraction = allocation.split(',') fraction = float(fraction) total_mem_bytes = mm.get_total_memory(torch.device(dev_name)) @@ -112,9 +125,6 @@ def analyze_ggml_loading(model, allocations_str): "alloc_gb": alloc_gb } - eq_line = "=" * 47 - dash_line = "-" * 47 - fmt_alloc = "{:<12}{:>10}{:>14}{:>10}" logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') logging.info(eq_line) logging.info(" DisTorch Analysis") @@ -122,7 +132,7 @@ def analyze_ggml_loading(model, allocations_str): logging.info(dash_line) logging.info(" DisTorch Device Allocations") logging.info(dash_line) - logging.info(fmt_alloc.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)")) + logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)")) logging.info(dash_line) sorted_devices = sorted(device_table.keys(), key=lambda d: (d == "cpu", d)) @@ -131,7 +141,7 @@ def analyze_ggml_loading(model, allocations_str): frac = device_table[dev]["fraction"] tot_gb = device_table[dev]["total_gb"] alloc_gb = device_table[dev]["alloc_gb"] - logging.info(fmt_alloc.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}")) + logging.info(fmt_assign.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}")) logging.info(dash_line) @@ -209,6 +219,104 @@ def analyze_ggml_loading(model, allocations_str): return {"device_assignments": device_assignments} +def calculate_vvram_allocation_string(model, virtual_vram_str): + + recipient_device, vram_amount, donors = virtual_vram_str.split(';') + virtual_vram_gb = float(vram_amount) + + eq_line = "=" * 47 + dash_line = "-" * 47 + fmt_assign = "{:<8} {:<6} {:<11} {:<7} {:<10}" + + logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') + logging.info(eq_line) + logging.info(" DisTorch Analysis") + logging.info(eq_line) + logging.info(dash_line) + logging.info(" DisTorch View VRAM Analysis") + logging.info(dash_line) + logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virtual(GB)")) + logging.info(dash_line) + + recipient_vram = mm.get_total_memory(torch.device(recipient_device)) / (1024**3) + recipient_virtual = recipient_vram + virtual_vram_gb + + logging.info(fmt_assign.format(recipient_device, 'recip', f" {recipient_vram:.2f}GB", f" {recipient_virtual:.2f}GB", f" +{virtual_vram_gb:.2f}GB")) + + ram_donors = [d for d in donors.split(',') if d != 'cpu'] + has_cpu = 'cpu' in donors.split(',') + donation_per_donor = virtual_vram_gb / (len(ram_donors) + (1 if has_cpu else 0)) + + donor_device_info = {} + for donor in ram_donors: + donor_vram = mm.get_total_memory(torch.device(donor)) / (1024**3) + donor_virtual = donor_vram - donation_per_donor + donor_device_info[donor] = (donor_vram, donor_virtual) + logging.info(fmt_assign.format(donor, 'donor', f" {donor_vram:.2f}GB", f" {donor_virtual:.2f}GB", f" -{donation_per_donor:.2f}GB")) + + if has_cpu: + system_dram_gb = mm.get_total_memory(torch.device('cpu')) / (1024**3) + cpu_virtual = system_dram_gb - donation_per_donor + logging.info(fmt_assign.format('cpu', 'donor', f" {system_dram_gb:.2f}GB", f" {cpu_virtual:.2f}GB", f" -{donation_per_donor:.2f}GB")) + + logging.info(dash_line) + + + layer_summary = {} + layer_list = [] + memory_by_type = defaultdict(int) + total_memory = 0 + + for name, module in model.named_modules(): + if hasattr(module, "weight"): + layer_type = type(module).__name__ + layer_summary[layer_type] = layer_summary.get(layer_type, 0) + 1 + layer_list.append((name, module, layer_type)) + layer_memory = 0 + if module.weight is not None: + layer_memory += module.weight.numel() * module.weight.element_size() + if hasattr(module, "bias") and module.bias is not None: + layer_memory += module.bias.numel() * module.bias.element_size() + memory_by_type[layer_type] += layer_memory + total_memory += layer_memory + + model_size_gb = total_memory / (1024**3) + + new_model_size_gb = model_size_gb - virtual_vram_gb + + logging.info(fmt_assign.format('model', 'model', f" {model_size_gb:.2f}GB", f" {new_model_size_gb:.2f}GB", f" -{virtual_vram_gb:.2f}GB")) + + if model_size_gb > (recipient_vram*0.9): + on_recipient = recipient_vram*0.9 + on_virtuals = model_size_gb - on_recipient + logging.info("Warning: Model size is greater than 90% of recipient VRAM.", on_virtuals, "GB of GGML Layers Offloaded Automatically to Virtual VRAM.") + else: + on_recipient = model_size_gb + on_virtuals = 0 + + new_on_recipient = max(0, on_recipient - virtual_vram_gb) + new_on_virtuals = min(model_size_gb, virtual_vram_gb / (len(ram_donors) + (1 if has_cpu else 0))) + + allocation_parts = [] + recipient_percent = new_on_recipient / recipient_vram if recipient_vram > 0 else 0.0 + allocation_parts.append(f"{recipient_device},{recipient_percent:.4f}") + + for donor in ram_donors: + donor_vram = donor_device_info[donor][0] + donor_percent = new_on_virtuals / donor_vram if donor_vram > 0 else 0.0 + allocation_parts.append(f"{donor},{donor_percent:.4f}") + if has_cpu: + cpu_percent = new_on_virtuals / system_dram_gb if system_dram_gb > 0 else 0.0 + allocation_parts.append(f"cpu,{cpu_percent:.4f}") + + allocation_string = ";".join(allocation_parts) + logging.info(dash_line) + fmt_mem = "{:<20}{:>20}" + logging.info(fmt_mem.format("Allocation String", allocation_string)) + logging.info(dash_line) + + return allocation_string + def get_device_list(): import torch return ["cpu"] + [f"cuda:{i}" for i in range(torch.cuda.device_count())] @@ -431,27 +539,33 @@ def override_class_with_distorch(cls): default_device = devices[1] if len(devices) > 1 else devices[0] inputs["optional"] = inputs.get("optional", {}) inputs["optional"]["device"] = (devices, {"default": default_device}) - inputs["optional"]["allocations"] = ("STRING", {"multiline": False, "default": "cuda:0,0.15;cpu,0.5"}) + inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 0.0, "min": 4.0, "max": 24.0, "step": 0.1}) + inputs["optional"]["allocations"] = ("STRING", {"multiline": False, "default": ""}) return inputs CATEGORY = "multigpu" FUNCTION = "override" - def override(self, *args, device=None, allocations=None, **kwargs): + def override(self, *args, device=None, allocations=None, virtual_vram_gb=0.0, **kwargs): global current_device if device is not None: current_device = device - + register_patched_ggufmodelpatcher() fn = getattr(super(), cls.FUNCTION) out = fn(*args, **kwargs) + vram_string = f"{device};{virtual_vram_gb};cpu" if virtual_vram_gb > 0 else "" + full_allocation = f"{allocations}#{vram_string}" if allocations or vram_string else "" + + logging.info(f"[DisTorch] Full allocation string: {full_allocation}") + if hasattr(out[0], 'model'): model_hash = create_model_hash(out[0], "override") - model_allocation_store[model_hash] = allocations + model_allocation_store[model_hash] = full_allocation elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): model_hash = create_model_hash(out[0].patcher, "override") - model_allocation_store[model_hash] = allocations + model_allocation_store[model_hash] = full_allocation return out