docs: Update architecture document for DisTorch SafeTensor
This commit is contained in:
+183
-314
@@ -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.
|
||||
|
||||
@@ -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,)
|
||||
Reference in New Issue
Block a user