From afd8fecd943d09f5e031c801ecbd38f2797a6d4e Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sun, 24 Aug 2025 05:37:04 -0500 Subject: [PATCH] Refactor DisTorch model patching logic for improved device assignment and FP8 casting --- distorch_2.py | 118 ++++++++++++++++---------------------------------- 1 file changed, 38 insertions(+), 80 deletions(-) diff --git a/distorch_2.py b/distorch_2.py index 30182e0..ecdd112 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -83,94 +83,52 @@ def register_patched_safetensor_modelpatcher(): # Parse allocation string and apply static assignment device_assignments = analyze_safetensor_loading(self, allocations) - model_type = type(self.model).__name__ + high_precision_loras = self.model._distorch_high_precision_loras + loading = self._load_list() + loading.sort(reverse=True) + for module_size, module_name, module_object, params in loading: + # Step 1: Write block/tensor to compute device first + module_object.to(device_to) - if model_type == "SDXL" or model_type == "LTXV": - on_compute_patching = False - else: - on_compute_patching = True + # Step 2: Apply LoRa patches while on compute device + weight_key = "{}.weight".format(module_name) + bias_key = "{}.bias".format(module_name) - if on_compute_patching: - high_precision_loras = self.model._distorch_high_precision_loras - loading = self._load_list() - loading.sort(reverse=True) - for module_size, module_name, module_object, params in loading: - # Step 1: Write block/tensor to compute device first - module_object.to(device_to) + if weight_key in self.patches: + self.patch_weight_to_device(weight_key, device_to=device_to) + if weight_key in self.weight_wrapper_patches: + module_object.weight_function.extend(self.weight_wrapper_patches[weight_key]) - # Step 2: Apply LoRa patches while on compute device - weight_key = "{}.weight".format(module_name) - bias_key = "{}.bias".format(module_name) + if bias_key in self.patches: + self.patch_weight_to_device(bias_key, device_to=device_to) + if bias_key in self.weight_wrapper_patches: + module_object.bias_function.extend(self.weight_wrapper_patches[bias_key]) - if weight_key in self.patches: - self.patch_weight_to_device(weight_key, device_to=device_to) - if weight_key in self.weight_wrapper_patches: - module_object.weight_function.extend(self.weight_wrapper_patches[weight_key]) + # Step 3: FP8 casting for CPU storage (if enabled) + block_target_device = device_assignments['block_assignments'].get(module_name, device_to) + has_patches = weight_key in self.patches or bias_key in self.patches - if bias_key in self.patches: - self.patch_weight_to_device(bias_key, device_to=device_to) - if bias_key in self.weight_wrapper_patches: - module_object.bias_function.extend(self.weight_wrapper_patches[bias_key]) + if not high_precision_loras and block_target_device == "cpu" and has_patches and model_original_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: + logger.info(f"[MultiGPU_DisTorch2] FP8 casting conditions met for {module_name}") + for param_name, param in module_object.named_parameters(): + if param.dtype.is_floating_point: + cast_data = comfy.float.stochastic_rounding(param.data, torch.float8_e4m3fn) + new_param = torch.nn.Parameter(cast_data.to(torch.float8_e4m3fn)) + new_param.requires_grad = param.requires_grad + setattr(module_object, param_name, new_param) + logger.debug(f"[MultiGPU_DisTorch2] Cast {module_name}.{param_name} to FP8 for CPU storage") - # Step 3: FP8 casting for CPU storage (if enabled) - block_target_device = device_assignments['block_assignments'].get(module_name, device_to) - has_patches = weight_key in self.patches or bias_key in self.patches - - logger.info(f"[MultiGPU_DisTorch2] Patch-on-Compute: Processing {module_name} -> block_target_device={block_target_device}") + # Step 4: Move to ultimate destination based on DisTorch assignment + if block_target_device != device_to: + logger.debug(f"[MultiGPU_DisTorch2] Moving {module_name} from {device_to} to {block_target_device}") + module_object.to(block_target_device) + module_object.comfy_cast_weights = True - if not high_precision_loras and block_target_device == "cpu" and has_patches and model_original_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: - logger.info(f"[MultiGPU_DisTorch2] FP8 casting conditions met for {module_name}") - for param_name, param in module_object.named_parameters(): - if param.dtype.is_floating_point: - cast_data = comfy.float.stochastic_rounding(param.data, torch.float8_e4m3fn) - new_param = torch.nn.Parameter(cast_data.to(torch.float8_e4m3fn)) - new_param.requires_grad = param.requires_grad - setattr(module_object, param_name, new_param) - logger.debug(f"[MultiGPU_DisTorch2] Cast {module_name}.{param_name} to FP8 for CPU storage") + # Mark as patched and update memory counter + module_object.comfy_patched_weights = True + mem_counter += module_size - # Step 4: Move to ultimate destination based on DisTorch assignment - if block_target_device != device_to: - logger.debug(f"[MultiGPU_DisTorch2] Moving {module_name} from {device_to} to {block_target_device}") - module_object.to(block_target_device) - - # Mark as patched and update memory counter - module_object.comfy_patched_weights = True - mem_counter += module_size - - logger.info(f"[MultiGPU_DisTorch2] DisTorch loading completed. Total memory: {mem_counter / (1024 * 1024):.2f}MB") - - else: - # Apply our static assignments instead of ComfyUI's dynamic ones - for block_name, target_device in device_assignments['block_assignments'].items(): - # Find the module by name - parts = block_name.split('.') - module = self.model - for part in parts: - if hasattr(module, part): - module = getattr(module, part) - else: - break - - if hasattr(module, 'weight') or hasattr(module, 'comfy_cast_weights'): - # Move to our assigned device - logger.info(f"[MultiGPU_DisTorch2] Patch-on-Device: Moving {block_name} to {target_device}") - module.to(target_device) - # Mark for ComfyUI's cast system if not already marked - if hasattr(module, 'comfy_cast_weights'): - module.comfy_cast_weights = True - - weight_key = "{}.weight".format(block_name) - bias_key = "{}.bias".format(block_name) - - if weight_key in self.patches: - self.patch_weight_to_device(weight_key, device_to=target_device) - if weight_key in self.weight_wrapper_patches: - module_object.weight_function.extend(self.weight_wrapper_patches[weight_key]) - - if bias_key in self.patches: - self.patch_weight_to_device(bias_key, device_to=target_device) - if bias_key in self.weight_wrapper_patches: - module_object.bias_function.extend(self.weight_wrapper_patches[bias_key]) + logger.info(f"[MultiGPU_DisTorch2] DisTorch loading completed. Total memory: {mem_counter / (1024 * 1024):.2f}MB") return 0