From 992b1c65871f3dce5e6dd6aeabc3853ae764f9fa Mon Sep 17 00:00:00 2001 From: John Pollock Date: Sat, 9 Aug 2025 10:46:05 -0500 Subject: [PATCH] docs: Update architecture document for DisTorch SafeTensor --- ARCHITECTURE_V2.0.0.md | 497 +++++++++++++++-------------------------- core/blockswap.py | 495 ---------------------------------------- 2 files changed, 183 insertions(+), 809 deletions(-) delete mode 100644 core/blockswap.py diff --git a/ARCHITECTURE_V2.0.0.md b/ARCHITECTURE_V2.0.0.md index 9f2d3ce..5cc7c2f 100755 --- a/ARCHITECTURE_V2.0.0.md +++ b/ARCHITECTURE_V2.0.0.md @@ -1,351 +1,220 @@ # ComfyUI-MultiGPU Architecture V2.0.0 +## DisTorch SafeTensor - Block Swap Memory Management -## Executive Summary +### ⚠️ BOOTSTRAP DOCUMENT - START HERE AFTER CONTEXT RESET ⚠️ -ComfyUI-MultiGPU provides intelligent model distribution across multiple GPUs and system RAM, optimizing for minimal VRAM usage while maintaining performance. Version 2.0 introduces a unified interface supporting both layer-by-layer transfers (DisTorch) and block swapping strategies. +**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. -## Core Concepts +**STATUS**: Implemented. The logic has been integrated into `__init__.py`. -### 1. Virtual VRAM -Virtual VRAM represents the extended memory pool available by offloading model components to other devices (CPU RAM or secondary GPUs). The system manages transfers between devices transparently during inference. +--- -### 2. Transfer Strategies +## STEP 1: UNDERSTAND THE EXISTING APPROACHES -#### DisTorch (Layer-by-Layer) -- **Mechanism**: Spoofs quantized tensors on offload device, dequantizes JIT to compute device -- **Transfer Size**: Single layer at a time (~100-200MB) -- **VRAM Usage**: Minimal (1 layer active) -- **Best For**: Video generation (long inference times) -- **Trade-off**: Many small PCIe transfers +### 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. + +### 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. + +### 1C. ComfyUI-MultiGPU Integration +**File**: `ComfyUI-MultiGPU/__init__.py` +**IMPLEMENTATION LOCATION**: All DisTorch logic is implemented within this single file to ensure portability and avoid external dependencies. + +--- + +## 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. + +**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). + +--- + +## STEP 3: DisTorch SafeTensor IMPLEMENTATION SPEC + +### The Wrapper Function +The core of the implementation is the `override_class_with_distorch_safetensor` function, which wraps existing ComfyUI model loaders. ```python -# DisTorch approach - minimal VRAM footprint -def forward_hook(module, input, output): - # Load single layer - load_layer_to_device(module, compute_device) - output = module.forward(input) - # Immediately offload - offload_layer(module, offload_device) - return output +def override_class_with_distorch_safetensor(cls): + """DisTorch wrapper for SafeTensor models, providing block-swap memory optimization.""" ``` -#### Block Swap (New in V2) -- **Mechanism**: Moves blocks of layers between devices -- **Transfer Size**: Configurable (1-8GB blocks) -- **VRAM Usage**: Reserved swap buffer -- **Best For**: Image generation (short inference times) -- **Trade-off**: Fewer, larger PCIe transfers +### 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. + +### 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 + +**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. + +--- + +## STEP 4: CURRENT IMPLEMENTATION CODE IN __init__.py + +The following code is a representation of the current implementation within `__init__.py`. ```python -# Block swap approach - batched transfers -def forward_hook(module, input, output): - if need_swap(module): - # Swap entire block - offload_block(current_block, offload_device) - load_block(next_block, compute_device) - return module.forward(input) -``` - -### 3. Unified Interface - -All strategies share common parameters: -```python -class VirtualVRAMConfig: - virtual_vram_gb: float # Total model size to offload - swap_space_gb: float # Reserved buffer on compute device - swap_device: str # Where to offload ("cpu", "cuda:1") +def override_class_with_distorch_safetensor(cls): + """DisTorch wrapper for SafeTensor models, providing block-swap memory optimization.""" - # Derived behavior - if swap_space_gb < min_layer_size: - use_distorch() # Layer-by-layer - else: - use_block_swap() # Block transfers -``` + 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 -## Implementation Architecture + CATEGORY = "multigpu" + FUNCTION = "override" -### Memory Management + 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}") -#### Size Calculation (Shared Utility) -```python -def calculate_model_size(model): - """Calculate actual memory footprint""" - total_bytes = 0 - for param in model.parameters(): - if hasattr(param, 'quant_type'): # GGUF - # Account for quantization - total_bytes += calculate_gguf_size(param) - else: # Safetensor - total_bytes += param.element_size() * param.nelement() - return total_bytes / (1024**3) # GB -``` - -#### Block Partitioning -```python -def partition_model(model, swap_space_gb): - """Divide model into swappable blocks""" - blocks = [] - current_block = [] - current_size = 0 - - for name, module in model.named_modules(): - module_size = get_module_size(module) - - if current_size + module_size > swap_space_gb: - # Start new block - blocks.append(current_block) - current_block = [module] - current_size = module_size - else: - current_block.append(module) - current_size += module_size - - return blocks -``` - -### Hook System - -#### Pre/Post Forward Hooks -```python -class ModelWrapper: - def __init__(self, model, config): - self.model = model - self.config = config - self.blocks = partition_model(model, config.swap_space_gb) - self.current_block_idx = -1 - - # Install hooks - for block_idx, block in enumerate(self.blocks): - for module in block: - module.register_forward_pre_hook( - lambda m, i: self.pre_forward(m, block_idx) + 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 ) - - def pre_forward(self, module, block_idx): - if block_idx != self.current_block_idx: - # Swap blocks - self.swap_blocks(self.current_block_idx, block_idx) - self.current_block_idx = block_idx + else: + logging.warning("[DisTorch SafeTensor] Loaded object does not have a 'model' attribute, skipping block swap.") + + return out + + return NodeOverrideDisTorchSafeTensor + + +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) ``` -### GGUF Handling +--- + +## STEP 5: REGISTRATION IN __init__.py + +The new DisTorch SafeTensor wrappers are registered for all relevant core ComfyUI nodes. -#### Quantized Tensor Management ```python -class GGUFHandler: - def handle_gguf_layer(self, layer): - if self.config.swap_space_gb < layer.size: - # Use DisTorch approach - dequantize JIT - return self.distorch_dequantize(layer) - else: - # Can move entire quantized block - return self.block_swap_quantized(layer) - - def distorch_dequantize(self, layer): - """Dequantize during transfer (COPY operation)""" - # Creates new tensor on compute device - return dequantize_to_device(layer, self.compute_device) - - def block_swap_quantized(self, layer): - """Move quantized tensor (SWAP operation)""" - # Moves existing tensor between devices - return layer.to(self.compute_device) +# Register the new DisTorch SafeTensor wrappers +NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"]) +NODE_CLASS_MAPPINGS["UNETLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"]) +# ... and so on for VAELoader, CLIPLoader, ControlNetLoader, etc. ``` -## Phase Implementation Plan +--- -### Phase 1: Block Swap for Safetensors (Current Focus) +## STEP 6: KEY DIFFERENCES -**Goal**: Implement configurable block swapping for non-quantized models. +### DisTorch (GGUF) +- **Granularity**: Per-layer. +- **Use Case**: Maximum memory saving on GGUF models, often with CPU offload. +- **Implementation**: Complex allocation strings and quantization handling. -**Implementation**: -```python -class DisTorchBlockSwap: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "virtual_vram_gb": ("FLOAT", { - "default": 4.0, - "min": 0.1, - "max": 64.0, - "step": 0.1 - }), - "swap_space_gb": ("FLOAT", { - "default": 1.0, - "min": 0.1, - "max": 16.0, - "step": 0.1 - }), - "swap_device": (["cpu", "cuda:0", "cuda:1"],), - } - } - - def apply(self, model, virtual_vram_gb, swap_space_gb, swap_device): - # Calculate model size - model_size = calculate_model_size(model) - - # Partition into blocks - blocks = partition_model(model, swap_space_gb) - - # Install swap hooks - wrapper = BlockSwapWrapper(model, blocks, swap_device) - - return (wrapper.model,) -``` +### DisTorch SafeTensor (NEW) +- **Granularity**: Per-block. +- **Use Case**: Balancing memory and speed for any SafeTensor model. +- **Implementation**: Simple forward hooks, model-agnostic. -### Phase 2: Unified GGUF Support +### WanVideo Block Swap +- **Granularity**: Per-block (model-specific). +- **Use Case**: Optimized specifically for WanVideo models. +- **Implementation**: Integrated directly into the custom model's forward pass. -**Goal**: Extend block swap to GGUF models, auto-selecting strategy. +--- -**Decision Logic**: -```python -def select_strategy(model, config): - if is_gguf(model): - min_layer = get_min_layer_size(model) - if config.swap_space_gb < min_layer: - return DisTorchStrategy() # Must dequantize JIT - else: - return BlockSwapStrategy() # Can move quantized blocks - else: - return BlockSwapStrategy() # Safetensors always use blocks -``` +## WHY THIS MATTERS -### Phase 3: Auto-Optimization +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. -**Goal**: Use empirical data to auto-configure optimal settings. +--- -See `DOE_OPTIMIZATION.md` for detailed benchmarking plan. +## STEP 7: IMPLEMENTATION CHECKLIST -**Auto Mode**: -```python -def auto_configure(model, workload): - # Detect hardware - pcie_gen = detect_pcie_generation() - gpu_bandwidth = detect_gpu_bandwidth() - - # Analyze workload - is_video = workload.frames > 1 - latent_size = workload.height * workload.width - - # Lookup optimal config from DOE results - if is_video: - return {"swap_space_gb": 0.1} # Minimize transfers - else: - return lookup_optimal_config( - model.size, latent_size, pcie_gen - ) -``` +- [X] Delete `blockswap.py` (DONE) +- [X] Document the approach (THIS DOCUMENT) +- [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) +- [ ] Test with SDXL checkpoint +- [ ] Test with Flux checkpoint +- [ ] Verify memory usage matches expectations +- [ ] Measure transfer overhead -## Performance Characteristics +--- -### Transfer Overhead Analysis +## FUTURE EXTENSIONS -| Strategy | Transfer Size | Frequency | PCIe Time | Best Case | -|----------|--------------|-----------|-----------|-----------| -| DisTorch | 100-200MB | Every layer | High | Video (long inference) | -| Block Swap (1GB) | 1GB | Every ~10 layers | Medium | Balanced | -| Block Swap (4GB) | 4GB | Every ~40 layers | Low | Image (short inference) | - -### Memory Usage Patterns - -``` -DisTorch (0.1GB swap): -|===| <- Active layer (100MB) -|...|...|...|...|...| <- Offloaded layers - -Block Swap (2GB swap): -|==========| <- Active block (2GB) -|..........|..........| <- Offloaded blocks -``` - -## Advantages Over Existing Solutions - -### vs. Sequential CPU Offload -- **Granular Control**: Configure exact offload amount -- **Multi-GPU Support**: Use secondary GPUs as fast swap -- **Quantization Aware**: Handles GGUF efficiently - -### vs. Model Parallelism -- **No Model Modification**: Works with any model -- **Dynamic**: Adjusts to available resources -- **Flexible**: User controls memory/speed trade-off - -## Code Organization - -``` -ComfyUI-MultiGPU/ -├── nodes.py # Node definitions -├── core/ -│ ├── distorch.py # Original layer-by-layer -│ ├── blockswap.py # New block swapping -│ ├── memory.py # Shared memory utilities -│ └── hooks.py # Hook management -├── strategies/ -│ ├── auto.py # Auto-optimization -│ ├── gguf.py # GGUF-specific handling -│ └── safetensor.py # Safetensor handling -└── benchmark/ - ├── doe.py # DOE test runner - └── profiles.py # Hardware profiles -``` - -## Testing Strategy - -### Unit Tests -- Memory calculation accuracy -- Block partitioning logic -- Hook installation/removal - -### Integration Tests -- Safetensor models (SDXL, Flux) -- GGUF models (quantized) -- Multi-GPU configurations - -### Performance Tests -- Measure transfer overhead -- Verify memory usage -- Benchmark vs baseline - -## Migration Path - -### For Existing Users -1. Current DisTorch nodes continue working -2. New unified node available alongside -3. Gradual migration as benefits proven - -### Configuration Migration -```python -# Old DisTorch -distorch_model = DisTorch(model, device_map) - -# New Unified (equivalent) -unified_model = VirtualVRAM( - model, - virtual_vram_gb=model_size, - swap_space_gb=0.1, # DisTorch-like - swap_device="cpu" -) -``` - -## Future Directions - -### Adaptive Strategies -- Monitor transfer patterns -- Adjust block size dynamically -- Predict optimal points - -### Pipeline Integration -- Coordinate with samplers -- Batch-aware swapping -- Multi-model orchestration - -### Hardware Acceleration -- Direct Storage API -- NVLink optimization -- CXL memory pooling - -## Conclusion - -Version 2.0 unifies memory management strategies under a coherent interface, providing users with fine-grained control over the memory/performance trade-off while maintaining backward compatibility and preparing for future optimizations. +1. **Auto-mode**: Automatically determine optimal settings based on available VRAM and model size. +2. **Dynamic Block Sizing**: Group layers into blocks dynamically instead of relying on the model's predefined block structure. +3. **Advanced Profiling**: Add tools to measure transfer overhead and help users optimize their settings. diff --git a/core/blockswap.py b/core/blockswap.py deleted file mode 100644 index 2146895..0000000 --- a/core/blockswap.py +++ /dev/null @@ -1,495 +0,0 @@ -""" -Block Swap implementation for ComfyUI-MultiGPU -Based on analysis of WanVideo's block swap mechanism -""" - -import torch -import logging -from typing import Dict, List, Tuple, Optional, Any -from dataclasses import dataclass -import gc -import os -from datetime import datetime -import traceback -import json - -# Set up file logging for DisTorch -log_dir = os.path.join(os.path.dirname(os.path.dirname(__file__)), "logs") -os.makedirs(log_dir, exist_ok=True) -log_file = os.path.join(log_dir, f"distorch_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log") - -# Configure file handler -file_handler = logging.FileHandler(log_file, mode='w') -file_handler.setLevel(logging.DEBUG) -file_formatter = logging.Formatter( - '%(asctime)s.%(msecs)03d - [%(name)s] - %(levelname)s - %(funcName)s:%(lineno)d - %(message)s', - datefmt='%Y-%m-%d %H:%M:%S' -) -file_handler.setFormatter(file_formatter) - -# Create DisTorch logger -distorch_logger = logging.getLogger("DisTorch") -distorch_logger.setLevel(logging.DEBUG) -distorch_logger.addHandler(file_handler) - -# Also add console handler for important messages -console_handler = logging.StreamHandler() -console_handler.setLevel(logging.INFO) -console_formatter = logging.Formatter('[DisTorch] %(levelname)s: %(message)s') -console_handler.setFormatter(console_formatter) -distorch_logger.addHandler(console_handler) - -distorch_logger.info(f"DisTorch logging initialized. Log file: {log_file}") -distorch_logger.debug("="*80) -distorch_logger.debug("DISTORCH BLOCK SWAP MODULE LOADED") -distorch_logger.debug(f"PyTorch version: {torch.__version__}") -distorch_logger.debug(f"CUDA available: {torch.cuda.is_available()}") -if torch.cuda.is_available(): - distorch_logger.debug(f"CUDA device count: {torch.cuda.device_count()}") - for i in range(torch.cuda.device_count()): - distorch_logger.debug(f" Device {i}: {torch.cuda.get_device_name(i)}") - distorch_logger.debug(f" Memory: {torch.cuda.get_device_properties(i).total_memory / (1024**3):.2f} GB") -distorch_logger.debug("="*80) - - -@dataclass -class BlockSwapConfig: - """Configuration for block swapping""" - virtual_vram_gb: float # Total model size to offload - swap_space_gb: float # Reserved buffer on compute device - swap_device: str # Where to offload ("cpu", "cuda:1", etc) - compute_device: str # Where to run computation (usually "cuda:0") - use_non_blocking: bool = False # Non-blocking transfers - - def __post_init__(self): - distorch_logger.debug(f"BlockSwapConfig.__post_init__ called") - distorch_logger.debug(f" Raw swap_device: {self.swap_device}") - distorch_logger.debug(f" Raw compute_device: {self.compute_device}") - distorch_logger.debug(f" virtual_vram_gb: {self.virtual_vram_gb}") - distorch_logger.debug(f" swap_space_gb: {self.swap_space_gb}") - distorch_logger.debug(f" use_non_blocking: {self.use_non_blocking}") - - try: - self.swap_device = torch.device(self.swap_device) - self.compute_device = torch.device(self.compute_device) - distorch_logger.debug(f" Converted swap_device: {self.swap_device}") - distorch_logger.debug(f" Converted compute_device: {self.compute_device}") - except Exception as e: - distorch_logger.error(f"Error converting devices: {e}") - distorch_logger.error(traceback.format_exc()) - raise - - -class BlockSwapManager: - """Manages block swapping for transformer models""" - - def __init__(self, model: torch.nn.Module, config: BlockSwapConfig): - distorch_logger.debug("BlockSwapManager.__init__ called") - distorch_logger.debug(f" Model type: {type(model).__name__}") - distorch_logger.debug(f" Model device: {next(model.parameters()).device if any(model.parameters()) else 'no params'}") - - self.model = model - self.config = config - self.blocks = [] - self.current_block_idx = -1 - self.hooks = [] - self.swap_count = 0 # Track number of swaps for debugging - - # Calculate model size - distorch_logger.debug("Calculating model size...") - self.model_size_gb = self._calculate_model_size() - distorch_logger.info(f"Model size: {self.model_size_gb:.2f} GB") - - # Log model structure - self._log_model_structure() - - # Partition model into blocks - distorch_logger.debug("Partitioning model into blocks...") - self._partition_model() - - # Install hooks - distorch_logger.debug("Installing forward hooks...") - self._install_hooks() - - distorch_logger.debug("BlockSwapManager initialization complete") - - def _log_model_structure(self): - """Log the model structure for debugging""" - distorch_logger.debug("Model structure analysis:") - module_count = 0 - param_count = 0 - module_types = {} - - for name, module in self.model.named_modules(): - module_count += 1 - module_type = type(module).__name__ - module_types[module_type] = module_types.get(module_type, 0) + 1 - - # Count parameters in this module - module_params = sum(p.numel() for p in module.parameters(recurse=False)) - if module_params > 0: - param_count += module_params - if module_count <= 10: # Log first 10 modules with params - distorch_logger.debug(f" {name}: {module_type} ({module_params:,} params)") - - distorch_logger.debug(f"Total modules: {module_count}") - distorch_logger.debug(f"Total parameters: {param_count:,}") - distorch_logger.debug("Module type distribution:") - for module_type, count in sorted(module_types.items(), key=lambda x: x[1], reverse=True)[:10]: - distorch_logger.debug(f" {module_type}: {count}") - - def _calculate_model_size(self) -> float: - """Calculate total model size in GB""" - total_bytes = 0 - param_count = 0 - - for name, param in self.model.named_parameters(): - if param.data is not None: - param_bytes = param.element_size() * param.nelement() - total_bytes += param_bytes - param_count += 1 - - if param_count <= 5: # Log first 5 parameters - distorch_logger.debug(f" Param {name}: shape={param.shape}, bytes={param_bytes:,}") - - distorch_logger.debug(f"Total parameters: {param_count}") - distorch_logger.debug(f"Total bytes: {total_bytes:,}") - - return total_bytes / (1024**3) - - def _get_module_size(self, module: torch.nn.Module) -> float: - """Calculate size of a module in GB""" - total_bytes = 0 - for param in module.parameters(recurse=False): - if param.data is not None: - total_bytes += param.element_size() * param.nelement() - return total_bytes / (1024**3) - - def _partition_model(self): - """Partition model into swappable blocks based on swap_space_gb""" - distorch_logger.debug("Starting model partitioning...") - - # Find transformer blocks (common patterns) - transformer = None - transformer_blocks = [] - - # Try to find transformer module - for name, module in self.model.named_modules(): - # Common transformer patterns - patterns = ['transformer', 'diffusion_model', 'unet', 'encoder', 'decoder'] - if any(pattern in name.lower() for pattern in patterns): - distorch_logger.debug(f"Found potential transformer module: {name} ({type(module).__name__})") - - # Check if it has sequential blocks - if hasattr(module, 'blocks'): - transformer = module - transformer_blocks = list(module.blocks) - distorch_logger.debug(f" Found {len(transformer_blocks)} blocks in {name}.blocks") - break - elif hasattr(module, 'layers'): - transformer = module - transformer_blocks = list(module.layers) - distorch_logger.debug(f" Found {len(transformer_blocks)} layers in {name}.layers") - break - elif hasattr(module, 'transformer_blocks'): - transformer = module - transformer_blocks = list(module.transformer_blocks) - distorch_logger.debug(f" Found {len(transformer_blocks)} transformer_blocks in {name}") - break - - if not transformer_blocks: - # Fallback: partition all modules - distorch_logger.warning("No transformer blocks found, using fallback partitioning") - self._partition_fallback() - return - - # Group blocks based on swap_space_gb - current_block = [] - current_size = 0 - swap_space_bytes = self.config.swap_space_gb * (1024**3) - - distorch_logger.debug(f"Grouping {len(transformer_blocks)} blocks with swap_space={self.config.swap_space_gb} GB") - - for idx, block in enumerate(transformer_blocks): - block_size = self._get_module_size(block) * (1024**3) # Convert to bytes - distorch_logger.debug(f" Block {idx}: size={block_size/(1024**3):.3f} GB") - - if current_size + block_size > swap_space_bytes and current_block: - # Start new block group - self.blocks.append(current_block) - distorch_logger.debug(f" Created block group {len(self.blocks)-1} with {len(current_block)} blocks, total size={current_size/(1024**3):.3f} GB") - current_block = [block] - current_size = block_size - else: - current_block.append(block) - current_size += block_size - - # Add remaining blocks - if current_block: - self.blocks.append(current_block) - distorch_logger.debug(f" Created final block group {len(self.blocks)-1} with {len(current_block)} blocks, total size={current_size/(1024**3):.3f} GB") - - distorch_logger.info(f"Partitioned into {len(self.blocks)} block groups") - for i, group in enumerate(self.blocks): - group_size = sum(self._get_module_size(b) for b in group) - distorch_logger.info(f" Block group {i}: {len(group)} blocks, {group_size:.3f} GB") - - def _partition_fallback(self): - """Fallback partitioning when transformer structure is not recognized""" - distorch_logger.debug("Using fallback partitioning strategy") - all_modules = [] - - # Collect all modules with parameters - for name, module in self.model.named_modules(): - param_count = sum(p.numel() for p in module.parameters(recurse=False)) - if param_count > 0: - module_size = self._get_module_size(module) - all_modules.append((name, module, module_size)) - if len(all_modules) <= 10: - distorch_logger.debug(f" Module {name}: {param_count:,} params, {module_size:.3f} GB") - - distorch_logger.debug(f"Found {len(all_modules)} modules with parameters") - - # Group by size - current_block = [] - current_size = 0 - swap_space_bytes = self.config.swap_space_gb * (1024**3) - - for name, module, module_size_gb in all_modules: - module_size = module_size_gb * (1024**3) - - if current_size + module_size > swap_space_bytes and current_block: - self.blocks.append([m for _, m, _ in current_block]) - distorch_logger.debug(f" Created block group with {len(current_block)} modules, size={current_size/(1024**3):.3f} GB") - current_block = [(name, module, module_size_gb)] - current_size = module_size - else: - current_block.append((name, module, module_size_gb)) - current_size += module_size - - if current_block: - self.blocks.append([m for _, m, _ in current_block]) - distorch_logger.debug(f" Created final block group with {len(current_block)} modules, size={current_size/(1024**3):.3f} GB") - - def _install_hooks(self): - """Install forward pre-hooks on blocks""" - distorch_logger.debug(f"Installing hooks on {len(self.blocks)} block groups") - - for block_idx, block_group in enumerate(self.blocks): - distorch_logger.debug(f" Installing hooks for block group {block_idx} ({len(block_group)} modules)") - for module in block_group: - hook = module.register_forward_pre_hook( - lambda m, i, bidx=block_idx: self._pre_forward_hook(m, i, bidx) - ) - self.hooks.append(hook) - - distorch_logger.debug(f"Installed {len(self.hooks)} hooks total") - - def _pre_forward_hook(self, module: torch.nn.Module, inputs: Tuple, block_idx: int): - """Hook called before forward pass of each block""" - if block_idx != self.current_block_idx: - distorch_logger.debug(f"Pre-forward hook triggered: current={self.current_block_idx}, needed={block_idx}") - self._swap_blocks(self.current_block_idx, block_idx) - self.current_block_idx = block_idx - return inputs - - def _swap_blocks(self, old_idx: int, new_idx: int): - """Swap blocks between devices""" - self.swap_count += 1 - distorch_logger.debug(f"[Swap #{self.swap_count}] Swapping from block {old_idx} to {new_idx}") - distorch_logger.debug(f" Swap device: {self.config.swap_device}") - distorch_logger.debug(f" Compute device: {self.config.compute_device}") - - start_time = datetime.now() - - # Offload old block - if old_idx >= 0 and old_idx < len(self.blocks): - distorch_logger.debug(f" Offloading block {old_idx} to {self.config.swap_device}") - for i, module in enumerate(self.blocks[old_idx]): - self._move_module(module, self.config.swap_device) - if i == 0: # Log first module movement - distorch_logger.debug(f" Moved module {type(module).__name__} to {self.config.swap_device}") - - # Load new block - if new_idx >= 0 and new_idx < len(self.blocks): - distorch_logger.debug(f" Loading block {new_idx} to {self.config.compute_device}") - for i, module in enumerate(self.blocks[new_idx]): - self._move_module(module, self.config.compute_device) - if i == 0: # Log first module movement - distorch_logger.debug(f" Moved module {type(module).__name__} to {self.config.compute_device}") - - # Clear cache if needed - if self.config.compute_device.type == 'cuda': - before_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3) - torch.cuda.empty_cache() - after_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3) - distorch_logger.debug(f" GPU cache cleared: {before_free:.2f} GB -> {after_free:.2f} GB") - - elapsed = (datetime.now() - start_time).total_seconds() - distorch_logger.debug(f" Swap completed in {elapsed:.3f} seconds") - - def _move_module(self, module: torch.nn.Module, device: torch.device): - """Move a module to specified device""" - try: - module.to(device, non_blocking=self.config.use_non_blocking) - except Exception as e: - distorch_logger.error(f"Error moving module {type(module).__name__} to {device}: {e}") - distorch_logger.error(traceback.format_exc()) - raise - - def prepare(self): - """Prepare model for inference by moving all blocks to swap device""" - distorch_logger.info(f"Preparing model: moving all blocks to {self.config.swap_device}") - - start_time = datetime.now() - - for i, block_group in enumerate(self.blocks): - distorch_logger.debug(f" Moving block group {i} ({len(block_group)} modules)") - for module in block_group: - self._move_module(module, self.config.swap_device) - - # Reset current block - self.current_block_idx = -1 - - # Clear GPU cache - if self.config.compute_device.type == 'cuda': - before_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3) - torch.cuda.empty_cache() - after_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3) - distorch_logger.debug(f"GPU memory after prepare: {before_free:.2f} GB -> {after_free:.2f} GB") - - gc.collect() - - elapsed = (datetime.now() - start_time).total_seconds() - distorch_logger.info(f"Model preparation completed in {elapsed:.3f} seconds") - - def cleanup(self): - """Remove hooks and cleanup""" - distorch_logger.info("Cleaning up BlockSwapManager") - distorch_logger.debug(f" Total swaps performed: {self.swap_count}") - - for hook in self.hooks: - hook.remove() - self.hooks.clear() - - distorch_logger.info("BlockSwap cleanup complete") - - -class DisTorch: - """ComfyUI node for block swap configuration""" - - @classmethod - def INPUT_TYPES(cls): - from .. import get_device_list - devices = get_device_list() - - distorch_logger.debug(f"DisTorch.INPUT_TYPES called, available devices: {devices}") - - return { - "required": { - "model": ("MODEL",), - "virtual_vram_gb": ("FLOAT", { - "default": 4.0, - "min": 0.1, - "max": 64.0, - "step": 0.1, - "tooltip": "Amount of model to offload to swap device" - }), - "swap_space_gb": ("FLOAT", { - "default": 1.0, - "min": 0.1, - "max": 16.0, - "step": 0.1, - "tooltip": "Size of buffer on compute device for active blocks" - }), - "swap_device": (devices, { - "default": "cpu", - "tooltip": "Device to offload inactive blocks to" - }), - "compute_device": (devices, { - "default": devices[1] if len(devices) > 1 else devices[0], - "tooltip": "Device to run computation on" - }), - }, - "optional": { - "use_non_blocking": ("BOOLEAN", { - "default": False, - "tooltip": "Use non-blocking memory transfers (faster but uses more RAM)" - }), - } - } - - RETURN_TYPES = ("MODEL",) - FUNCTION = "apply_block_swap" - CATEGORY = "multigpu" - - def apply_block_swap(self, model, virtual_vram_gb: float, swap_space_gb: float, - swap_device: str, compute_device: str, use_non_blocking: bool = False): - """Apply block swap configuration to model""" - - distorch_logger.info("="*60) - distorch_logger.info("DisTorch.apply_block_swap called") - distorch_logger.info("="*60) - - distorch_logger.info(f"Configuration:") - distorch_logger.info(f" Virtual VRAM: {virtual_vram_gb} GB") - distorch_logger.info(f" Swap space: {swap_space_gb} GB") - distorch_logger.info(f" Swap device: {swap_device}") - distorch_logger.info(f" Compute device: {compute_device}") - distorch_logger.info(f" Non-blocking: {use_non_blocking}") - - distorch_logger.debug(f"Input model type: {type(model)}") - distorch_logger.debug(f"Model attributes: {dir(model)[:10]}...") # Log first 10 attributes - - try: - # Create config - config = BlockSwapConfig( - virtual_vram_gb=virtual_vram_gb, - swap_space_gb=swap_space_gb, - swap_device=swap_device, - compute_device=compute_device, - use_non_blocking=use_non_blocking - ) - - # Get the actual model (handle ModelPatcher) - if hasattr(model, 'model'): - actual_model = model.model - distorch_logger.debug(f"Extracted actual model from ModelPatcher: {type(actual_model)}") - else: - actual_model = model - distorch_logger.debug(f"Using model directly: {type(actual_model)}") - - # Check if model has diffusion_model (common pattern) - if hasattr(actual_model, 'diffusion_model'): - target_model = actual_model.diffusion_model - distorch_logger.debug(f"Found diffusion_model: {type(target_model)}") - else: - target_model = actual_model - distorch_logger.debug(f"No diffusion_model found, using model as-is") - - # Create block swap manager - distorch_logger.debug("Creating BlockSwapManager...") - manager = BlockSwapManager(target_model, config) - - # Prepare model (move blocks to swap device) - distorch_logger.debug("Preparing model...") - manager.prepare() - - # Store manager on model for later access - model._block_swap_manager = manager - distorch_logger.debug("Stored BlockSwapManager on model._block_swap_manager") - - # Also set the load_device attribute if it exists - if hasattr(model, 'load_device'): - model.load_device = config.compute_device - distorch_logger.debug(f"Set model.load_device to {config.compute_device}") - - distorch_logger.info("Block swap configuration applied successfully") - distorch_logger.info("="*60) - - except Exception as e: - distorch_logger.error(f"Error in apply_block_swap: {e}") - distorch_logger.error(traceback.format_exc()) - raise - - return (model,)