From 996f298b4e501b7546516ef832db86b0230f39a0 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sat, 9 Aug 2025 17:51:15 -0500 Subject: [PATCH] feat: Implement robust block discovery and GGUF-style logging for SafeTensor DisTorch --- ARCHITECTURE_V2.0.0.md | 147 +++++++++-------------------------------- __init__.py | 125 ++++++++++++++++++++++++++++------- 2 files changed, 133 insertions(+), 139 deletions(-) diff --git a/ARCHITECTURE_V2.0.0.md b/ARCHITECTURE_V2.0.0.md index 5cc7c2f..3d9ab58 100755 --- a/ARCHITECTURE_V2.0.0.md +++ b/ARCHITECTURE_V2.0.0.md @@ -5,19 +5,19 @@ **PURPOSE**: This document captures the exact understanding and implementation plan for DisTorch SafeTensor, which generalizes the block-swap concept from WanVideoWrapper for any SafeTensor model. -**STATUS**: Implemented. The logic has been integrated into `__init__.py`. +**STATUS**: Implemented. The logic has been integrated into `__init__.py` with robust block detection and detailed logging. --- ## STEP 1: UNDERSTAND THE EXISTING APPROACHES ### 1A. ComfyUI-GGUF DisTorch Implementation -**File**: `ComfyUI-GGUF/__init__.py` and `gguf_model_patcher.py` -**Mechanism**: Distributes individual quantized layers across multiple devices. Dequantizes layers just-in-time for computation. Optimized for maximum memory savings with GGUF models. +**File**: `ComfyUI-GGUF/__init__.py` +**Mechanism**: Distributes individual quantized layers across multiple devices. Dequantizes layers just-in-time for computation. Optimized for maximum memory savings with GGUF models. Provides detailed, formatted logging tables. ### 1B. ComfyUI-WanVideoWrapper Block Swap **File**: `ComfyUI-WanVideoWrapper/nodes_model_loading.py` -**Mechanism**: Swaps entire, pre-defined model blocks (e.g., ResNet blocks, Attention blocks) between a compute device and a swap device during inference. It is highly effective but tailored specifically for the WanVideo model architecture. +**Mechanism**: Swaps entire, pre-defined model blocks between a compute device and a swap device. It is highly effective but tailored specifically for the WanVideo model architecture. ### 1C. ComfyUI-MultiGPU Integration **File**: `ComfyUI-MultiGPU/__init__.py` @@ -27,9 +27,9 @@ ## STEP 2: THE EXACT PROBLEM WE'RE SOLVING -Users have large SafeTensor models that do not fit into a single GPU's VRAM. We provide them with a solution that is more flexible than single-layer offloading and more general-purpose than WanVideo's integrated approach. +Users have large SafeTensor models (like SDXL, FLUX, etc.) that do not fit into a single GPU's VRAM. The goal is to provide a memory management solution that is flexible, model-agnostic, and provides clear, informative logging, on par with the GGUF DisTorch implementation. -**DisTorch SafeTensor (NEW)**: A memory management solution that intelligently swaps large, contiguous blocks of a model between a primary compute GPU and a secondary swap device (another GPU or system RAM). +**DisTorch SafeTensor (NEW)**: A memory management solution that intelligently discovers and swaps large, contiguous blocks of a model between a primary compute GPU and a secondary swap device (another GPU or system RAM). --- @@ -38,38 +38,26 @@ Users have large SafeTensor models that do not fit into a single GPU's VRAM. We ### The Wrapper Function The core of the implementation is the `override_class_with_distorch_safetensor` function, which wraps existing ComfyUI model loaders. -```python -def override_class_with_distorch_safetensor(cls): - """DisTorch wrapper for SafeTensor models, providing block-swap memory optimization.""" -``` - ### UI Parameters (What Users See) The node provides four key parameters to control the memory swapping behavior, ordered for intuitive use: 1. **`compute_device`**: The primary GPU where computations will occur (e.g., `cuda:0`). 2. **`compute_reserved_swap_gb`**: The amount of VRAM (in GB) to keep reserved on the `compute_device` for active blocks. This acts as a hot-cache. 3. **`virtualram_swap_device`**: The device to offload inactive blocks to (e.g., `cpu` or `cuda:1`). -4. **`virtualram_gb`**: The total size (in GB) of model blocks to offload to the `virtualram_swap_device`. This effectively creates "virtual VRAM" on your compute device. +4. **`virtualram_gb`**: The total size (in GB) of model blocks to offload to the `virtualram_swap_device`. -### The Math (20GB Model Example) -- **Model**: 20GB total size. -- **`compute_device`**: `cuda:0` (24GB VRAM) -- **`compute_reserved_swap_gb`**: `1.0` GB -- **`virtualram_swap_device`**: `cpu` -- **`virtualram_gb`**: `4.0` GB +### Intelligent Block Discovery +The `apply_block_swap` function now uses a multi-stage process to find swappable blocks, ensuring compatibility with various model architectures, including FLUX. +1. **Standard UNet Structure**: Checks for `input_blocks`, `middle_block`, `output_blocks`. +2. **Generic `blocks` Attribute**: Looks for a `model.blocks` or `diffusion_model.blocks` list. +3. **Generic `layers` Attribute**: Looks for a `model.layers` or `diffusion_model.layers` list. +4. **Fallback to ModuleList**: Scans for any top-level `torch.nn.ModuleList` as a last resort. -**Result**: -- **4GB** of the model's blocks are immediately moved to the `cpu`. -- The remaining **16GB** of blocks are loaded onto `cuda:0`. -- During inference, blocks are swapped as needed, but a buffer of at least **1GB** (`compute_reserved_swap_gb`) worth of blocks is kept on the compute device if possible. - -### Operation Flow -1. The user selects a model using a DisTorch-wrapped loader (e.g., `CheckpointLoaderSimpleDisTorchMultiGPU`). -2. The model is loaded normally by the underlying ComfyUI loader. -3. The `apply_block_swap` function analyzes the model to identify swappable blocks (e.g., input, middle, and output blocks of a UNet). -4. Based on `virtualram_gb`, a number of blocks are moved to the `virtualram_swap_device`. -5. The `forward` method of each offloaded block is patched with a hook. -6. When the model runs, the hook moves the required block to the `compute_device` just before it's needed and moves it back to the `virtualram_swap_device` afterward, using non-blocking transfers for efficiency. +### High-Quality Datalogging +A new `analyze_safetensor_distorch` function generates detailed, formatted tables in the console, identical in style to the GGUF implementation, showing: +- Device Allocations +- Block Analysis (Type, Count, Memory) +- Final Block Assignments --- @@ -78,89 +66,23 @@ The node provides four key parameters to control the memory swapping behavior, o The following code is a representation of the current implementation within `__init__.py`. ```python -def override_class_with_distorch_safetensor(cls): - """DisTorch wrapper for SafeTensor models, providing block-swap memory optimization.""" - - class NodeOverrideDisTorchSafeTensor(cls): - @classmethod - def INPUT_TYPES(s): - inputs = copy.deepcopy(cls.INPUT_TYPES()) - devices = get_device_list() - compute_device = devices[1] if len(devices) > 1 else devices[0] - - inputs["optional"] = inputs.get("optional", {}) - - # Reordered and renamed parameters - inputs["optional"]["compute_device"] = (devices, { - "default": compute_device, - "tooltip": "Primary device for computation." - }) - inputs["optional"]["compute_reserved_swap_gb"] = ("FLOAT", { - "default": 1.0, - "min": 0.1, - "max": 16.0, - "step": 0.1, - "tooltip": "GB of VRAM to keep reserved on the compute device." - }) - inputs["optional"]["virtualram_swap_device"] = (devices, { - "default": "cpu", - "tooltip": "Device to offload inactive model blocks to." - }) - inputs["optional"]["virtualram_gb"] = ("FLOAT", { - "default": 4.0, - "min": 0.1, - "max": 64.0, - "step": 0.1, - "tooltip": "Amount of VRAM (in GB) to offload to the swap device." - }) - return inputs - - CATEGORY = "multigpu" - FUNCTION = "override" - - def override(self, *args, compute_device=None, compute_reserved_swap_gb=1.0, - virtualram_swap_device="cpu", virtualram_gb=4.0, **kwargs): - global current_device - - logging.info(f"[DisTorch SafeTensor] Override called with: compute_device={compute_device}, swap_device={virtualram_swap_device}, virtualram_gb={virtualram_gb}, reserved_gb={compute_reserved_swap_gb}") - - if compute_device is not None: - current_device = compute_device - - fn = getattr(super(), cls.FUNCTION) - out = fn(*args, **kwargs) - - model = out[0] - if hasattr(model, 'model'): - logging.info("[DisTorch SafeTensor] Model has 'model' attribute, applying block swap.") - apply_block_swap( - model, - compute_device=compute_device, - swap_device=virtualram_swap_device, - virtual_vram_gb=virtualram_gb, - reserved_swap_gb=compute_reserved_swap_gb - ) - else: - logging.warning("[DisTorch SafeTensor] Loaded object does not have a 'model' attribute, skipping block swap.") - - return out - - return NodeOverrideDisTorchSafeTensor - +def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks): + """Provides a detailed analysis of the block swap configuration, mimicking the GGUF DisTorch style.""" + # ... (Full implementation in __init__.py) def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu", virtual_vram_gb=4.0, reserved_swap_gb=1.0): """ Applies WanVideo-style block swapping by patching the forward method of individual model blocks. """ - # ... (Full implementation in __init__.py) + # ... (Full implementation with intelligent block discovery in __init__.py) ``` --- ## STEP 5: REGISTRATION IN __init__.py -The new DisTorch SafeTensor wrappers are registered for all relevant core ComfyUI nodes. +The DisTorch SafeTensor wrappers are registered for all relevant core ComfyUI nodes under the `...DisTorchMultiGPU` alias. ```python # Register the new DisTorch SafeTensor wrappers @@ -175,27 +97,18 @@ NODE_CLASS_MAPPINGS["UNETLoaderDisTorchMultiGPU"] = override_class_with_distorch ### DisTorch (GGUF) - **Granularity**: Per-layer. -- **Use Case**: Maximum memory saving on GGUF models, often with CPU offload. -- **Implementation**: Complex allocation strings and quantization handling. +- **Logging**: High-quality, formatted tables. +- **Use Case**: Maximum memory saving on GGUF models. ### DisTorch SafeTensor (NEW) -- **Granularity**: Per-block. -- **Use Case**: Balancing memory and speed for any SafeTensor model. -- **Implementation**: Simple forward hooks, model-agnostic. +- **Granularity**: Per-block (intelligently discovered). +- **Logging**: High-quality, formatted tables (matches GGUF style). +- **Use Case**: Balancing memory and speed for **any** SafeTensor model. ### WanVideo Block Swap - **Granularity**: Per-block (model-specific). +- **Logging**: Basic. - **Use Case**: Optimized specifically for WanVideo models. -- **Implementation**: Integrated directly into the custom model's forward pass. - ---- - -## WHY THIS MATTERS - -1. **Flexibility**: Enables running models that are larger than a single GPU's VRAM. -2. **Control**: Users can fine-tune the memory vs. speed trade-off. -3. **Compatibility**: Works with all standard SafeTensor models loaded through core ComfyUI nodes. -4. **Simplicity**: All logic is self-contained within the `ComfyUI-MultiGPU` custom node. --- @@ -206,6 +119,8 @@ NODE_CLASS_MAPPINGS["UNETLoaderDisTorchMultiGPU"] = override_class_with_distorch - [X] Implement `override_class_with_distorch_safetensor` in `__init__.py` (DONE) - [X] Rename and reorder UI parameters (DONE) - [X] Expand coverage to all core ComfyUI nodes (DONE) +- [X] **Implement robust, multi-stage block discovery to support models like FLUX (DONE)** +- [X] **Overhaul logging to match the high-quality, formatted style of the GGUF DisTorch implementation (DONE)** - [ ] Test with SDXL checkpoint - [ ] Test with Flux checkpoint - [ ] Verify memory usage matches expectations diff --git a/__init__.py b/__init__.py index 410d253..adfb927 100644 --- a/__init__.py +++ b/__init__.py @@ -532,6 +532,69 @@ def override_class_with_distorch_safetensor(cls): return NodeOverrideDisTorchSafeTensor +def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks): + """Provides a detailed analysis of the block swap configuration, mimicking the GGUF DisTorch style.""" + + eq_line = "=" * 60 + dash_line = "-" * 60 + + logging.info(eq_line) + logging.info(" DisTorch SafeTensor Memory Analysis") + logging.info(eq_line) + + # Device Allocation Table + fmt_assign = "{:<12}{:>15}{:>15}{:>15}" + logging.info(fmt_assign.format("Device", "Role", "Total Mem (GB)", "Config (GB)")) + logging.info(dash_line) + + compute_total_gb = mm.get_total_memory(torch.device(compute_device)) / (1024**3) + swap_total_gb = mm.get_total_memory(torch.device(swap_device)) / (1024**3) + + logging.info(fmt_assign.format(compute_device, "Compute", f"{compute_total_gb:.2f}", f"Reserve: {reserved_swap_gb:.2f}")) + logging.info(fmt_assign.format(swap_device, "Swap", f"{swap_total_gb:.2f}", f"Offload: {virtual_vram_gb:.2f}")) + logging.info(dash_line) + + # Block Analysis Table + block_summary = defaultdict(lambda: {'count': 0, 'memory': 0}) + total_memory = 0 + + for block in all_blocks: + block_type = type(block).__name__ + block_memory = sum(p.numel() * p.element_size() for p in block.parameters()) + block_summary[block_type]['count'] += 1 + block_summary[block_type]['memory'] += block_memory + total_memory += block_memory + + logging.info(" DisTorch SafeTensor Block Analysis") + logging.info(dash_line) + fmt_layer = "{:<20}{:>10}{:>15}{:>12}" + logging.info(fmt_layer.format("Block Type", "Count", "Memory (MB)", "% Total")) + logging.info(dash_line) + + sorted_blocks = sorted(block_summary.items(), key=lambda x: x[1]['memory'], reverse=True) + + for block_type, data in sorted_blocks: + mem_mb = data['memory'] / (1024 * 1024) + mem_percent = (data['memory'] / total_memory) * 100 if total_memory > 0 else 0 + logging.info(fmt_layer.format(block_type, str(data['count']), f"{mem_mb:.2f}", f"{mem_percent:.1f}%")) + logging.info(dash_line) + + # Final Assignment Table + model_size_gb = total_memory / (1024**3) + block_size_gb = model_size_gb / len(all_blocks) if all_blocks else 0 + blocks_to_offload = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0 + blocks_on_compute = len(all_blocks) - blocks_to_offload + + logging.info(" DisTorch Final Block Assignments") + logging.info(dash_line) + fmt_final = "{:<20}{:>15}" + logging.info(fmt_final.format("Total Model Size (GB):", f"{model_size_gb:.2f}")) + logging.info(fmt_final.format("Average Block Size (MB):", f"{block_size_gb * 1024:.2f}" if all_blocks else "N/A")) + logging.info(dash_line) + logging.info(fmt_final.format("Blocks on Compute:", f"{blocks_on_compute}")) + logging.info(fmt_final.format("Blocks on Swap:", f"{blocks_to_offload}")) + logging.info(eq_line) + def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu", virtual_vram_gb=4.0, reserved_swap_gb=1.0): """ @@ -543,57 +606,73 @@ def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu", model_to_patch = None if hasattr(model_patcher, 'model') and hasattr(model_patcher.model, 'diffusion_model'): model_to_patch = model_patcher.model.diffusion_model + logging.info("[DisTorch SafeTensor] Found 'diffusion_model' attribute for patching.") elif hasattr(model_patcher, 'model'): model_to_patch = model_patcher.model + logging.info("[DisTorch SafeTensor] Found 'model' attribute for patching.") else: logging.error("[DisTorch SafeTensor] Could not find a valid model to patch for block swapping.") return all_blocks = [] - if hasattr(model_to_patch, 'input_blocks'): + # 1. Standard UNet Structure + if hasattr(model_to_patch, 'input_blocks') and hasattr(model_to_patch, 'middle_block') and hasattr(model_to_patch, 'output_blocks'): + logging.info("[DisTorch SafeTensor] Found standard UNet structure ('input_blocks', 'middle_block', 'output_blocks').") all_blocks.extend(model_to_patch.input_blocks) - if hasattr(model_to_patch, 'middle_block'): if isinstance(model_to_patch.middle_block, torch.nn.Module): all_blocks.append(model_to_patch.middle_block) - if hasattr(model_to_patch, 'output_blocks'): all_blocks.extend(model_to_patch.output_blocks) + # 2. Simple 'blocks' attribute + elif hasattr(model_to_patch, 'blocks') and isinstance(model_to_patch.blocks, torch.nn.ModuleList): + logging.info("[DisTorch SafeTensor] Found 'blocks' attribute of type ModuleList.") + all_blocks.extend(model_to_patch.blocks) + # 3. Simple 'layers' attribute + elif hasattr(model_to_patch, 'layers') and isinstance(model_to_patch.layers, torch.nn.ModuleList): + logging.info("[DisTorch SafeTensor] Found 'layers' attribute of type ModuleList.") + all_blocks.extend(model_to_patch.layers) + # 4. Fallback to top-level ModuleLists + else: + logging.info("[DisTorch SafeTensor] No standard structure found. Falling back to searching for top-level ModuleLists.") + for child in model_to_patch.children(): + if isinstance(child, torch.nn.ModuleList): + logging.info(f"[DisTorch SafeTensor] Found top-level ModuleList with {len(child)} modules. Adding them as blocks.") + all_blocks.extend(child) if not all_blocks: - logging.warning("[DisTorch SafeTensor] No swappable blocks were found in the model.") + logging.error("[DisTorch SafeTensor] CRITICAL: No swappable blocks were found in the model. Block swap cannot be applied.") return - logging.info(f"[DisTorch SafeTensor] Found {len(all_blocks)} swappable blocks in the model.") + logging.info(f"[DisTorch SafeTensor] Successfully identified {len(all_blocks)} swappable blocks.") + + # Run and display the analysis + analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks) model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3) - if len(all_blocks) > 0: - block_size_gb = model_size_gb / len(all_blocks) - blocks_to_offload = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0 - blocks_on_compute = len(all_blocks) - blocks_to_offload - else: - blocks_to_offload = 0 - blocks_on_compute = 0 - - logging.info(f"[DisTorch SafeTensor] Model size: {model_size_gb:.2f} GB, Avg block size: {block_size_gb * 1024:.2f} MB") - logging.info(f"[DisTorch SafeTensor] Offloading {blocks_to_offload} blocks to {swap_device}. Keeping {blocks_on_compute} blocks on {compute_device}.") + block_size_gb = model_size_gb / len(all_blocks) if all_blocks else 0 + blocks_to_offload = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0 + blocks_on_compute = len(all_blocks) - blocks_to_offload for i, block in enumerate(all_blocks): - if i < blocks_on_compute: - block.to(compute_device) - else: - block.to(swap_device) - + # Determine target device for this block + target_device = compute_device if i < blocks_on_compute else swap_device + block.to(target_device) + + # Patch the forward method only if the block is on the swap device + if target_device == swap_device: original_forward = block.forward - def create_patched_forward(original_f, b, cd, sd): + def create_patched_forward(original_f, b, block_index, cd, sd): def patched_forward(*args, **kwargs): + logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} to {cd} for computation.") b.to(cd, non_blocking=True) result = original_f(*args, **kwargs) + logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} back to {sd}.") b.to(sd, non_blocking=True) return result return patched_forward - block.forward = create_patched_forward(original_forward, block, torch.device(compute_device), torch.device(swap_device)) - logging.info(f"[DisTorch SafeTensor] Patched forward method for block {i}.") + block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device)) + logging.info(f"[DisTorch SafeTensor] Patched forward method for block {i} on {swap_device}.") logging.info("[DisTorch SafeTensor] Block swap setup complete.") # For backwards compatibility, keep the old name pointing to the new safetensor wrapper