docs: Update architecture document for DisTorch SafeTensor

This commit is contained in:
John Pollock
2025-08-09 10:46:05 -05:00
parent aee0987779
commit 992b1c6587
2 changed files with 183 additions and 809 deletions
+183 -314
View File
@@ -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.
-495
View File
@@ -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,)