Virtual VRAM "automatic" mode for DisTorch, WIP but working

This commit is contained in:
John Pollock
2025-02-07 15:05:08 -06:00
parent 5a403e638c
commit 3a4c6d50c8
+125 -11
View File
@@ -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