From 2481b520644b423d4206de263d16191c1c6e3138 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Tue, 26 Aug 2025 16:43:04 -0500 Subject: [PATCH] Refactor: Improve block allocation and expert string parsing This commit refactors the DisTorch safetensor loading and allocation logic for improved performance and correctness. The main changes to the block assignment are: - The primary compute device is now included in the pool of "donor" devices, allowing for more holistic memory quota calculation across all available GPUs. - Unassigned "orphan" blocks are now allocated to the compute device instead of the CPU. This keeps more of the model in VRAM, reducing potential bottlenecks. Additionally, this commit: - Fixes a bug in the byte expert string parser where the wildcard `*` was incorrectly checked in the device name instead of the value. - Standardizes variable names like `allocations_string` for better code clarity and consistency. --- distorch_2.py | 31 +++++++++++++------------------ 1 file changed, 13 insertions(+), 18 deletions(-) diff --git a/distorch_2.py b/distorch_2.py index db5bdf8..9725ad0 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -130,16 +130,16 @@ def register_patched_safetensor_modelpatcher(): logger.info("[MultiGPU_DisTorch2] Successfully patched ModelPatcher.partially_load") -def analyze_safetensor_loading(model_patcher, allocations_str): +def analyze_safetensor_loading(model_patcher, allocations_string): """ Analyze and distribute safetensor model blocks across devices """ DEVICE_RATIOS_DISTORCH = {} device_table = {} - distorch_alloc = allocations_str + distorch_alloc = allocations_string virtual_vram_gb = 0.0 - distorch_alloc, virtual_vram_str = allocations_str.split('#') + distorch_alloc, virtual_vram_str = allocations_string.split('#') compute_device = virtual_vram_str.split(';')[0] logger.info(f"[MultiGPU_DisTorch2] Compute Device: {compute_device}") @@ -260,7 +260,7 @@ def analyze_safetensor_loading(model_patcher, allocations_str): block_assignments = {} # Create a memory quota for each donor device based on its calculated allocation. - donor_devices = [d for d in sorted_devices if d != compute_device] + donor_devices = [d for d in sorted_devices] donor_quotas = { dev: device_table[dev]["alloc_gb"] * (1024**3) for dev in donor_devices @@ -276,8 +276,8 @@ def analyze_safetensor_loading(model_patcher, allocations_str): assigned_to_donor = True break # Move to the next block - if not assigned_to_donor: - block_assignments[block_name] = "cpu" + if not assigned_to_donor: #Note - small rounding errors and tensor-fitting on devices make a block occasionally an orphan. We treat orphans the same as tiny_block_list as they are generally small rounding errors + block_assignments[block_name] = compute_device if tiny_block_list: for block_name, module, block_type, block_memory in tiny_block_list: @@ -372,10 +372,9 @@ def calculate_fraction_from_byte_expert_string(model_patcher, byte_str): if ',' not in allocation: continue dev_name, val_str = allocation.split(',', 1) - is_wildcard = '*' in dev_name + is_wildcard = '*' in 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: @@ -414,9 +413,9 @@ def calculate_fraction_from_byte_expert_string(model_patcher, byte_str): fraction = bytes_alloc / total_device_vram allocation_parts.append(f"{dev},{fraction:.4f}") - result_string = ";".join(allocation_parts) + allocations_string = ";".join(allocation_parts) - return result_string + return allocations_string def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str): """ @@ -464,9 +463,9 @@ def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str): logger.info(f"[MultiGPU_DisTorch2] Ratio(%) Mode - {ratio_str} -> {ratio_string} ratio, put {put_part}") - result_string = ";".join(allocation_parts) + allocations_string = ";".join(allocation_parts) - return result_string + return allocations_string def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str): """Calculate virtual VRAM allocation string for distributed safetensor loading""" @@ -545,12 +544,8 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str): donor_percent = donor_allocations[donor] / donor_vram allocation_parts.append(f"{donor},{donor_percent:.4f}") - allocation_string = ";".join(allocation_parts) - - fmt_mem = "{:<20}{:>20}" - logger.info(fmt_mem.format("[MultiGPU_DisTorch2] Virtual VRAM Expert String", allocation_string)) - - return allocation_string + allocations_string = ";".join(allocation_parts) + return allocations_string def override_class_with_distorch_safetensor_v2(cls): """DisTorch 2.0 wrapper for safetensor models"""