Refactor: Improve DisTorch2 allocation logic and robustness

This commit refactors several aspects of the DisTorch2 device allocation logic to make it more robust, predictable, and easier to debug.

Key changes:
- Rework the byte-based allocation string parser (`calculate_fraction_from_byte_expert_string`). The new implementation correctly respects the user-defined device order and more cleanly handles the wildcard (`*`) for assigning remaining model parts.
-revert the "improvements" to the analyze safetensor loading routine causing it to catestrophically fail
This commit is contained in:
John Pollock
2025-08-26 11:13:38 -05:00
parent 6c2a3d5b15
commit 40cccdf01d
+61 -78
View File
@@ -149,6 +149,8 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
mode = "ratio"
distorch_alloc = calculate_fraction_from_ratio_expert_string(model_patcher, distorch_alloc)
logger.info(f"[MultiGPU_DisTorch2] Final Allocation String: {distorch_alloc}")
eq_line = "=" * 50
dash_line = "-" * 50
fmt_assign = "{:<18}{:>7}{:>14}{:>10}"
@@ -176,17 +178,15 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
logger.info(fmt_rosetta.format("Device", "VRAM GB", "Dev %", "Model GB", "Dist %"))
logger.info(dash_line)
from .nodes import get_device_list
all_devices_list = get_device_list()
sorted_devices = sorted(all_devices_list, key=lambda d: (d == "cpu", d))
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:
total_dev_gb = mm.get_total_memory(torch.device(dev)) / (1024**3)
alloc_fraction = device_table.get(dev, {}).get("fraction", 0.0)
alloc_gb = device_table.get(dev, {}).get("alloc_gb", 0.0)
total_dev_gb = device_table[dev]["total_gb"]
alloc_fraction = device_table[dev]["fraction"]
alloc_gb = device_table[dev]["alloc_gb"]
# 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
@@ -247,18 +247,21 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
logger.info(dash_line)
# Distribute blocks sequentially from the tail of the model
from .nodes import get_device_list
all_devices = get_device_list()
device_assignments = {dev: [] for dev in all_devices}
device_assignments = {device: [] for device in DEVICE_RATIOS_DISTORCH.keys()}
block_assignments = {}
compute_device = str(current_device)
# Determine the primary compute device (first non-cpu device)
compute_device = "cuda:0" # Fallback
for dev in sorted_devices:
if dev != "cpu":
compute_device = dev
break
# Create a memory quota for each donor device based on its calculated allocation.
donor_devices = [d for d in all_devices_list if d != compute_device]
donor_devices = [d for d in sorted_devices if d != compute_device]
donor_quotas = {
dev: device_table.get(dev, {}).get("alloc_gb", 0.0) * (1024**3)
for dev in all_devices_list
dev: device_table[dev]["alloc_gb"] * (1024**3)
for dev in donor_devices
}
# Iterate from the TAIL of the model, assigning blocks to donors until their quotas are filled.
@@ -272,10 +275,9 @@ def analyze_safetensor_loading(model_patcher, allocations_str):
assigned_to_donor = True
break # Move to the next block
# If no donor had enough quota, assign it to the CPU as a fallback.
# If no donor had enough quota, assign it to the primary compute device.
if not assigned_to_donor:
block_assignments[block_name] = "cpu"
logger.info(f"[MultiGPU_DisTorch2] WARNING: Unaccounted for block '{block_name}' fell back to CPU. This may indicate a malformed allocation string.")
block_assignments[block_name] = compute_device
# Explicitly assign tiny blocks to the compute device
if tiny_block_list:
@@ -359,85 +361,66 @@ def parse_memory_string(mem_str):
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).
Converts a user-provided byte string (e.g., "cuda:1,4gb;cpu,*") into a
fractional VRAM allocation string that the main assignment logic can use.
This function strictly respects device order and byte quotas.
"""
raw_block_list = model_patcher._load_list()
total_model_memory = sum(module_size for module_size, _, _, _ in raw_block_list)
remaining_model_bytes = total_model_memory
# Use a list of tuples to preserve the user-defined order
parsed_allocations = []
wildcard_device = "cpu" # Default wildcard device
raw_parsed = {}
wildcard_device = "cpu"
for allocation in byte_str.split(';'):
if ',' not in allocation: continue
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
is_wildcard = '*' in dev_name
raw_parsed[dev_name] = parse_memory_string(val_str)
if is_wildcard:
dev_name = dev_name.replace('*', '').strip()
wildcard_device = dev_name
# Don't add wildcard to the priority list yet
else:
byte_val = parse_memory_string(val_str)
parsed_allocations.append({'device': dev_name, 'bytes': byte_val})
# Handle allocation logic
total_requested_bytes = sum(raw_parsed.values())
final_allocations = {}
final_byte_allocations = defaultdict(int)
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}'.")
# Process devices with specific byte allocations first, in order
for alloc in parsed_allocations:
dev = alloc['device']
requested_bytes = alloc['bytes']
# Convert byte allocations to fractions of device VRAM
# Determine the actual bytes to allocate to this device
bytes_to_assign = min(requested_bytes, remaining_model_bytes)
if bytes_to_assign > 0:
final_byte_allocations[dev] = bytes_to_assign
remaining_model_bytes -= bytes_to_assign
logger.info(f"[MultiGPU_DisTorch2] Assigning {bytes_to_assign / (1024**2):.2f}MB of model to {dev} (requested {requested_bytes / (1024**2):.2f}MB).")
if remaining_model_bytes <= 0:
logger.info("[MultiGPU_DisTorch2] All model blocks have been allocated. Subsequent devices in the string will receive no assignment.")
break
# Assign any leftover model bytes to the wildcard device
if remaining_model_bytes > 0:
final_byte_allocations[wildcard_device] += remaining_model_bytes
logger.info(f"[MultiGPU_DisTorch2] Assigning remaining {remaining_model_bytes / (1024**2):.2f}MB of model to wildcard device '{wildcard_device}'.")
# Convert the final byte allocations to VRAM fractions
allocation_parts = []
for dev, bytes_alloc in final_allocations.items():
for dev, bytes_alloc in final_byte_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}")
logger.info(f"[MultiGPU_DisTorch2] Converted byte string '{byte_str}' to final fraction string: '{result_string}'")
return result_string
def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str):