Merge pull request #31 from AInVFX/blockswap

BlockSwap support, thanks To @adrientoupet
This commit is contained in:
NumZ
2025-07-09 22:10:00 +02:00
committed by GitHub
9 changed files with 1336 additions and 100 deletions
+26
View File
@@ -137,6 +137,32 @@ Of course, the output resolution also has an impact, so if your hardware doesn't
- `new_resolution`: New desired short edge in px, will keep ratio on other edge
- `batch_size`: VERY IMPORTANT!, this model consume a lot of VRAM, All your VRAM, even for the 3B model, so for GPU under 24GB VRAM keep this value Low, good value is "1" without temporal consistency, "5" for temporal consistency, but higher is this value better is the result.
- `preserve_vram`: for VRAM < 24GB, If true, It will unload unused models during process, longer but works, otherwise probably OOM with
4. 🧩 **BlockSwap Configuration (Optional - For Limited VRAM)**
<img src="docs/BlockSwap.png" width="100%">
BlockSwap enables running large models on GPUs with limited VRAM by dynamically swapping transformer blocks between GPU and CPU memory during inference.
**To enable BlockSwap:**
- Add the **SEEDVR2 BlockSwap Config** node to your workflow
- Connect its output to the `block_swap_config` input of the SeedVR2 Video Upscaler node
**BlockSwap parameters:**
- `blocks_to_swap` (0-32 for 3B model, 0-36 for 7B model): Number of transformer blocks to offload
- 0 = Disabled (fastest, highest VRAM)
- 16 = Balanced (moderate speed/VRAM trade-off)
- 32-36 = Maximum savings (slowest, lowest VRAM)
- `offload_io_components`: Move embeddings/IO layers to CPU (additional VRAM savings, slower)
- `use_non_blocking`: Asynchronous GPU transfers (keep True for better performance)
- `cache_model`: Keep model in RAM between runs (faster for batch processing)
- `enable_debug`: Show detailed memory usage and timing
**Finding optimal settings:**
- Start with `blocks_to_swap=16`, increase if you get OOM errors, decrease if you have spare VRAM
- Enable debug mode to monitor memory usage
- The first 1-2 blocks might show longer swap times - this is normal
- Combine with `preserve_vram=True` for maximum memory savings
## 🖥️ Run as Standalone
Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

