Merge pull request #31 from AInVFX/blockswap
BlockSwap support, thanks To @adrientoupet
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user