feat: Add byte-based model allocation mode

Introduces a new expert allocation mode allowing users to define model distribution using absolute memory values (e.g., "8g", "512m"). This provides more direct and predictable control over how a model is split across devices compared to the percentage-based method.

The new allocation string format is `device,size;device,size;...`, for example: `"cuda:0,8g;cuda:1,4g;cpu*,2g"`.

Key features:
- A wildcard `*` designates a device to receive any remaining unallocated model parts.
- If requested allocations exceed the model size, they are pro-rated down.
- A new `parse_memory_string` utility handles flexible memory unit parsing (g, m, k, b).

Additionally, the device allocation summary table has been improved to be more descriptive, now showing total VRAM, percentage of device VRAM used, absolute model GB allocated, and the model distribution percentage.
This commit is contained in:
John Pollock
2025-08-26 06:45:55 -05:00
parent c58ffaeb05
commit 56b8dd233e
+151 -30
View File
@@ -7,6 +7,7 @@ import sys
import torch
import logging
import hashlib
import re
logger = logging.getLogger("MultiGPU")
import copy
@@ -143,13 +144,11 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
distorch_alloc = calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str)
elif any(c in distorch_alloc.lower() for c in ['g', 'm', 'k', 'b']):
mode = "byte"
distorch_alloc = calculate_fraction_from_byte_expert_string(model_patcher, distorch_alloc)
elif "%" in distorch_alloc:
mode = "ratio"
distorch_alloc = calculate_fraction_from_ratio_expert_string(model_patcher, distorch_alloc)
logger.info(f"[MultiGPU_DisTorch2] Detected allocation mode: {mode}")
eq_line = "=" * 50
dash_line = "-" * 50
fmt_assign = "{:<18}{:>7}{:>14}{:>10}"
@@ -168,20 +167,36 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
"alloc_gb": alloc_gb
}
# Final Allocation Table
logger.info(eq_line)
logger.info(" DisTorch2 Model Device Allocations")
logger.info(eq_line)
logger.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
fmt_rosetta = "{:<8}{:>9}{:>9}{:>11}{:>10}"
logger.info(fmt_rosetta.format("Device", "VRAM GB", "Dev %", "Model GB", "Dist %"))
logger.info(dash_line)
sorted_devices = sorted(device_table.keys(), key=lambda d: (d == "cpu", d))
# Calculate total allocated model size for ratio calculation
total_allocated_model_bytes = sum(d["alloc_gb"] * (1024**3) for d in device_table.values())
for dev in sorted_devices:
frac = device_table[dev]["fraction"]
tot_gb = device_table[dev]["total_gb"]
total_dev_gb = device_table[dev]["total_gb"]
alloc_fraction = device_table[dev]["fraction"]
alloc_gb = device_table[dev]["alloc_gb"]
logger.info(fmt_assign.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}"))
# Calculate the distribution ratio percentage
dist_ratio_percent = (alloc_gb * (1024**3) / total_allocated_model_bytes) * 100 if total_allocated_model_bytes > 0 else 0
logger.info(fmt_rosetta.format(
dev,
f"{total_dev_gb:.2f}",
f"{alloc_fraction*100:.1f}%",
f"{alloc_gb:.2f}",
f"{dist_ratio_percent:.1f}%"
))
logger.info(dash_line)
block_summary = {}
@@ -323,18 +338,116 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
"block_assignments": block_assignments
}
def parse_memory_string(mem_str):
"""Parses a memory string (e.g., '4.0g', '512M') and returns bytes."""
mem_str = mem_str.strip().lower()
match = re.match(r'(\d+\.?\d*)\s*([gmkb]?)', mem_str)
if not match:
raise ValueError(f"Invalid memory string format: {mem_str}")
val, unit = match.groups()
val = float(val)
if unit == 'g':
return val * (1024**3)
elif unit == 'm':
return val * (1024**2)
elif unit == 'k':
return val * 1024
else: # b or no unit
return val
def calculate_fraction_from_byte_expert_string(model_patcher, byte_str):
"""
Converts a user-provided byte string (which describes how to split the MODEL)
into a fraction string (which describes the fraction of DEVICE VRAM to use).
"""
raw_block_list = model_patcher._load_list()
total_model_memory = sum(module_size for module_size, _, _, _ in raw_block_list)
raw_parsed = {}
wildcard_device = "cpu"
for allocation in byte_str.split(';'):
if ',' not in allocation: continue
dev_name, val_str = allocation.split(',', 1)
if '*' in dev_name:
dev_name = dev_name.replace('*','').strip()
wildcard_device = dev_name
raw_parsed[dev_name] = parse_memory_string(val_str)
# Handle allocation logic
total_requested_bytes = sum(raw_parsed.values())
final_allocations = {}
if total_requested_bytes > total_model_memory:
logger.info(f"[MultiGPU_DisTorch2] Over-allocation: Requested {total_requested_bytes/(1024**3):.2f}GB, but model is {total_model_memory/(1024**3):.2f}GB. Pro-rating allocations.")
for dev, val in raw_parsed.items():
final_allocations[dev] = (val / total_requested_bytes) * total_model_memory
else:
final_allocations = raw_parsed
remaining_bytes = total_model_memory - total_requested_bytes
if wildcard_device not in final_allocations:
final_allocations[wildcard_device] = 0
final_allocations[wildcard_device] += remaining_bytes
if remaining_bytes > 0:
logger.info(f"[MultiGPU_DisTorch2] Under-allocation: {remaining_bytes/(1024**2):.2f}MB of model unallocated. Assigning to wildcard device '{wildcard_device}'.")
# Convert byte allocations to fractions of device VRAM
allocation_parts = []
for dev, bytes_alloc in final_allocations.items():
total_device_vram = mm.get_total_memory(torch.device(dev))
if total_device_vram > 0:
fraction = bytes_alloc / total_device_vram
allocation_parts.append(f"{dev},{fraction:.4f}")
# Add user-facing logging
original_parts = []
original_wildcard_device = None
for allocation in byte_str.split(';'):
if ',' not in allocation: continue
dev_name, val_str = allocation.split(',', 1)
if '*' in dev_name:
dev_name = dev_name.replace('*','').strip()
original_wildcard_device = dev_name
original_parts.append((dev_name, val_str.strip()))
if original_parts:
formatted_parts = []
for dev_name, val_str in original_parts:
if 'mb' in val_str.lower():
mb_val = float(val_str.lower().replace('mb', ''))
gb_val = mb_val / 1024
formatted_parts.append(f"{gb_val:.2f}gb on {dev_name}")
elif 'gb' in val_str.lower() or 'g' in val_str.lower():
val_num = float(''.join(filter(lambda x: x.isdigit() or x == '.', val_str)))
formatted_parts.append(f"{val_num:.2f}gb on {dev_name}")
else:
formatted_parts.append(f"{val_str} on {dev_name}")
if formatted_parts:
if len(formatted_parts) == 1:
put_part = formatted_parts[0]
elif len(formatted_parts) == 2:
put_part = f"{formatted_parts[0]} and {formatted_parts[1]}"
else:
put_part = ", ".join(formatted_parts[:-1]) + f", and {formatted_parts[-1]}"
wildcard_dev = original_wildcard_device if original_wildcard_device else "cpu"
logger.info(f"[MultiGPU_DisTorch2] Direct(byte) Mode - {byte_str} -> '*' {wildcard_dev} = over/underflow device, put {put_part}")
result_string = ";".join(allocation_parts)
logger.info(f"[MultiGPU_DisTorch2] Converted byte string to fraction string: {result_string}")
return result_string
def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str):
"""
Converts a user-provided ratio string (which describes how to split the MODEL)
into a fraction string (which describes the fraction of DEVICE VRAM to use).
This is the correct bridge between the user-facing 'ratio' mode and the
internal 'fraction' system.
"""
# 1. Get the model's total size in bytes. This is what we are splitting.
raw_block_list = model_patcher._load_list()
total_model_memory = sum(module_size for module_size, _, _, _ in raw_block_list)
# 2. Parse the user's ratio string (e.g., "cuda:0,75;cpu,25") into a dictionary.
raw_ratios = {}
for allocation in ratio_str.split(';'):
if ',' not in allocation: continue
@@ -343,27 +456,36 @@ def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str):
value = float(val_str.replace('%','').strip())
raw_ratios[dev_name] = value
# 3. Sum the total ratio parts to normalize against (e.g., 75 + 25 = 100).
total_ratio_parts = sum(raw_ratios.values())
# 4. For each device, calculate the fraction of its VRAM required to hold its piece of the model.
allocation_parts = []
if total_ratio_parts > 0:
for dev, ratio_val in raw_ratios.items():
# a. Calculate how many bytes of the MODEL this device is responsible for.
# e.g., (75 / 100) * 10GB_model = 7.5GB of the model goes on this device.
bytes_of_model_for_device = (ratio_val / total_ratio_parts) * total_model_memory
# b. Get the total available VRAM for this specific device.
total_vram_of_device = mm.get_total_memory(torch.device(dev))
# c. The internal 'fraction' is the portion of the device's VRAM we need to use.
# e.g., 7.5GB_model_portion / 24GB_device_vram = 0.3125
if total_vram_of_device > 0:
required_fraction = bytes_of_model_for_device / total_vram_of_device
allocation_parts.append(f"{dev},{required_fraction:.4f}")
# 5. Return the newly constructed fraction string (e.g., "cuda:0,0.3125;cpu,0.0195").
for dev, ratio_val in raw_ratios.items():
bytes_of_model_for_device = (ratio_val / total_ratio_parts) * total_model_memory
total_vram_of_device = mm.get_total_memory(torch.device(dev))
if total_vram_of_device > 0:
required_fraction = bytes_of_model_for_device / total_vram_of_device
allocation_parts.append(f"{dev},{required_fraction:.4f}")
ratio_values = [str(v) for v in raw_ratios.values()]
ratio_string = ":".join(ratio_values)
normalized_pcts = [(v / total_ratio_parts) * 100 for v in raw_ratios.values()]
put_parts = []
for i, dev_name in enumerate(raw_ratios.keys()):
put_parts.append(f"{int(normalized_pcts[i])}% on {dev_name}")
if len(put_parts) == 1:
put_part = put_parts[0]
elif len(put_parts) == 2:
put_part = f"{put_parts[0]} and {put_parts[1]}"
else:
put_part = ", ".join(put_parts[:-1]) + f", and {put_parts[-1]}"
logger.info(f"[MultiGPU_DisTorch2] Ratio(%) Mode - {ratio_str} -> {ratio_string} ratio, put {put_part}")
result_string = ";".join(allocation_parts)
logger.info(f"[MultiGPU_DisTorch2] Converted ratio string to fraction string: {result_string}")
return result_string
@@ -373,7 +495,6 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
recipient_device, vram_amount, donors = virtual_vram_str.split(';')
virtual_vram_gb = float(vram_amount)
# EXACT SAME FORMATTING AS GGML
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<8} {:<6} {:>11} {:>9} {:>9}"