Virtual VRAM "automatic" mode for DisTorch, WIP but working
This commit is contained in:
+125
-11
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user