+36 -6
View File
@@ -17,6 +17,7 @@ Key Features:
"""
import os
import gc
import torch
import time
from torchvision.transforms import Compose, Lambda, Normalize
@@ -120,14 +121,19 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
conditions = [condition]
t = time.time()
# Check if BlockSwap is active
use_blockswap = hasattr(runner, "_blockswap_active") and runner._blockswap_active
# Use adaptive autocast for optimal performance
with torch.no_grad():
with torch.autocast("cuda", autocast_dtype, enabled=True):
video_tensors = runner.inference(
noises=noises,
conditions=conditions,
preserve_vram=preserve_vram, # Memory offload optimization
preserve_vram=preserve_vram # Memory offload optimization
and not use_blockswap, # Disable dit_offload if BlockSwap active
temporal_overlap=temporal_overlap,
use_blockswap=use_blockswap,
**text_embeds_dict,
)
@@ -141,7 +147,7 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
cond_latents = cond_latents[0].to("cpu")
conditions = conditions[0].to("cpu")
condition = condition.to("cpu")
return samples #, last_latents
@@ -175,7 +181,7 @@ def cut_videos(videos):
return result
def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, preserve_vram=False, temporal_overlap=0, debug=False, progress_callback=None):
def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, preserve_vram=False, temporal_overlap=0, debug=False, block_swap_config=None, progress_callback=None):
"""
Main generation loop with context-aware temporal processing
@@ -188,6 +194,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
batch_size (int): Batch size for processing
preserve_vram (str/bool): VRAM preservation mode
temporal_overlap (int): Frames for temporal continuity
debug (bool): Debug mode
block_swap_config (dict): Optional BlockSwap configuration
progress_callback (callable): Optional callback for progress reporting
Returns:
@@ -202,7 +210,13 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
- Real-time progress reporting
"""
device = "cuda" if torch.cuda.is_available() else "cpu"
# Log BlockSwap status
if block_swap_config:
blocks_to_swap = block_swap_config.get("blocks_to_swap", 0)
if blocks_to_swap > 0:
print(f"🔄 Generation starting with BlockSwap: {blocks_to_swap} blocks")
# Adaptive model dtype detection for maximum performance
model_dtype = None
try:
@@ -296,6 +310,12 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
#images = images.to("cpu")
#print(f"🔄 Images to CPU time: {time.time() - t} seconds")
# Use existing debugger from runner if available
debugger = None
if hasattr(runner, '_blockswap_debugger') and runner._blockswap_debugger is not None:
debugger = runner._blockswap_debugger
debugger.clear_history()
try:
# Main processing loop with context awareness
for batch_count, batch_idx in enumerate(range(0, len(images), step)):
@@ -414,12 +434,19 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
if debug:
print(f"🔄 Time batch: {time.time() - tps_loop} seconds")
if preserve_vram:
# Clean VRAM after each batch when preserve_vram is active (but not with blockswap)
if preserve_vram and not (block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0):
torch.cuda.empty_cache()
#del transformed_video
#clear_vram_cache()
# Log memory state at the end of each batch
if debugger:
debugger.log_memory_state(f"Batch {batch_number} - Memory", show_tensors=True)
finally:
if debugger:
debugger.log("🧹 Generation loop cleanup")
# Final cleanup of embeddings
text_pos_embeds = text_pos_embeds.to("cpu")
text_neg_embeds = text_neg_embeds.to("cpu")
@@ -428,7 +455,10 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
torch.cuda.empty_cache()
#del text_pos_embeds, text_neg_embeds
#clear_vram_cache()
# Log final memory state
if debugger:
debugger.log_memory_state("Generation loop - After cleanup", show_tensors=True)
# OPTIMISATION ULTIME : Pré-allocation et copie directe (évite les torch.cat multiples)
print(f"💾 Processing {len(batch_samples)} batch_samples with memory-optimized pre-allocation")
+10 -4
View File
@@ -290,6 +290,7 @@ class VideoDiffusionInfer():
cfg_scale: Optional[float] = None,
preserve_vram: bool = False,
temporal_overlap: int = 0,
use_blockswap: bool = False,
) -> List[Tensor]:
assert len(noises) == len(conditions) == len(texts_pos) == len(texts_neg)
batch_size = len(noises)
@@ -363,10 +364,15 @@ class VideoDiffusionInfer():
self.vae = self.vae.to("cpu")
if self.debug:
print(f"🔄 VAE to CPU time: {time.time() - t} seconds")
t = time.time()
self.dit = self.dit.to(get_device())
if self.debug:
print(f"🔄 Dit to GPU time: {time.time() - t} seconds")
# Before sampling, check if BlockSwap is active
if not use_blockswap and not hasattr(self, "_blockswap_active"):
t = time.time()
self.dit = self.dit.to(get_device())
if self.debug:
print(f"🔄 Dit to GPU time: {time.time() - t} seconds")
else:
# BlockSwap manages device placement
pass
t = time.time()
+115 -23
View File
@@ -30,22 +30,26 @@ except ImportError:
from src.optimization.memory_manager import get_basic_vram_info, clear_vram_cache
from src.optimization.compatibility import FP8CompatibleDiT
from src.optimization.memory_manager import preinitialize_rope_cache
from src.optimization.memory_manager import preinitialize_rope_cache, clear_rope_lru_caches
from src.common.config import load_config, create_object
from src.core.infer import VideoDiffusionInfer
from src.optimization.blockswap import apply_block_swap_to_dit
# Get script directory for config paths
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False, block_swap_config=None, cached_runner=None):
"""
Configure and create a VideoDiffusionInfer runner for the specified model
Args:
model (str): Model filename (e.g., "seedvr2_ema_3b_fp16.safetensors")
base_cache_dir (str): Base directory containing model files
preserve_vram (bool): Whether to preserve VRAM
debug (bool): Enable debug logging
block_swap_config (dict): Optional BlockSwap configuration
cached_runner: Optional cached runner to reuse entirely (not just DiT)
Returns:
VideoDiffusionInfer: Configured runner instance ready for inference
@@ -56,6 +60,58 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
- VAE configuration with proper parameter handling
- Memory optimization and RoPE cache pre-initialization
"""
# Check if we can fully reuse the cached runner
if cached_runner and block_swap_config and block_swap_config.get("cache_model", False):
# Clear RoPE caches before reuse
if hasattr(cached_runner, 'dit'):
dit_model = cached_runner.dit
if hasattr(dit_model, 'dit_model'):
dit_model = dit_model.dit_model
clear_rope_lru_caches(dit_model)
print(f"♻️ Reusing cached runner for {model}")
# Check if blockswap needs to be applied
blockswap_needed = block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
if blockswap_needed:
# Check if we have cached configuration
has_cached_config = hasattr(cached_runner, "_cached_blockswap_config")
if has_cached_config:
# Compare configurations
cached_config = cached_runner._cached_blockswap_config
config_matches = (
cached_config.get("blocks_to_swap") == block_swap_config.get("blocks_to_swap") and
cached_config.get("offload_io_components") == block_swap_config.get("offload_io_components", False) and
cached_config.get("use_non_blocking") == block_swap_config.get("use_non_blocking", True)
)
if config_matches:
# Configuration matches - fast re-application
print("✅ BlockSwap config matches, performing fast re-application")
# Mark as active before applying
cached_runner._blockswap_active = True
# Apply BlockSwap (will be fast since model structure is intact)
apply_block_swap_to_dit(cached_runner, block_swap_config)
else:
# Configuration changed - apply new config
print("🔄 BlockSwap configuration changed, applying new config")
apply_block_swap_to_dit(cached_runner, block_swap_config)
else:
# No cached config - apply fresh
print("🔄 Applying BlockSwap to cached runner")
apply_block_swap_to_dit(cached_runner, block_swap_config)
return cached_runner
else:
# No BlockSwap needed
return cached_runner
# If we reach here, create a new runner
t = time.time()
vram_info = get_basic_vram_info()
if debug:
@@ -76,7 +132,7 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
print(f"🔄 RUNNER : CONFIG LOAD TIME: {time.time() - t} seconds")
# DiT model configuration is now handled directly in the YAML config files
# No need for dynamic path resolution here anymore!
# Load and configure VAE with additional parameters
vae_config_path = os.path.join(script_directory, 'src/models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml')
t = time.time()
@@ -89,11 +145,11 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
spatial_downsample_factor = vae_config.get('spatial_downsample_factor', 8)
temporal_downsample_factor = vae_config.get('temporal_downsample_factor', 4)
vae_config.spatial_downsample_factor = spatial_downsample_factor
vae_config.temporal_downsample_factor = temporal_downsample_factor
if debug:
print(f"🔄 RUNNER : VAE CONFIG SET TIME: {time.time() - t} seconds")
# Merge additional VAE config with main config (preserving __object__ from main config)
t = time.time()
config.vae.model = OmegaConf.merge(config.vae.model, vae_config)
@@ -104,20 +160,24 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
# Create runner
runner = VideoDiffusionInfer(config, debug)
OmegaConf.set_readonly(runner.config, False)
# Store model name for cache validation
runner._model_name = model
if debug:
print(f"🔄 RUNNER : RUNNER VIDEO DIFFUSION INFER TIME: {time.time() - t} seconds")
# Set device
device = "cuda" if torch.cuda.is_available() else "cpu"
# Configure models
checkpoint_path = os.path.join(base_cache_dir, f'./{model}')
t = time.time()
runner = configure_dit_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
runner = configure_dit_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
if debug:
print(f"🔄 RUNNER : DIT MODEL INFERENCE TIME: {time.time() - t} seconds")
t = time.time()
checkpoint_path = os.path.join(base_cache_dir, f'./{config.vae.checkpoint}')
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug, block_swap_config)
if debug:
print(f"🔄 RUNNER : VAE MODEL INFERENCE TIME: {time.time() - t} seconds")
@@ -126,11 +186,25 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
runner.vae.set_memory_limit(**runner.config.vae.memory_limit)
if debug:
print(f"🔄 RUNNER : VAE MEMORY LIMIT TIME: {time.time() - t} seconds")
# Pre-initialize RoPE cache for optimal performance
t = time.time()
preinitialize_rope_cache(runner)
if debug:
print(f"🔄 RUNNER : ROPE CACHE PREINITIALIZE TIME: {time.time() - t} seconds")
# Check if BlockSwap is active
blockswap_active = (
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
)
# Pre-initialize RoPE cache for optimal performance if BlockSwap is NOT active
if not blockswap_active:
t = time.time()
preinitialize_rope_cache(runner)
if debug:
print(f"🔄 RUNNER : ROPE CACHE PREINITIALIZE TIME: {time.time() - t} seconds")
else:
if debug:
print(f"🔄 RUNNER : Skipping RoPE pre-init due to BlockSwap")
# Apply BlockSwap if configured
if blockswap_active:
apply_block_swap_to_dit(runner, block_swap_config)
#clear_vram_cache()
return runner
@@ -196,7 +270,7 @@ def load_quantized_state_dict(checkpoint_path, device="cpu", keep_native_fp8=Tru
def configure_dit_model_inference(runner, device, checkpoint, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False):
def configure_dit_model_inference(runner, device, checkpoint, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False, block_swap_config=None):
"""
Configure DiT model for inference without distributed decorators
@@ -205,20 +279,30 @@ def configure_dit_model_inference(runner, device, checkpoint, config, preserve_v
device (str): Target device
checkpoint (str): Path to model checkpoint
config: Model configuration
block_swap_config (dict): Optional BlockSwap configuration
Features:
- Automatic format detection and optimal loading
- Native FP8 support with universal compatibility wrapper
- Gradient checkpointing configuration
- Intelligent dtype handling for all model architectures
- BlockSwap support for low VRAM systems
"""
# Create dit model
t = time.time()
loading_device = "cpu" if preserve_vram else device
with torch.device(device):
# Check if BlockSwap is active
blockswap_active = (
block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0
)
loading_device = "cpu" if (preserve_vram or blockswap_active) else device
if blockswap_active and debug:
print(f"🔄 CONFIG DIT : BlockSwap active - creating model on CPU")
with torch.device(loading_device):
runner.dit = create_object(config.dit.model)
# Passer les opérations au modèle
@@ -241,7 +325,7 @@ def configure_dit_model_inference(runner, device, checkpoint, config, preserve_v
if 'state' in locals():
del state
if debug:
print(f"🔄 CONFIG DIT : DiT load time: {time.time() - t} seconds")
#state.to("cpu")
@@ -250,24 +334,31 @@ def configure_dit_model_inference(runner, device, checkpoint, config, preserve_v
# Apply universal compatibility wrapper to ALL models
# This ensures RoPE compatibility and optimal performance across all architectures
t = time.time()
runner.dit = FP8CompatibleDiT(runner.dit)
# Check if already wrapped to avoid double wrapping
if not isinstance(runner.dit, FP8CompatibleDiT):
runner.dit = FP8CompatibleDiT(runner.dit, skip_conversion=False)
if debug:
print(f"🔄 CONFIG DIT : FP8CompatibleDiT time: {time.time() - t} seconds")
# Move DiT to CPU to prevent VRAM leaks (especially for 3B model with complex RoPE)
if preserve_vram:
if preserve_vram and not blockswap_active:
if debug:
print(f"🔄 CONFIG DIT : dit to cpu cause preserve_vram: {preserve_vram}")
runner.dit = runner.dit.to("cpu")
if "7b" in model_weight:
clear_vram_cache()
else:
if state_loading_device == "cpu":
if state_loading_device == "cpu" and not blockswap_active:
runner.dit.to(device)
# Log BlockSwap status if active
if blockswap_active and debug:
print(f"🔄 CONFIG DIT : BlockSwap active ({block_swap_config.get('blocks_to_swap', 0)} blocks) - placement handled by BlockSwap")
return runner
def configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False):
def configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False, block_swap_config=None):
"""
Configure VAE model for inference without distributed decorators
@@ -275,6 +366,7 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config, prese
runner: VideoDiffusionInfer instance
config: Model configuration
device (str): Target device
block_swap_config (dict): Optional BlockSwap configuration
Features:
- Dynamic path resolution for VAE checkpoints
@@ -356,4 +448,4 @@ def configure_vae_model_inference(runner, device, checkpoint_path, config, prese
print(f"🔄 CONFIG VAE : VAE SET CAUSAL SLICING TIME: {time.time() - t} seconds")
return runner
#runner.vae.to("cpu")
#runner.vae.to("cpu")
+251 -47
View File
@@ -2,9 +2,7 @@
# Clean interface for SeedVR2 VideoUpscaler integration with ComfyUI
# Extracted from original seedvr2.py lines 1731-1812
from datetime import datetime
import os
import gc
import time
import torch
from typing import Tuple, Dict, Any
@@ -12,7 +10,14 @@ from typing import Tuple, Dict, Any
from src.utils.downloads import download_weight, get_base_cache_dir
from src.core.model_manager import configure_runner
from src.core.generation import generation_loop
from src.optimization.memory_manager import clear_rope_lru_caches, fast_model_cleanup
from src.optimization.memory_manager import fast_model_cleanup, fast_ram_cleanup
from src.optimization.blockswap import cleanup_blockswap
from src.optimization.memory_manager import (
clear_rope_lru_caches,
fast_model_cleanup,
fast_ram_cleanup,
clear_all_caches
)
# Import ComfyUI progress reporting
from server import PromptServer
@@ -81,6 +86,12 @@ class SeedVR2:
}),
"preserve_vram": ("BOOLEAN", {"default": False}),
},
"optional": {
"block_swap_config": (
"block_swap_config",
{"tooltip": "Optional BlockSwap configuration for low VRAM mode"},
),
},
}
# Define return types for ComfyUI
@@ -90,7 +101,7 @@ class SeedVR2:
CATEGORY = "SEEDVR2"
def execute(self, images: torch.Tensor, model: str, seed: int, new_resolution: int,
batch_size: int, preserve_vram: bool) -> Tuple[torch.Tensor]:
batch_size: int, preserve_vram: bool, block_swap_config=None) -> Tuple[torch.Tensor]:
"""Execute SeedVR2 video upscaling with progress reporting"""
temporal_overlap = 0
@@ -100,50 +111,103 @@ class SeedVR2:
debug = False
cfg_scale = 1.0
try:
return self._internal_execute(images, model, seed, new_resolution, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug)
return self._internal_execute(images, model, seed, new_resolution, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug, block_swap_config)
except Exception as e:
self.cleanup(force_ram_cleanup=True)
raise e
def cleanup(self, force_ram_cleanup: bool = True):
"""Fast cleanup with minimal logging"""
def cleanup(self, force_ram_cleanup: bool = True, keep_model_cached: bool = False, block_swap_config=None):
"""
Comprehensive cleanup with memory tracking
Args:
force_ram_cleanup (bool): Whether to perform aggressive RAM cleanup
keep_model_cached (bool): Whether to keep the model in RAM (only applies with BlockSwap)
block_swap_config: Block swap configuration with enable_debug flag
"""
# Determine if we should keep model cached
should_keep_model = False
if self.runner and keep_model_cached:
is_blockswap_active = (
hasattr(self.runner, "_blockswap_active")
and self.runner._blockswap_active
)
# Check if cache_model is enabled in config
cache_model_enabled = block_swap_config and block_swap_config.get("cache_model", False)
should_keep_model = is_blockswap_active and cache_model_enabled
if self.runner:
# Clear cache
if hasattr(self.runner, 'cache') and hasattr(self.runner.cache, 'cache'):
for key, value in list(self.runner.cache.cache.items()):
if hasattr(value, 'cpu'):
value.cpu()
if hasattr(value, 'detach'):
value.detach()
del value
self.runner.cache.cache.clear()
# Use existing debugger if available
debugger = None
if self.runner and hasattr(self.runner, '_blockswap_debugger'):
debugger = self.runner._blockswap_debugger
debugger.clear_history()
# Perform partial or full cleanup based on model caching
if should_keep_model:
debugger.log("🧹 Partial cleanup - keeping model in RAM")
# Clear DiT model
if hasattr(self.runner, 'dit') and self.runner.dit is not None:
# Clean BlockSwap with state preservation
if hasattr(self.runner, "_blockswap_active") and self.runner._blockswap_active:
cleanup_blockswap(self.runner, keep_state_for_cache=True)
# Clear all caches
if self.runner:
clear_all_caches(self.runner, debugger)
else:
# Full cleanup - existing implementation
if debugger:
debugger.log("🧹 Full cleanup - clearing everything")
if self.runner:
# Clean BlockSwap if active
if hasattr(self.runner, "_blockswap_active") and self.runner._blockswap_active:
cleanup_blockswap(self.runner, keep_state_for_cache=False)
clear_rope_lru_caches(self.runner.dit)
fast_model_cleanup(self.runner.dit)
del self.runner.dit
self.runner.dit = None
# Clear cache
if hasattr(self.runner, 'cache') and hasattr(self.runner.cache, 'cache'):
for key, value in list(self.runner.cache.cache.items()):
if hasattr(value, 'cpu'):
value.cpu()
if hasattr(value, 'detach'):
value.detach()
del value
self.runner.cache.cache.clear()
# Clear DiT model
if hasattr(self.runner, 'dit') and self.runner.dit is not None:
# Handle FP8CompatibleDiT wrapper
if hasattr(self.runner.dit, 'dit_model'):
# Clean inner model first
clear_rope_lru_caches(self.runner.dit.dit_model)
fast_model_cleanup(self.runner.dit.dit_model)
# Break reference from wrapper to model
self.runner.dit.dit_model = None
else:
# Direct model cleanup
clear_rope_lru_caches(self.runner.dit)
fast_model_cleanup(self.runner.dit)
del self.runner.dit
self.runner.dit = None
# Clear VAE model
if hasattr(self.runner, 'vae') and self.runner.vae is not None:
#from src.optimization.memory_manager import fast_model_cleanup
fast_model_cleanup(self.runner.vae)
del self.runner.vae
self.runner.vae = None
# Clear other components
for component in ['sampler', 'sampling_timesteps', 'schedule', 'config']:
if hasattr(self.runner, component):
setattr(self.runner, component, None)
del self.runner
self.runner = None
# Clear VAE model
if hasattr(self.runner, 'vae') and self.runner.vae is not None:
#from src.optimization.memory_manager import fast_model_cleanup
fast_model_cleanup(self.runner.vae)
del self.runner.vae
self.runner.vae = None
# Clear other components
for component in ['sampler', 'sampling_timesteps', 'schedule', 'config']:
if hasattr(self.runner, component):
setattr(self.runner, component, None)
del self.runner
self.runner = None
# Clear embeddings
if self.text_pos_embeds is not None:
if hasattr(self.text_pos_embeds, 'cpu'):
@@ -161,40 +225,71 @@ class SeedVR2:
# Fast RAM cleanup
if force_ram_cleanup:
from src.optimization.memory_manager import fast_ram_cleanup
fast_ram_cleanup()
# BlockSwap debugger memory state
if debugger is not None:
cleanup_stage = "partial" if should_keep_model else "full"
debugger.log_memory_state(f"After {cleanup_stage} cleanup", show_tensors=True)
def _internal_execute(self, images, model, seed, new_resolution, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug):
def _internal_execute(self, images, model, seed, new_resolution, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug, block_swap_config):
"""Internal execution logic with progress tracking"""
total_start_time = time.time()
# Check if we should use model caching
use_cache = (
block_swap_config
and block_swap_config.get("blocks_to_swap", 0) > 0
and block_swap_config.get("cache_model", False)
)
if self.runner is not None:
current_model = getattr(self.runner, '_model_name', None)
model_changed = current_model != model
if model_changed and self.runner is not None:
print(
f"🔄 Model changed from {self.current_model_name} to {model}, clearing cache..."
)
self.cleanup(
force_ram_cleanup=True,
keep_model_cached=False,
block_swap_config=block_swap_config,
)
self.runner = None
# Configure runner
if debug:
print("🔄 Configuring inference runner...")
runner_start = time.time()
self.runner = configure_runner(model, get_base_cache_dir(), preserve_vram, debug)
self.runner = configure_runner(
model, get_base_cache_dir(), preserve_vram, debug,
block_swap_config=block_swap_config,
cached_runner=self.runner # Pass existing runner if any
)
self.current_model_name = model
if debug:
print(f"🔄 Runner configuration time: {time.time() - runner_start:.2f}s")
if debug:
print("🚀 Starting video upscaling generation...")
# Execute generation with progress callback
sample = generation_loop(
self.runner, images, cfg_scale, seed, new_resolution,
batch_size, preserve_vram, temporal_overlap, debug,
block_swap_config=block_swap_config,
progress_callback=self._progress_callback
)
print(f"✅ Video upscaling completed successfully!")
print(f"✅ Video upscaling completed successfully!")
# Cleanup
print(f"🔄 Total execution time: {time.time() - total_start_time:.2f}s")
self.cleanup(force_ram_cleanup=True)
self.cleanup(force_ram_cleanup=True, keep_model_cached=use_cache, block_swap_config=block_swap_config)
return (sample,)
def _progress_callback(self, batch_idx, total_batches, current_batch_frames, message=""):
@@ -212,19 +307,128 @@ class SeedVR2:
def __del__(self):
"""Destructor"""
try:
self.cleanup(force_ram_cleanup=True)
self.cleanup(force_ram_cleanup=True, keep_model_cached=False, block_swap_config=None)
except:
pass
class SeedVR2BlockSwap:
"""Configure block swapping to reduce VRAM usage"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"blocks_to_swap": (
"INT",
{
"default": 16,
"min": 0,
"max": 36,
"step": 1,
"tooltip": "Number of transformer blocks to swap to CPU. Start with 16 and increase until OOM errors stop. 0=disabled",
},
),
"use_non_blocking": (
"BOOLEAN",
{
"default": True,
"tooltip": "Use non-blocking GPU transfers for better performance.",
},
),
"offload_io_components": (
"BOOLEAN",
{
"default": False,
"tooltip": "Offload embeddings and I/O layers to CPU. Enable if you need additional VRAM savings beyond block swapping",
},
),
"cache_model": (
"BOOLEAN",
{
"default": False,
"tooltip": "Keep model in RAM between runs to avoid model loading time. Useful for batch processing",
},
),
"enable_debug": (
"BOOLEAN",
{
"default": False,
"tooltip": "Show detailed memory usage and timing information during inference",
},
),
}
}
RETURN_TYPES = ("block_swap_config",)
FUNCTION = "create_config"
CATEGORY = "SEEDVR2"
DESCRIPTION = """Configure block swapping to reduce VRAM usage during video upscaling.
BlockSwap dynamically moves transformer blocks between GPU and CPU/RAM during inference, enabling large models to run on limited VRAM systems with minimal performance impact.
Configuration Guidelines:
- blocks_to_swap=0: Disabled (fastest, highest VRAM usage)
- blocks_to_swap=16: Balanced mode (moderate speed/VRAM trade-off)
- blocks_to_swap=32-36: Maximum savings (slowest, lowest VRAM)
Advanced Options:
- use_non_blocking: Enables asynchronous GPU transfers for better performance
- offload_io_components: Moves embeddings and I/O layers to CPU for additional VRAM savings (slower)
- cache_model: Keeps model in RAM between runs (avoids model loading time on subsequent generations)
- enable_debug: Shows detailed memory usage and timing information
Performance Tips:
- Start with blocks_to_swap=16 and increase until you no longer get OOM errors or decrease if you have spare VRAM
- Enable offload_io_components if you still need additional VRAM savings
- Note: Even if inference succeeds, you may still OOM during VAE decoding - combine BlockSwap with VAE tiling if needed (feature in development)
- Enable cache_model for batch processing to skip model reloading between runs
- Keep non_blocking=True for better performance (default)
- Combine with smaller batch_size for maximum VRAM savings
The actual memory savings depend on your specific model architecture and will be shown in the debug output when enabled.
"""
def create_config(
self,
blocks_to_swap,
use_non_blocking,
offload_io_components,
cache_model,
enable_debug,
):
if blocks_to_swap > 0 or offload_io_components:
configs = []
if blocks_to_swap > 0:
configs.append(f"{blocks_to_swap} blocks")
if use_non_blocking:
configs.append("non blocking")
if offload_io_components:
configs.append("I/O components")
if cache_model and blocks_to_swap > 0:
configs.append("model caching")
print(f"🔄 BlockSwap configured: {', '.join(configs)}")
return (
{
"blocks_to_swap": blocks_to_swap,
"use_non_blocking": use_non_blocking,
"offload_io_components": offload_io_components,
"cache_model": cache_model,
"enable_debug": enable_debug,
},
)
# ComfyUI Node Mappings
NODE_CLASS_MAPPINGS = {
"SeedVR2": SeedVR2,
"SeedVR2BlockSwap": SeedVR2BlockSwap,
}
# Human-readable node display names
NODE_DISPLAY_NAME_MAPPINGS = {
"SeedVR2": "SeedVR2 Video Upscaler",
"SeedVR2BlockSwap": "SeedVR2 BlockSwap Config",
}
# Export version and metadata
+735
View File
@@ -0,0 +1,735 @@
"""
BlockSwap Module for SeedVR2
This module implements dynamic block swapping between GPU and CPU memory
to enable running large models on limited VRAM systems.
Key Features:
- Dynamic transformer block offloading during inference
- Non-blocking GPU transfers for optimal performance
- RoPE computation fallback to CPU on OOM
- Minimal performance overhead with intelligent caching
- I/O component offloading for maximum memory savings
"""
import time
import types
import torch
import weakref
import psutil
import gc
import comfy.model_management as mm
from typing import Dict, Any, List, Tuple, Optional, Union
from src.optimization.memory_manager import get_vram_usage
def get_module_memory_mb(module: torch.nn.Module) -> float:
"""
Calculate memory usage of a module in MB.
Args:
module: PyTorch module to measure
Returns:
Memory usage in megabytes
"""
total_bytes = sum(
param.nelement() * param.element_size()
for param in module.parameters()
if param.data is not None
)
return total_bytes / (1024 * 1024)
class BlockSwapDebugger:
"""
Debug logger for BlockSwap operations.
Tracks memory usage, swap timings, and provides detailed logging
for debugging and performance analysis of block swapping operations.
"""
def __init__(self, enabled: bool = False):
"""
Initialize the debugger.
Args:
enabled: Whether debug logging is enabled
"""
self.enabled = enabled
self.swap_times: List[Tuple[int, float, str]] = []
self.vram_history: List[float] = []
def log(self, message: str, level: str = "INFO") -> None:
"""Log a message if debugging is enabled."""
if self.enabled:
print(f"[{level}] {message}")
def log_swap_time(self, component_id, duration: float, component_type: str = "block", direction: str = "compute") -> None:
"""
Log swap timing information for any component (blocks or I/O).
Args:
component_id: Block index (int) or I/O component name (str)
duration: Time taken for the swap operation
component_type: Type of component ("block" or "io")
direction: Direction of swap ("compute" or "offload")
"""
# Store timing data with component info
self.swap_times.append({
'component_id': component_id,
'component_type': component_type,
'duration': duration,
'direction': direction
})
if self.enabled:
# Format message based on component type
if component_type == "block":
message = f"Block {component_id} swap {direction}: {duration*1000:.1f}ms"
elif component_type == "io":
message = f"I/O {component_id} swap {direction}: {duration*1000:.1f}ms"
else:
message = f"{component_type} {component_id} swap {direction}: {duration*1000:.1f}ms"
self.log(message, "SWAP")
def log_memory_state(self, stage: str, show_tensors: bool = False) -> None:
"""Log current memory state for debugging."""
# GPU Memory
if torch.cuda.is_available():
allocated_gb, reserved_gb, peak_gb = get_vram_usage()
vram_info = f"VRAM: {allocated_gb:.2f}/{reserved_gb:.2f}GB (peak: {peak_gb:.2f}GB)"
self.vram_history.append(allocated_gb)
else:
vram_info = "VRAM: CPU mode"
# RAM Memory
ram_info = ""
if psutil:
try:
process = psutil.Process()
ram_gb = process.memory_info().rss / (1024**3)
ram_info = f" | RAM: {ram_gb:.1f}GB"
except Exception:
pass
# Tensor count (optional - expensive operation)
tensor_info = ""
if show_tensors:
tensor_count = sum(1 for obj in gc.get_objects() if torch.is_tensor(obj))
tensor_info = f" | Tensors: {tensor_count}"
self.log(f"🧮 {stage}: {vram_info}{ram_info}{tensor_info}")
def clear_history(self) -> None:
"""Clear accumulated history."""
self.swap_times.clear()
self.vram_history.clear()
def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any]) -> None:
"""
Apply block swapping configuration to a DIT model with OOM protection.
This is the main entry point for configuring block swapping on a model.
It handles block selection, I/O component offloading, and device placement.
Args:
runner: VideoDiffusionInfer instance containing the model
block_swap_config: Configuration dictionary with keys:
- blocks_to_swap: Number of blocks to swap (from the start)
- offload_io_components: Whether to offload I/O components
- use_non_blocking: Whether to use non-blocking transfers
- enable_debug: Whether to enable debug logging
"""
if not block_swap_config:
return
blocks_to_swap = block_swap_config.get("blocks_to_swap", 0)
if blocks_to_swap <= 0:
return
# Always use fresh debugger for clean state
enable_debug = block_swap_config.get("enable_debug", False)
# Clean up old debugger if exists
if hasattr(runner, '_blockswap_debugger'):
old_debugger = runner._blockswap_debugger
if old_debugger:
old_debugger.clear_history()
delattr(runner, '_blockswap_debugger')
# Create new debugger
debugger = BlockSwapDebugger(enabled=enable_debug)
runner._blockswap_debugger = debugger
# Get the actual model (handle FP8CompatibleDiT wrapper)
model = runner.dit
if hasattr(model, "dit_model"):
model = model.dit_model
# Determine devices
device = "cuda" if torch.cuda.is_available() else "cpu"
offload_device = str(mm.unet_offload_device())
use_non_blocking = block_swap_config.get("use_non_blocking", True)
# Validate model structure
if not hasattr(model, "blocks"):
debugger.log("Model doesn't have 'blocks' attribute for BlockSwap", "WARN")
return
total_blocks = len(model.blocks)
debugger.log(f"Model has {total_blocks} blocks total")
blocks_to_swap = min(blocks_to_swap, total_blocks)
# Configure model with blockswap attributes
model.blocks_to_swap = blocks_to_swap - 1 # Convert to 0-indexed
model.main_device = device
model.offload_device = offload_device
model.use_non_blocking = use_non_blocking
debugger.log(f"Configuring: {blocks_to_swap}/{total_blocks} blocks for swapping")
debugger.log_memory_state("Before BlockSwap", show_tensors=True)
# Configure I/O components
offload_io_components = block_swap_config.get("offload_io_components", False)
io_components_offloaded = _configure_io_components(model, device, offload_device, use_non_blocking,
offload_io_components, debugger)
# Configure block placement and memory tracking
memory_stats = _configure_blocks(model, device, offload_device, use_non_blocking, debugger)
memory_stats['io_components'] = io_components_offloaded
# Log memory summary
_log_memory_summary(memory_stats, offload_device, device, offload_io_components,
use_non_blocking, debugger)
# Wrap block forward methods for dynamic swapping
for b, block in enumerate(model.blocks):
if b <= model.blocks_to_swap:
_wrap_block_forward(block, b, model, debugger)
# Patch RoPE modules for robust error handling
_patch_rope_for_blockswap(model, debugger)
# Mark BlockSwap as active
runner._blockswap_active = True
# Store configuration for debugging and cleanup
runner._block_swap_config = {
"blocks_swapped": blocks_to_swap,
"offload_io_components": offload_io_components,
"total_blocks": total_blocks,
"use_non_blocking": use_non_blocking,
"offload_device": offload_device,
"main_device": device,
"enable_debug": block_swap_config.get("enable_debug", False),
"offload_memory": memory_stats['offload_memory'],
"main_memory": memory_stats['main_memory']
}
# Protect model from being moved entirely
_protect_model_from_move(model, runner, debugger)
debugger.log_memory_state("After BlockSwap", show_tensors=True)
debugger.log("✅ BlockSwap configuration complete")
def _configure_io_components(model, device: str, offload_device: str,
use_non_blocking: bool, offload_io_components: bool,
debugger: BlockSwapDebugger) -> List[str]:
"""Configure I/O component placement and wrapping."""
io_components_offloaded = []
# Process non-block parameters
for name, param in model.named_parameters():
if "block" not in name:
target_device = offload_device if offload_io_components else device
param.data = param.data.to(target_device, non_blocking=use_non_blocking)
status = "(offloaded)" if offload_io_components else ""
debugger.log(f" {name} → {target_device} {status}")
# Handle I/O modules with dynamic swapping
for name, module in model.named_children():
if name != "blocks":
if offload_io_components:
module.to(offload_device)
_wrap_io_forward(module, name, model, debugger)
io_components_offloaded.append(name)
debugger.log(f" {name} → {offload_device} (with dynamic swapping)")
else:
module.to(device)
debugger.log(f" {name} → {device}")
return io_components_offloaded
def _configure_blocks(model, device: str, offload_device: str,
use_non_blocking: bool, debugger: BlockSwapDebugger) -> Dict[str, float]:
"""Configure block placement and calculate memory statistics."""
total_offload_memory = 0.0
total_main_memory = 0.0
# Move blocks based on swap configuration
for b, block in enumerate(model.blocks):
block_memory = get_module_memory_mb(block)
if b > model.blocks_to_swap:
block.to(device)
total_main_memory += block_memory
else:
block.to(offload_device, non_blocking=use_non_blocking)
total_offload_memory += block_memory
# Ensure all buffers match their containing module's device
for b, block in enumerate(model.blocks):
target_device = device if b > model.blocks_to_swap else offload_device
for name, buffer in block.named_buffers():
if buffer.device != torch.device(target_device):
buffer.data = buffer.data.to(target_device)
# Clean up memory
mm.soft_empty_cache()
gc.collect()
return {
"offload_memory": total_offload_memory,
"main_memory": total_main_memory,
"io_components": [] # Will be populated by caller
}
def _log_memory_summary(memory_stats: Dict[str, float], offload_device: str,
device: str, offload_io_components: bool,
use_non_blocking: bool, debugger: BlockSwapDebugger) -> None:
"""Log memory usage summary."""
debugger.log("----------------------")
debugger.log("Block swap memory summary:")
debugger.log(f"Transformer blocks on {offload_device}: {memory_stats['offload_memory']:.2f}MB")
debugger.log(f"Transformer blocks on {device}: {memory_stats['main_memory']:.2f}MB")
total_memory = memory_stats['offload_memory'] + memory_stats['main_memory']
debugger.log(f"Total memory used by transformer blocks: {total_memory:.2f}MB")
if offload_io_components and memory_stats.get('io_components'):
debugger.log(f"I/O components offloaded: {', '.join(memory_stats['io_components'])}")
debugger.log(f"Non-blocking memory transfer: {use_non_blocking}")
debugger.log("----------------------")
def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn.Module, debugger: BlockSwapDebugger) -> None:
"""Wrap individual block forward to handle device movement using weak references to prevent leaks."""
if hasattr(block, '_original_forward'):
return # Already wrapped
# Store original forward method
original_forward = block.forward
# Create weak references
model_ref = weakref.ref(model)
debugger_ref = weakref.ref(debugger)
# Store block_idx on the block itself to avoid closure issues
block._block_idx = block_idx
def wrapped_forward(self, *args, **kwargs):
# Retrieve weak references
model = model_ref()
debugger = debugger_ref()
if not model:
# Model has been garbage collected, fall back to original
return original_forward(*args, **kwargs)
# Check if block swap is active for this block
if hasattr(model, 'blocks_to_swap') and self._block_idx <= model.blocks_to_swap:
t_start = time.time() if debugger and debugger.enabled else None
# Only move to GPU if necessary
current_device = next(self.parameters()).device
target_device = torch.device(model.main_device)
if current_device != target_device:
self.to(model.main_device, non_blocking=model.use_non_blocking)
# Synchronize if needed
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
torch.cuda.synchronize()
# Execute forward pass with OOM protection
output = original_forward(*args, **kwargs)
# Move back to offload device
self.to(model.offload_device, non_blocking=model.use_non_blocking)
# Log timing if debugger is available
if debugger and t_start is not None:
debugger.log_swap_time(
component_id=self._block_idx,
duration=time.time() - t_start,
component_type="block",
direction="compute"
)
# Only clear cache under memory pressure
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
else:
output = original_forward(*args, **kwargs)
return output
# Bind the wrapped function as a method to the block
block.forward = types.MethodType(wrapped_forward, block)
# Store reference to original forward for cleanup
block._original_forward = original_forward
def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn.Module, debugger: BlockSwapDebugger) -> None:
"""Wrap I/O component forward using weak references to prevent memory leaks."""
if hasattr(module, '_is_io_wrapped') and module._is_io_wrapped:
return # Already wrapped
# Store original forward method
original_forward = module.forward
# Create weak references
model_ref = weakref.ref(model)
debugger_ref = weakref.ref(debugger) if debugger else lambda: None
# Store module name on the module itself
module._module_name = module_name
module._original_forward = original_forward
def wrapped_io_forward(self, *args, **kwargs):
# Retrieve weak references
model = model_ref()
debugger = debugger_ref()
if not model:
# Model has been garbage collected, fall back to original
return self._original_forward(*args, **kwargs)
t_start = time.time() if debugger and debugger.enabled else None
# Check current device to avoid unnecessary moves
current_device = next(self.parameters()).device
target_device = torch.device(model.main_device)
# Move to GPU for computation if needed
if current_device != target_device:
self.to(model.main_device)
# Synchronize if not using non-blocking transfers
if hasattr(model, 'use_non_blocking') and not model.use_non_blocking:
torch.cuda.synchronize()
# Execute forward pass
output = self._original_forward(*args, **kwargs)
# Move back to offload device
self.to(model.offload_device, non_blocking=model.use_non_blocking)
# Log timing if debugger is available
if debugger and t_start is not None:
debugger.log_swap_time(
component_id=self._module_name,
duration=time.time() - t_start,
component_type="io",
direction="compute"
)
# Only clear cache under memory pressure
if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9:
mm.soft_empty_cache()
return output
# Bind as a method
module.forward = types.MethodType(wrapped_io_forward, module)
module._is_io_wrapped = True
# Store module reference for restoration
if not hasattr(model, '_io_swappers'):
model._io_swappers = []
model._io_swappers.append((module, module_name))
def _patch_rope_for_blockswap(model, debugger: BlockSwapDebugger) -> None:
"""
Patch RoPE modules to handle device mismatches gracefully.
RoPE (Rotary Position Embeddings) can cause device mismatches when
blocks are on different devices. This patches the get_axial_freqs
method to handle these cases robustly.
"""
rope_patches = []
for name, module in model.named_modules():
if "rope" in name.lower() and hasattr(module, "get_axial_freqs"):
original_method = module.get_axial_freqs
def robust_rope_wrapper(self, *args, **kwargs):
try:
return original_method(*args, **kwargs)
except (RuntimeError, KeyError) as e:
error_msg = str(e).lower()
if "device" in error_msg or "memory" in error_msg or "allocation" in error_msg:
debugger.log(f"RoPE issue for {name}: {e}")
# Get current device from parameters
current_device = "cuda"
if list(self.parameters()):
current_device = next(self.parameters()).device
# Try with cleared cache first
if hasattr(original_method, 'cache_clear'):
original_method.cache_clear()
try:
return original_method(*args, **kwargs)
except:
pass
# Fallback to CPU computation
debugger.log(f"RoPE fallback to CPU for {name}")
self.cpu()
try:
result = original_method(*args, **kwargs)
# Move module back to original device
self.to(current_device)
# Move result to appropriate device if it's a tensor
if hasattr(result, 'to'):
if len(args) > 0 and hasattr(args[0], 'device'):
return result.to(args[0].device)
return result.to(current_device)
return result
except Exception as cpu_error:
# Always restore device even on error
self.to(current_device)
raise cpu_error
else:
raise
module.get_axial_freqs = types.MethodType(robust_rope_wrapper, module)
rope_patches.append((module, original_method))
if rope_patches:
model._rope_patches = rope_patches
debugger.log(f"✅ Patched {len(rope_patches)} RoPE modules with robust device handling")
def _protect_model_from_move(model, runner, debugger: BlockSwapDebugger) -> None:
"""
Protect model from being moved entirely to GPU when BlockSwap is active.
This prevents other code from accidentally moving the entire model to GPU
which would defeat the purpose of block swapping.
"""
if not hasattr(model, '_original_to'):
# Store runner reference as weak reference to avoid circular refs
model._blockswap_runner_ref = weakref.ref(runner)
model._original_to = model.to
# Define the protected method without closures
def protected_model_to(self, device, *args, **kwargs):
# Check blockswap status using weak reference
if str(device) != "cpu":
runner_ref = getattr(self, '_blockswap_runner_ref', None)
if runner_ref:
runner_obj = runner_ref()
if runner_obj and hasattr(runner_obj, "_blockswap_active") and runner_obj._blockswap_active:
print("[INFO] ⚠️ Blocked attempt to move blockswapped model to GPU")
return self
# Use original method stored as attribute
if hasattr(self, '_original_to'):
return self._original_to(device, *args, **kwargs)
else:
# This shouldn't happen, but fallback to super().to()
return super(type(self), self).to(device, *args, **kwargs)
# Bind as a method to the model instance
model.to = types.MethodType(protected_model_to, model)
def cleanup_blockswap(runner, keep_state_for_cache: bool = False) -> None:
"""
Clean up BlockSwap configurations and restore original methods.
This should be called when BlockSwap is no longer needed to restore
the model to its original state and free up any resources.
Args:
runner: VideoDiffusionInfer instance to clean up
keep_state_for_cache: If True, stores configuration for fast re-application
"""
# Early return if BlockSwap not active
if not hasattr(runner, "_blockswap_active") or not runner._blockswap_active:
print("[INFO] ⚠️ BlockSwap not active, skipping cleanup")
return
# Use existing debugger if available
debugger = getattr(runner, '_blockswap_debugger', None)
if debugger is None:
# Create new debugger only if none exists
debugger = BlockSwapDebugger(enabled=runner._block_swap_config.get("enable_debug", False))
runner._blockswap_debugger = debugger
else:
debugger.clear_history()
debugger.log("🧹 Starting BlockSwap cleanup")
# Get the actual model (handle FP8CompatibleDiT wrapper)
model = runner.dit
if hasattr(model, "dit_model"):
model = model.dit_model
# Store configuration BEFORE cleanup if caching
cached_config = None
if keep_state_for_cache and hasattr(runner, "_block_swap_config"):
cached_config = {
"blocks_to_swap": runner._block_swap_config.get("blocks_swapped"),
"offload_io_components": runner._block_swap_config.get("offload_io_components"),
"use_non_blocking": runner._block_swap_config.get("use_non_blocking"),
"offload_device": runner._block_swap_config.get("offload_device"),
"main_device": runner._block_swap_config.get("main_device"),
"enable_debug": runner._block_swap_config.get("enable_debug", False),
}
runner._cached_blockswap_config = cached_config
debugger.log("📦 Storing configuration for fast re-application")
# Restore block forward methods
if hasattr(model, 'blocks'):
restored_count = 0
for idx, block in enumerate(model.blocks):
if hasattr(block, '_original_forward'):
block.forward = block._original_forward
delattr(block, '_original_forward')
restored_count += 1
# Clean up ALL wrapper attributes
attrs_to_clean = ['_block_idx', '_model_ref', '_debugger_ref', '_blockswap_wrapped']
for attr in attrs_to_clean:
if hasattr(block, attr):
delattr(block, attr)
# Clear gradients to free memory
block.zero_grad(set_to_none=True)
# Move block to CPU and ensure all buffers follow
if not keep_state_for_cache:
block.to("cpu")
# Force memory deallocation for all parameters and buffers
for param in block.parameters():
if param.data.numel() > 0:
param.data.set_()
for buffer in block.buffers():
if buffer.data.numel() > 0:
buffer.data.set_()
if restored_count > 0:
debugger.log(f"✅ Restored original forward for {restored_count} blocks")
# Restore RoPE methods and clear LRU caches
if hasattr(model, '_rope_patches'):
for module, original_method in model._rope_patches:
# Clear the LRU cache before restoring
if hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
module.get_axial_freqs = original_method
debugger.log(f"✅ Restored {len(model._rope_patches)} RoPE modules")
delattr(model, '_rope_patches')
else:
# Fallback: Clear RoPE caches without restoration
cleared_count = 0
for module in model.modules():
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
cleared_count += 1
if cleared_count > 0:
debugger.log(f"✅ Cleared {cleared_count} RoPE LRU caches")
# Restore I/O component forward methods
if hasattr(model, '_io_swappers'):
for module, module_name in model._io_swappers:
if hasattr(module, '_is_io_wrapped') and hasattr(module, '_original_forward'):
module.forward = module._original_forward
# Clean up wrapper attributes
attrs_to_clean = ['_original_forward', '_model_ref', '_debugger_ref',
'_module_name', '_is_io_wrapped']
for attr in attrs_to_clean:
if hasattr(module, attr):
delattr(module, attr)
debugger.log(f"✅ Restored {len(model._io_swappers)} I/O component wrappers")
delattr(model, '_io_swappers')
# Restore original .to() method
if hasattr(model, '_original_to'):
model.to = model._original_to
delattr(model, '_original_to')
debugger.log("✅ Restored original .to() method")
# Clean up weak reference on model
if hasattr(model, '_blockswap_runner_ref'):
delattr(model, '_blockswap_runner_ref')
# Clean up BlockSwap attributes from model
attrs_to_remove = ["blocks_to_swap", "main_device", "offload_device", "use_non_blocking"]
for attr in attrs_to_remove:
if hasattr(model, attr):
delattr(model, attr)
# Mark model as not configured
if hasattr(model, '_blockswap_configured'):
delattr(model, '_blockswap_configured')
# Move model to CPU to free VRAM (safe now that wrappers are removed)
if not keep_state_for_cache:
model.to("cpu")
debugger.log("📦 Moved model to CPU")
# Clean up runner attributes
runner._blockswap_active = False
# Remove all config attributes if not caching
if not cached_config:
if hasattr(runner, "_cached_blockswap_config"):
delattr(runner, "_cached_blockswap_config")
if hasattr(runner, "_block_swap_config"):
delattr(runner, "_block_swap_config")
# Clear debugger reference (only if not caching)
if not keep_state_for_cache and hasattr(runner, '_blockswap_debugger'):
delattr(runner, '_blockswap_debugger')
# Clear local debugger reference
debugger = None
# Force garbage collection (multiple passes for thorough cleanup)
gc.collect(2) # Full collection including oldest generation
gc.collect()
gc.collect()
# Final memory cleanup
mm.soft_empty_cache()
+20 -18
View File
@@ -19,31 +19,33 @@ class FP8CompatibleDiT(torch.nn.Module):
- Flash Attention: Automatic optimization of attention layers
"""
def __init__(self, dit_model):
def __init__(self, dit_model, skip_conversion=False):
super().__init__()
self.dit_model = dit_model
self.model_dtype = self._detect_model_dtype()
self.is_fp8_model = self.model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
self.is_fp16_model = self.model_dtype == torch.float16
# Detect model type
is_nadit_7b = self._is_nadit_model() # NaDiT 7B (dit/nadit)
is_nadit_v2_3b = self._is_nadit_v2_model() # NaDiT v2 3B (dit_v2/nadit)
if is_nadit_7b:
# 🎯 CRITICAL FIX: ALL NaDiT 7B models (FP8 AND FP16) require BFloat16 conversion
# 7B architecture has dtype compatibility issues regardless of storage format
if self.is_fp8_model:
print("🎯 Detected NaDiT 7B FP8 - Converting all parameters to BFloat16")
# Only convert if not already done (e.g., when reusing cached weights)
if not skip_conversion:
# Detect model type
is_nadit_7b = self._is_nadit_model() # NaDiT 7B (dit/nadit)
is_nadit_v2_3b = self._is_nadit_v2_model() # NaDiT v2 3B (dit_v2/nadit)
if is_nadit_7b:
# 🎯 CRITICAL FIX: ALL NaDiT 7B models (FP8 AND FP16) require BFloat16 conversion
# 7B architecture has dtype compatibility issues regardless of storage format
if self.is_fp8_model:
print("🎯 Detected NaDiT 7B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
else:
print("🎯 Detected NaDiT 7B FP16")
elif self.is_fp8_model and is_nadit_v2_3b:
# For NaDiT v2 3B FP8: Convert ALL model to BFloat16
print("🎯 Detected NaDiT v2 3B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
else:
print("🎯 Detected NaDiT 7B FP16")
elif self.is_fp8_model and is_nadit_v2_3b:
# For NaDiT v2 3B FP8: Convert ALL model to BFloat16
print("🎯 Detected NaDiT v2 3B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
# 🚀 FLASH ATTENTION OPTIMIZATION (Phase 2)
self._apply_flash_attention_optimization()
+143 -2
View File
@@ -12,7 +12,7 @@ import time
from typing import Tuple, Optional
from src.common.cache import Cache
from src.models.dit_v2.rope import RotaryEmbeddingBase
from comfy import model_management as mm
def get_basic_vram_info():
"""🔍 Méthode basique avec PyTorch natif"""
@@ -248,4 +248,145 @@ def fast_ram_cleanup():
try:
torch._C._clear_cache()
except:
pass
pass
def clear_all_caches(runner, debugger=None) -> int:
"""
Aggressively clear all caches from runner and model.
Optimized to only process what's necessary.
Args:
runner: The runner instance to clear caches from
debugger: Optional BlockSwapDebugger instance for logging
"""
if not runner:
return 0
# Try to get debugger from runner if not provided
if debugger is None and hasattr(runner, '_blockswap_debugger'):
debugger = runner._blockswap_debugger
# Helper function for logging
def log_message(message, level="INFO"):
if debugger and debugger.enabled:
debugger.log(message, level)
else:
print(f" {message}")
cleaned_items = 0
# Early exit if no caches to clear
has_cache = hasattr(runner, 'cache') and hasattr(runner.cache, 'cache')
if not has_cache and not hasattr(runner, 'dit'):
return 0
# Clear main runner cache efficiently
if has_cache and runner.cache.cache:
cache_entries = len(runner.cache.cache)
# Process all cache items to properly free memory
for key, value in list(runner.cache.cache.items()):
if torch.is_tensor(value):
# Force deallocation of tensor storage
if value.is_cuda:
value.data = value.data.cpu()
value.grad = None
if value.numel() > 0:
value.data.set_() # Release underlying storage
elif isinstance(value, (list, tuple)):
for item in value:
if torch.is_tensor(item):
if item.is_cuda:
item.data = item.data.cpu()
item.grad = None
if item.numel() > 0:
item.data.set_()
# Clear the cache after processing
runner.cache.cache.clear()
cleaned_items += cache_entries
log_message(f"✅ Cleared {cache_entries} cache entries")
# Clear any accumulated state in blocks
if hasattr(runner, 'dit'):
model = runner.dit
if hasattr(model, 'dit_model'):
model = model.dit_model
# Clear RoPE LRU caches
rope_caches_cleared = clear_rope_lru_caches(model)
cleaned_items += rope_caches_cleared
if rope_caches_cleared > 0:
log_message(f"✅ Cleared {rope_caches_cleared} RoPE LRU caches")
# Clear block attributes if needed
if hasattr(model, 'blocks'):
block_attrs_cleared = 0
# Define PyTorch's essential attributes that must NOT be deleted
essential_attrs = {
'_modules', '_parameters', '_buffers',
'_forward_hooks', '_forward_pre_hooks',
'_backward_hooks', '_backward_pre_hooks',
'_state_dict_hooks', '_state_dict_pre_hooks',
'_load_state_dict_pre_hooks', '_load_state_dict_post_hooks',
'_non_persistent_buffers_set', '_version',
'_is_full_backward_hook', 'training',
'_original_forward', # BlockSwap attribute
'_is_io_wrapped', # BlockSwap attribute
'_block_idx', # BlockSwap attribute
}
for idx, block in enumerate(model.blocks):
# Get all attributes that look like caches
attrs_to_remove = []
for attr_name in list(block.__dict__.keys()):
# Only remove cache-like attributes, not essential PyTorch attributes
if (attr_name not in essential_attrs and
('cache' in attr_name or
'temp' in attr_name or
(attr_name.startswith('_') and
not attr_name.startswith('__') and
attr_name not in essential_attrs))):
attrs_to_remove.append(attr_name)
# Remove the identified attributes
for attr_name in attrs_to_remove:
try:
delattr(block, attr_name)
block_attrs_cleared += 1
except AttributeError:
pass # Already deleted or doesn't exist
if block_attrs_cleared > 0:
log_message(f"✅ Cleared {block_attrs_cleared} temporary attributes from blocks")
# Clear any temporary attributes that might accumulate
temp_attrs = ['_temp_cache', '_block_cache', '_swap_cache', '_generation_cache',
'_rope_cache', '_intermediate_cache', '_backward_cache']
# Check both runner and model for these attributes
for obj in [runner, getattr(runner, 'dit', None)]:
if obj is None:
continue
# Handle wrapped models
if hasattr(obj, 'dit_model'):
obj = obj.dit_model
for attr in temp_attrs:
if hasattr(obj, attr):
delattr(obj, attr)
cleaned_items += 1
log_message(f"✅ Cleared {attr} from {type(obj).__name__}")
# Force garbage collection
gc.collect(2) # Collect all generations
# Clear CUDA cache
if torch.cuda.is_available():
torch.cuda.empty_cache()
mm.soft_empty_cache()
return cleaned_items