refactor: modularize node architecture and enhance BlockSwap flexibility

Core Changes:
- Split monolithic comfyui_node.py into modular per-node files:
  * video_upscaler.py - main upscaler node
  * dit_model_loader.py - DiT model configuration
  * vae_model_loader.py - VAE model configuration
  * torch_compile_settings.py - torch.compile settings
  * __init__.py - centralized node registry

BlockSwap Enhancements:
- Fixed swap_io_components to work independently of blocks_to_swap
- Both features can now be enabled separately or together
- Updated validation logic throughout (model_manager, blockswap, generation)
- Improved config description to reflect combined capabilities

Logging Improvements:
- Standardized indentation and improved debug messages clarity
- Added footer with support links (YouTube, GitHub) for community engagement
This commit is contained in:
Adrien Toupet
2025-10-11 17:37:58 -04:00
parent bf46dfd674
commit d6e04aa108
10 changed files with 544 additions and 539 deletions
+1 -1
View File
@@ -1,3 +1,3 @@
from .src.interfaces.comfyui_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .src.interfaces import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+1 -1
View File
@@ -1190,7 +1190,7 @@ def postprocess_all_batches(ctx: Optional[Dict[str, Any]] = None,
# Format channel info for readability
channels_str = "RGBA" if Cf == 4 else "RGB" if Cf == 3 else f"{Cf}-channel"
debug.log(f"Final video assembled: Total frames: {total_frames}, Resolution: {Wf}x{Hf}px, Channels: {channels_str}", category="video", force=True)
debug.log(f"Final video assembled: Total frames: {total_frames}, Resolution: {Wf}x{Hf}px, Channels: {channels_str}", category="generation", force=True)
else:
ctx['final_video'] = torch.empty((0, 0, 0, 0), dtype=torch.float16)
debug.log("No frame to assemble", level="WARNING", category="video", force=True)
+25 -15
View File
@@ -93,13 +93,16 @@ def _describe_blockswap_config(config: Optional[Dict[str, Any]]) -> str:
if config is None:
return "disabled"
blocks = config.get("blocks_to_swap", 0)
if blocks <= 0:
blocks_to_swap = config.get("blocks_to_swap", 0)
swap_io_components = config.get("swap_io_components", False)
# Early return only if both block swap and I/O swap are disabled
if blocks_to_swap <= 0 and not swap_io_components:
return "disabled"
block_text = "block" if blocks == 1 else "blocks"
parts = [f"{blocks} {block_text}"]
if config.get("swap_io_components", False):
block_text = "block" if blocks_to_swap <= 1 else "blocks"
parts = [f"{blocks_to_swap} {block_text}"]
if swap_io_components:
parts.append("I/O offload")
return f"enabled ({', '.join(parts)})"
@@ -338,11 +341,17 @@ def _update_dit_config_inplace(
if blockswap_changed:
# Determine change type from config comparison
old_blocks = cached_blockswap_config.get("blocks_to_swap", 0) if cached_blockswap_config else 0
old_swap_io = cached_blockswap_config.get("swap_io_components", False) if cached_blockswap_config else False
new_blocks = block_swap_config.get("blocks_to_swap", 0) if block_swap_config else 0
new_swap_io = block_swap_config.get("swap_io_components", False) if block_swap_config else False
# If old config had BlockSwap, clean it up first
if old_blocks > 0:
if new_blocks == 0:
# Check if old config had ANY BlockSwap features (blocks or I/O components)
had_blockswap = old_blocks > 0 or old_swap_io
has_blockswap = new_blocks > 0 or new_swap_io
# If old config had BlockSwap features, clean them up first
if had_blockswap:
if not has_blockswap:
# Disabling BlockSwap completely
debug.log("Disabling BlockSwap completely", category="blockswap")
cleanup_blockswap(runner=cached_runner, keep_state_for_cache=False)
@@ -354,7 +363,8 @@ def _update_dit_config_inplace(
cached_runner._blockswap_active = False
# Apply new BlockSwap if configured (before torch.compile)
if block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0:
if block_swap_config and (block_swap_config.get("blocks_to_swap", 0) > 0 or
block_swap_config.get("swap_io_components", False)):
cached_runner.dit = model # Temporarily set for BlockSwap
apply_block_swap_to_dit(cached_runner, block_swap_config, debug)
model = cached_runner.dit # Get potentially wrapped model
@@ -735,7 +745,7 @@ def _load_gguf_state(checkpoint_path: str, device: str, debug: Optional[Debug],
# Progress reporting
if (i + 1) % 100 == 0:
debug.log(f" Loaded {i+1}/{total_tensors} tensors...", category="dit")
debug.log(f" Loaded {i+1}/{total_tensors} tensors...", category="dit")
debug.log(f"Successfully loaded {len(state_dict)} tensors to {device}", category="success")
@@ -842,9 +852,9 @@ class GGUFTensor(torch.Tensor):
return final_result
except Exception as e:
self.debug.log(f"Numpy fallback also failed: {e}", level="WARNING", category="dit", force=True)
self.debug.log(f" Tensor type: {self.tensor_type}", level="WARNING", category="dit", force=True)
self.debug.log(f" Shape: {self.shape}", level="WARNING", category="dit", force=True)
self.debug.log(f" Target shape: {self.tensor_shape}", level="WARNING", category="dit", force=True)
self.debug.log(f" Tensor type: {self.tensor_type}", level="WARNING", category="dit", force=True)
self.debug.log(f" Shape: {self.shape}", level="WARNING", category="dit", force=True)
self.debug.log(f" Target shape: {self.tensor_shape}", level="WARNING", category="dit", force=True)
traceback.print_exc()
# Return regular tensor as last resort
@@ -1405,14 +1415,14 @@ def _report_parameter_mismatches(state: Dict[str, torch.Tensor],
if unmatched:
debug.log(f"Warning: {len(unmatched)} parameters from GGUF not found in model",
level="WARNING", category="dit", force=True)
debug.log(f" First few unmatched: {unmatched[:5]}", level="WARNING", category="dit", force=True)
debug.log(f" First few unmatched: {unmatched[:5]}", level="WARNING", category="dit", force=True)
# Check for missing model parameters
missing = [name for name in model_state if name not in loaded_names]
if missing:
debug.log(f"Warning: {len(missing)} model parameters not loaded from GGUF",
level="WARNING", category="dit", force=True)
debug.log(f" First few missing: {missing[:5]}", level="WARNING", category="dit", force=True)
debug.log(f" First few missing: {missing[:5]}", level="WARNING", category="dit", force=True)
def _initialize_meta_buffers_wrapped(model: torch.nn.Module, target_device: str, debug: Debug) -> None:
+38
View File
@@ -0,0 +1,38 @@
"""
SeedVR2 ComfyUI Nodes
Central registry for all SeedVR2 nodes
"""
from .video_upscaler import SeedVR2VideoUpscaler
from .dit_model_loader import SeedVR2LoadDiTModel
from .vae_model_loader import SeedVR2LoadVAEModel
from .torch_compile_settings import SeedVR2TorchCompileSettings
# ComfyUI Node Mappings -
NODE_CLASS_MAPPINGS = {
"SeedVR2VideoUpscaler": SeedVR2VideoUpscaler,
"SeedVR2LoadDiTModel": SeedVR2LoadDiTModel,
"SeedVR2LoadVAEModel": SeedVR2LoadVAEModel,
"SeedVR2TorchCompileSettings": SeedVR2TorchCompileSettings,
}
# Human-readable node display names - unchanged
NODE_DISPLAY_NAME_MAPPINGS = {
"SeedVR2VideoUpscaler": "SeedVR2 Video Upscaler",
"SeedVR2LoadDiTModel": "SeedVR2 (Down)Load DiT Model",
"SeedVR2LoadVAEModel": "SeedVR2 (Down)Load VAE Model",
"SeedVR2TorchCompileSettings": "SeedVR2 Torch Compile Settings",
}
__version__ = "2.0.0"
__author__ = "numz, adrientoupet"
__description__ = "ComfyUI integration of ByteDance-Seed's SeedVR2: One-step diffusion-based video/image upscaling with memory-efficient inference"
__all__ = [
'SeedVR2VideoUpscaler',
'SeedVR2LoadDiTModel',
'SeedVR2LoadVAEModel',
'SeedVR2TorchCompileSettings',
'NODE_CLASS_MAPPINGS',
'NODE_DISPLAY_NAME_MAPPINGS'
]
+95
View File
@@ -0,0 +1,95 @@
"""
SeedVR2 DiT Model Loader Node
Configure DiT (Diffusion Transformer) model with memory optimization
"""
from typing import Dict, Any, Tuple
from ..utils.model_registry import get_available_dit_models, DEFAULT_DIT
from ..optimization.memory_manager import get_device_list
class SeedVR2LoadDiTModel:
"""
Configure DiT (Diffusion Transformer) model loader with memory optimization
Provides configuration for:
- Model selection and device placement
- BlockSwap memory optimization for limited VRAM
- Model caching between runs
- Optional torch.compile integration
Returns:
SEEDVR2_DIT configuration dictionary for main upscaler node
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
devices = get_device_list()
return {
"required": {
"model": (get_available_dit_models(), {
"default": DEFAULT_DIT,
"tooltip": "DiT model for upscaling. Models will automatically download on first use. Additional models can be added to the ComfyUI models folder."
}),
"device": (devices, {
"default": devices[0],
"tooltip": "Device to use for DiT processing"
}),
},
"optional": {
"blocks_to_swap": ("INT", {
"default": 0,
"min": 0,
"max": 36,
"step": 1,
"tooltip": "Number of transformer blocks to swap to CPU. 0=disabled (fastest, most VRAM). Higher values save VRAM but are slower. 3B model: 0-32 blocks, 7B model: 0-36 blocks."
}),
"swap_io_components": ("BOOLEAN", {
"default": False,
"tooltip": "Offload input/output embeddings and norm layers to CPU for additional VRAM savings (slower)"
}),
"cache_in_ram": ("BOOLEAN", {
"default": False,
"tooltip": "Keep model in RAM between runs for faster reuse. Useful for batch processing."
}),
"torch_compile_args": ("TORCH_COMPILE_ARGS", {
"tooltip": "Optional torch.compile settings from SeedVR2 Torch Compile Settings node for speedup"
}),
}
}
RETURN_TYPES = ("SEEDVR2_DIT",)
FUNCTION = "create_config"
CATEGORY = "SEEDVR2"
DESCRIPTION = (
"Configure DiT model for SeedVR2 upscaling. Supports BlockSwap for limited VRAM, "
"model caching, and torch.compile optimization. Connect output to SeedVR2 Video Upscaler."
)
def create_config(self, model: str, device: str, blocks_to_swap: int = 0,
swap_io_components: bool = False, cache_in_ram: bool = False,
torch_compile_args: Dict[str, Any] = None) -> Tuple[Dict[str, Any]]:
"""
Create DiT model configuration for SeedVR2 main node
Args:
model: Model filename to load
device: Target device for model execution
blocks_to_swap: Number of transformer blocks to swap to CPU (0=disabled)
swap_io_components: Whether to offload I/O components to CPU
cache_in_ram: Whether to keep model in RAM between runs
torch_compile_args: Optional torch.compile configuration from settings node
Returns:
Tuple containing configuration dictionary for SeedVR2 main node
"""
config = {
"model": model,
"device": device,
"blocks_to_swap": blocks_to_swap,
"swap_io_components": swap_io_components,
"cache_in_ram": cache_in_ram,
"torch_compile_args": torch_compile_args,
}
return (config,)
+91
View File
@@ -0,0 +1,91 @@
"""
SeedVR2 Torch Compile Settings Node
Configure torch.compile optimization for DiT and VAE models
"""
from typing import Dict, Any, Tuple
class SeedVR2TorchCompileSettings:
"""Configure torch.compile optimization for DiT and VAE models"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
return {
"required": {
"backend": (["inductor", "cudagraphs"], {
"default": "inductor",
"tooltip": (
"Compilation backend:\n"
"• inductor: Full optimization with Triton kernel generation and fusion\n"
"• cudagraphs: Lightweight, only wraps model with CUDA graphs, no kernel optimization"
)
}),
"mode": (["default", "reduce-overhead", "max-autotune", "max-autotune-no-cudagraphs"], {
"default": "default",
"tooltip": (
"Optimization level (compilation time vs runtime speed):\n"
"• default: Fast compilation, good speedup\n"
"• reduce-overhead: Lower overhead, better for smaller models\n"
"• max-autotune: Slowest compilation, best runtime (recommended for production)\n"
"• max-autotune-no-cudagraphs: Like max-autotune but without cudagraphs"
)
}),
"fullgraph": ("BOOLEAN", {
"default": False,
"tooltip": "Compile entire model as single graph (faster but less flexible). May fail with dynamic shapes."
}),
"dynamic": ("BOOLEAN", {
"default": False,
"tooltip": "Handle varying input shapes without recompilation. Useful for different resolutions/batch sizes."
}),
"dynamo_cache_size_limit": ("INT", {
"default": 64,
"min": 0,
"max": 1024,
"step": 1,
"tooltip": "Maximum cached compiled versions per function. Increase if using many different input shapes."
}),
"dynamo_recompile_limit": ("INT", {
"default": 128,
"min": 0,
"max": 1024,
"step": 1,
"tooltip": "Maximum recompilation attempts before fallback to eager mode. Only increase if you see recompile_limit warnings"
}),
}
}
RETURN_TYPES = ("TORCH_COMPILE_ARGS",)
FUNCTION = "create_args"
CATEGORY = "SEEDVR2"
DESCRIPTION = (
"Configure torch.compile optimization for DiT and VAE speedup. "
"Connect to DiT and/or VAE model loader. Requires PyTorch 2.0+ and Triton."
)
def create_args(self, backend: str, mode: str, fullgraph: bool, dynamic: bool,
dynamo_cache_size_limit: int, dynamo_recompile_limit: int = 128) -> Tuple[Dict[str, Any]]:
"""
Create torch.compile configuration for model optimization
Args:
backend: Compilation backend ("inductor" or "cudagraphs")
mode: Optimization mode ("default", "reduce-overhead", "max-autotune", etc.)
fullgraph: Whether to compile entire model as single graph
dynamic: Whether to handle varying input shapes without recompilation
dynamo_cache_size_limit: Maximum cached compiled versions per function
dynamo_recompile_limit: Maximum recompilation attempts before fallback
Returns:
Tuple containing torch.compile configuration dictionary
"""
compile_args = {
"backend": backend,
"mode": mode,
"fullgraph": fullgraph,
"dynamic": dynamic,
"dynamo_cache_size_limit": dynamo_cache_size_limit,
"dynamo_recompile_limit": dynamo_recompile_limit,
}
return (compile_args,)
+127
View File
@@ -0,0 +1,127 @@
"""
SeedVR2 VAE Model Loader Node
Configure VAE (Variational Autoencoder) model with tiling support
"""
from typing import Dict, Any, Tuple
from ..utils.model_registry import get_available_vae_models, DEFAULT_VAE
from ..optimization.memory_manager import get_device_list
class SeedVR2LoadVAEModel:
"""
Configure VAE (Variational Autoencoder) model loader with tiling support
Provides configuration for:
- Model selection and device placement
- Tiled encoding/decoding for VRAM reduction
- Tile size and overlap control
- Model caching between runs
- Optional torch.compile integration
Returns:
SEEDVR2_VAE configuration dictionary for main upscaler node
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
devices = get_device_list()
return {
"required": {
"model": (get_available_vae_models(), {
"default": DEFAULT_VAE,
"tooltip": "VAE model for encoding/decoding. Models will automatically download on first use. Additional models can be added to the ComfyUI models folder."
}),
"device": (devices, {
"default": devices[0],
"tooltip": "Device to use for VAE processing"
}),
},
"optional": {
"encode_tiled": ("BOOLEAN", {
"default": False,
"tooltip": "Enable tiled encoding to reduce VRAM during encoding"
}),
"encode_tile_size": ("INT", {
"default": 512,
"min": 64,
"step": 32,
"tooltip": "Size of encoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster."
}),
"encode_tile_overlap": ("INT", {
"default": 64,
"min": 0,
"step": 32,
"tooltip": "Pixel overlap between encoding tiles to reduce visible seams. Higher = better blending but slower processing."
}),
"decode_tiled": ("BOOLEAN", {
"default": False,
"tooltip": "Enable tiled decoding to reduce VRAM during decoding"
}),
"decode_tile_size": ("INT", {
"default": 512,
"min": 64,
"step": 32,
"tooltip": "Size of decoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster."
}),
"decode_tile_overlap": ("INT", {
"default": 64,
"min": 0,
"step": 32,
"tooltip": "Pixel overlap between decoding tiles to reduce visible seams. Higher = better blending but slower processing."
}),
"cache_in_ram": ("BOOLEAN", {
"default": False,
"tooltip": "Keep model in RAM between runs for faster reuse. Useful for batch processing."
}),
"torch_compile_args": ("TORCH_COMPILE_ARGS", {
"tooltip": "Optional torch.compile settings from SeedVR2 Torch Compile Settings node for speedup"
}),
}
}
RETURN_TYPES = ("SEEDVR2_VAE",)
FUNCTION = "create_config"
CATEGORY = "SEEDVR2"
DESCRIPTION = (
"Configure VAE model for SeedVR2 encoding/decoding. Supports tiled processing for VRAM reduction, "
"model caching, and torch.compile optimization. Connect output to SeedVR2 Video Upscaler."
)
def create_config(self, model: str, device: str, encode_tiled: bool = False,
encode_tile_size: int = 512, encode_tile_overlap: int = 64,
decode_tiled: bool = False, decode_tile_size: int = 512,
decode_tile_overlap: int = 64, cache_in_ram: bool = False,
torch_compile_args: Dict[str, Any] = None) -> Tuple[Dict[str, Any]]:
"""
Create VAE model configuration for SeedVR2 main node
Args:
model: Model filename to load
device: Target device for model execution
encode_tiled: Enable tiled encoding
encode_tile_size: Tile size for encoding
encode_tile_overlap: Tile overlap for encoding
decode_tiled: Enable tiled decoding
decode_tile_size: Tile size for decoding
decode_tile_overlap: Tile overlap for decoding
cache_in_ram: Whether to keep model in RAM between runs
torch_compile_args: Optional torch.compile configuration from settings node
Returns:
Tuple containing configuration dictionary for SeedVR2 main node
"""
config = {
"model": model,
"device": device,
"encode_tiled": encode_tiled,
"encode_tile_size": encode_tile_size,
"encode_tile_overlap": encode_tile_overlap,
"decode_tiled": decode_tiled,
"decode_tile_size": decode_tile_size,
"decode_tile_overlap": decode_tile_overlap,
"cache_in_ram": cache_in_ram,
"torch_compile_args": torch_compile_args,
}
return (config,)
@@ -1,22 +1,14 @@
# ComfyUI Node Interface
# Clean interface for SeedVR2 VideoUpscaler integration with ComfyUI
# Extracted from original seedvr2.py lines 1731-1812
"""
SeedVR2 Video Upscaler Node
Main ComfyUI node for high-quality video upscaling using diffusion models
"""
import os
import time
import torch
from typing import Tuple, Dict, Any, Optional
from ..utils.constants import get_base_cache_dir, get_script_directory
from ..utils.constants import get_base_cache_dir
from ..utils.downloads import download_weight
from ..utils.model_registry import (
get_available_vae_models,
get_available_dit_models,
DEFAULT_DIT,
DEFAULT_VAE
)
from ..utils.debug import Debug
from ..core.model_manager import configure_runner
from ..core.generation import (
setup_device_environment,
prepare_generation_context,
@@ -29,8 +21,7 @@ from ..core.generation import (
)
from ..optimization.memory_manager import (
cleanup_text_embeddings,
complete_cleanup,
get_device_list
complete_cleanup
)
# Import ComfyUI progress reporting
@@ -39,9 +30,8 @@ try:
except ImportError:
ProgressBar = None
script_directory = get_script_directory()
class SeedVR2:
class SeedVR2VideoUpscaler:
"""
SeedVR2 Video Upscaler ComfyUI Node
@@ -54,14 +44,13 @@ class SeedVR2:
"""
def __init__(self):
"""Initialize SeedVR2 node"""
"""Initialize SeedVR2VideoUpscaler node"""
self.runner = None
self.ctx = None
self.debug = None
self._dit_model_name = ""
self._vae_model_name = ""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
"""
@@ -133,7 +122,6 @@ class SeedVR2:
}
}
# Define return types for ComfyUI
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "SEEDVR2"
@@ -152,24 +140,8 @@ class SeedVR2:
Args:
pixels: Input video frames as tensor (N, H, W, C) in [0, 1] range
dit: DiT model configuration from SeedVR2LoadDiTModel node containing:
- model: Model filename
- device: Target device
- blocks_to_swap: BlockSwap configuration
- swap_io_components: I/O component swap flag
- cache_in_ram: Model caching flag
- torch_compile_args: Optional compilation settings
vae: VAE model configuration from SeedVR2LoadVAEModel node containing:
- model: Model filename
- device: Target device
- encode_tiled: Enable tiled encoding
- encode_tile_size: Encoding tile size
- encode_tile_overlap: Encoding tile overlap
- decode_tiled: Enable tiled decoding
- decode_tile_size: Decoding tile size
- decode_tile_overlap: Decoding tile overlap
- cache_in_ram: Model caching flag
- torch_compile_args: Optional compilation settings
dit: DiT model configuration from SeedVR2LoadDiTModel node
vae: VAE model configuration from SeedVR2LoadVAEModel node
seed: Random seed for reproducible generation
new_resolution: Target resolution for shortest edge (maintains aspect ratio)
batch_size: Frames per batch (minimum 5 for temporal consistency)
@@ -184,70 +156,63 @@ class SeedVR2:
Raises:
ValueError: If model files cannot be downloaded or configuration is invalid
RuntimeError: If generation pipeline fails
Note:
- Models are automatically downloaded on first use
- Minimum batch_size of 5 recommended for temporal consistency
- Higher batch_size improves quality but requires more VRAM
RuntimeError: If generation fails
"""
# Unpack DiT configuration
dit_model = dit["model"]
dit_device = dit["device"]
blocks_to_swap = dit.get("blocks_to_swap", 0)
swap_io_components = dit.get("swap_io_components", False)
cache_model_dit = dit.get("cache_in_ram", False)
dit_torch_compile_args = dit.get("torch_compile_args", None)
# Unpack VAE configuration
vae_model = vae["model"]
vae_device = vae["device"]
encode_tiled = vae.get("encode_tiled", False)
encode_tile_size = vae.get("encode_tile_size", 512)
encode_tile_overlap = vae.get("encode_tile_overlap", 64)
decode_tiled = vae.get("decode_tiled", False)
decode_tile_size = vae.get("decode_tile_size", 512)
decode_tile_overlap = vae.get("decode_tile_overlap", 64)
cache_model_vae = vae.get("cache_in_ram", False)
vae_torch_compile_args = vae.get("torch_compile_args", None)
# Create block_swap_config if blocks_to_swap > 0
block_swap_config = None
if blocks_to_swap > 0:
block_swap_config = {
"blocks_to_swap": blocks_to_swap,
"swap_io_components": swap_io_components,
}
# Fixed parameters (could be exposed in future)
temporal_overlap = 0
# Initialize debug
# Initialize debug instance
self.debug = Debug(enabled=enable_debug)
debug = self.debug
# Extract configuration from dict inputs
dit_model = dit["model"]
vae_model = vae["model"]
dit_device = dit["device"]
vae_device = vae["device"]
cache_model_dit = dit["cache_in_ram"]
cache_model_vae = vae["cache_in_ram"]
dit_torch_compile_args = dit.get("torch_compile_args")
vae_torch_compile_args = vae.get("torch_compile_args")
# Extract VAE tiling configuration
encode_tiled = vae["encode_tiled"]
encode_tile_size = vae["encode_tile_size"]
encode_tile_overlap = vae["encode_tile_overlap"]
decode_tiled = vae["decode_tiled"]
decode_tile_size = vae["decode_tile_size"]
decode_tile_overlap = vae["decode_tile_overlap"]
# Extract DiT BlockSwap configuration
block_swap_config = None
if dit.get("blocks_to_swap", 0) > 0 or dit.get("swap_io_components", False):
block_swap_config = {
"blocks_to_swap": dit["blocks_to_swap"],
"swap_io_components": dit["swap_io_components"],
}
# Fixed parameters
temporal_overlap = 0
# Intro logo
debug.log("", category="none", force=True)
debug.log(" ╔══════════════════════════════════════════════════════════╗", category="none", force=True)
debug.log(" ║ ███████ ███████ ███████ ██████ ██ ██ ██████ ███████ ║", category="none", force=True)
debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True)
debug.log(" ║ ███████ █████ █████ ██ ██ ██ ██ ██████ █████ ║", category="none", force=True)
debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True)
debug.log(" ║ ███████ ███████ ███████ ██████ ████ ██ ██ ███████ ║", category="none", force=True)
debug.log(" ║ © ByteDance Seed · NumZ · AInVFX ║", category="none", force=True)
debug.log(" ╚══════════════════════════════════════════════════════════╝", category="none", force=True)
debug.log("", category="none", force=True)
self.debug.log("", category="none", force=True)
self.debug.log(" ╔══════════════════════════════════════════════════════════╗", category="none", force=True)
self.debug.log(" ║ ███████ ███████ ███████ ██████ ██ ██ ██████ ███████ ║", category="none", force=True)
self.debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True)
self.debug.log(" ║ ███████ █████ █████ ██ ██ ██ ██ ██████ █████ ║", category="none", force=True)
self.debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True)
self.debug.log(" ║ ███████ ███████ ███████ ██████ ████ ██ ██ ███████ ║", category="none", force=True)
self.debug.log(" ║ © ByteDance Seed · NumZ · AInVFX ║", category="none", force=True)
self.debug.log(" ╚══════════════════════════════════════════════════════════╝", category="none", force=True)
self.debug.log("", category="none", force=True)
debug.start_timer("total_execution", force=True)
self.debug.start_timer("total_execution", force=True)
self.debug.log("━━━━━━━━━ Model Preparation ━━━━━━━━━", category="none")
debug.log("━━━━━━━━━ Model Preparation ━━━━━━━━━", category="none")
# Initial memory state
self.debug.log_memory_state("Before model preparation", show_tensors=False, detailed_tensors=False)
self.debug.start_timer("model_preparation")
debug.log_memory_state("Before model preparation", show_tensors=False, detailed_tensors=False)
debug.start_timer("model_preparation")
# Check if download succeeded
if not download_weight(dit_model=dit_model, vae_model=vae_model, debug=self.debug):
if not download_weight(dit_model=dit_model, vae_model=vae_model, debug=debug):
raise RuntimeError(
f"Failed to download required model files. "
f"DiT model: {dit_model}, VAE model: {vae_model}. "
@@ -261,7 +226,7 @@ class SeedVR2:
# RGBA detected: ComfyUI inverts alpha when loading images
# (ComfyUI mask convention: 1=hidden, 0=visible vs standard alpha: 1=opaque, 0=transparent)
# Invert alpha in-place: multiply by -1 then add 1 (equivalent to 1.0 - alpha)
self.debug.log("RGBA input detected - inverting alpha channel to match mask convention", category="info")
debug.log("RGBA input detected - inverting alpha channel to match mask convention", category="info")
pixels[..., 3].mul_(-1).add_(1)
return self._internal_execute(pixels, dit_model, vae_model, seed, new_resolution, cfg_scale,
@@ -272,45 +237,8 @@ class SeedVR2:
cache_model_dit, cache_model_vae, dit_device, vae_device,
dit_torch_compile_args, vae_torch_compile_args, block_swap_config)
except Exception as e:
self.cleanup(cache_model_dit=cache_model_dit, cache_model_vae=cache_model_vae, debug=self.debug)
self.cleanup(cache_model_dit=cache_model_dit, cache_model_vae=cache_model_vae, debug=debug)
raise e
def cleanup(self, cache_model_dit: bool = False, cache_model_vae: bool = False,
debug: Optional['Debug'] = None) -> None:
"""
Cleanup runner and free memory
Args:
cache_model_dit: If True, keep DiT model in RAM for future runs
cache_model_vae: If True, keep VAE model in RAM for future runs
debug: Optional debug instance for logging cleanup operations
"""
# Clear progress bar if it exists
if hasattr(self, '_pbar') and self._pbar is not None:
self._pbar.update_absolute(0, 100)
self._pbar = None
# Get debug from runner if not provided
if debug is None and self.runner and hasattr(self.runner, 'debug'):
debug = self.runner.debug
# Use complete_cleanup for all cleanup operations
if self.runner:
complete_cleanup(runner=self.runner, debug=debug, keep_dit_in_ram=cache_model_dit, keep_vae_in_ram=cache_model_vae)
# Delete runner only if neither model is cached
if not (cache_model_dit or cache_model_vae):
del self.runner
self.runner = None
self._dit_model_name = ""
self._vae_model_name = ""
# Clean up context text embeddings if they exist
if self.ctx:
cleanup_text_embeddings(self.ctx, debug)
self.ctx = None
def _internal_execute(
self,
@@ -392,7 +320,6 @@ class SeedVR2:
This method manages the complete lifecycle including model loading,
inference, and cleanup based on caching preferences.
"""
debug = self.debug
# Initialize progress bar
@@ -453,11 +380,11 @@ class SeedVR2:
# Transform outputs CTHW format, get dimensions from last two axes
output_h, output_w = transformed_sample.shape[-2:] # Get actual output size
debug.log(f" Total frames: {total_frames}, Input: {input_w}x{input_h}px → Output: {output_w}x{output_h}px, Batch size: {batch_size}, Channels: {channels_info}", category="generation", force=True)
debug.log(f" Total frames: {total_frames}, Input: {input_w}x{input_h}px → Output: {output_w}x{output_h}px, Batch size: {batch_size}, Channels: {channels_info}", category="generation", force=True)
# Store transform in context for reuse during encoding (avoids recreating it)
ctx['video_transform'] = video_transform
# Phase 1: Encode all batches
ctx = encode_all_batches(
self.runner,
@@ -512,7 +439,7 @@ class SeedVR2:
sample = sample.cpu()
debug.log("", category="none", force=True)
debug.log("Video upscaling completed successfully!", category="generation", force=True)
debug.log("Video upscaling completed successfully!", category="success", force=True)
debug.end_timer("generation", "Video generation")
@@ -548,8 +475,13 @@ class SeedVR2:
if total_execution_time > 0:
fps = total_frames / total_execution_time
debug.log(f"Average FPS: {fps:.2f} frames/sec", category="timing", force=True)
debug.log("━━━━━━━━━━━━━━━━━━", category="none")
# Footer with support info
debug.log("", category="none", force=True)
debug.log("━━━━━━━━━━━━━━━━━━", category="none", force=True)
debug.log("Questions? Updates? Watch the videos, star the repo & join us!", category="dialogue", force=True)
debug.log("https://www.youtube.com/@AInVFX", category="generation", force=True)
debug.log("https://github.com/numz/ComfyUI-SeedVR2_VideoUpscaler", category="star", force=True)
# Clear history for next run (do this last, after all logging)
debug.clear_history()
@@ -561,31 +493,38 @@ class SeedVR2:
return (sample,)
def _progress_callback(self, current_step: int, total_steps: int,
current_batch_frames: int, phase_name: str = "") -> None:
current_frames: int, phase_name: str) -> None:
"""
Progress callback for generation phases
Update progress bar based on pipeline phase
Args:
current_step: Current step number within the phase
total_steps: Total steps in the current phase
current_batch_frames: Number of frames in current batch
phase_name: Name of the current phase (e.g., "Phase 1: Encoding")
current_step: Current step within phase
total_steps: Total steps in phase
current_frames: Number of frames being processed
phase_name: Name of current phase
"""
if not hasattr(self, '_pbar') or self._pbar is None:
if self._pbar is None:
return
# Calculate overall progress across all phases
# Phase weights: Encode=20%, Upscale=60%, Decode=20%
phase_weights = {"Encoding": 0.2, "Upscaling": 0.6, "Decoding": 0.2}
phase_offset = {"Encoding": 0.0, "Upscaling": 0.2, "Decoding": 0.8}
# Define phase weights and offsets for overall progress
phase_weights = {
"Phase 1: Encoding": 0.25,
"Phase 2: Upscaling": 0.50,
"Phase 3: Decoding": 0.20,
"Phase 4: Post-processing": 0.05
}
# Extract the phase type from "Phase X: Type" format
if ":" in phase_name:
phase_key = phase_name.split(":")[1].strip()
else:
phase_key = "Upscaling"
phase_offset = {
"Phase 1: Encoding": 0.0,
"Phase 2: Upscaling": 0.25,
"Phase 3: Decoding": 0.75,
"Phase 4: Post-processing": 0.95
}
# Extract phase key from phase_name
phase_key = phase_name.split(" (")[0] if " (" in phase_name else phase_name
# Get weight and offset
weight = phase_weights.get(phase_key, 1.0)
offset = phase_offset.get(phase_key, 0.0)
@@ -597,6 +536,41 @@ class SeedVR2:
progress_value = int(overall_progress * 100)
self._pbar.update_absolute(progress_value, 100)
def cleanup(self, cache_model_dit: bool = False, cache_model_vae: bool = False,
debug: Optional[Debug] = None) -> None:
"""
Cleanup runner and free memory
Args:
cache_model_dit: If True, keep DiT model in RAM for future runs
cache_model_vae: If True, keep VAE model in RAM for future runs
debug: Optional debug instance for logging cleanup operations
"""
# Clear progress bar if it exists
if hasattr(self, '_pbar') and self._pbar is not None:
self._pbar.update_absolute(0, 100)
self._pbar = None
# Get debug from runner if not provided
if debug is None and self.runner and hasattr(self.runner, 'debug'):
debug = self.runner.debug
# Use complete_cleanup for all cleanup operations
if self.runner:
complete_cleanup(runner=self.runner, debug=debug, keep_dit_in_ram=cache_model_dit, keep_vae_in_ram=cache_model_vae)
# Delete runner only if neither model is cached
if not (cache_model_dit or cache_model_vae):
del self.runner
self.runner = None
self._dit_model_name = ""
self._vae_model_name = ""
# Clean up context text embeddings if they exist
if self.ctx:
cleanup_text_embeddings(self.ctx, debug)
self.ctx = None
def __del__(self):
"""Destructor"""
try:
@@ -612,341 +586,4 @@ class SeedVR2:
if hasattr(self, attr):
delattr(self, attr)
except:
pass
class SeedVR2LoadDiTModel:
"""
Configure DiT (Diffusion Transformer) model loader with memory optimization
Provides configuration for:
- Model selection and device placement
- BlockSwap memory optimization for limited VRAM
- Model caching between runs
- Optional torch.compile integration
Returns:
SEEDVR2_DIT configuration dictionary for main upscaler node
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
devices = get_device_list()
return {
"required": {
"model": (get_available_dit_models(), {
"default": DEFAULT_DIT,
"tooltip": "DiT model for upscaling. Models will automatically download on first use. Additional models can be added to the ComfyUI models folder."
}),
"device": (devices, {
"default": devices[0],
"tooltip": "Device to use for DiT upscaling"
}),
},
"optional": {
"blocks_to_swap": ("INT", {
"default": 0,
"min": 0,
"max": 36,
"step": 1,
"tooltip": "BlockSwap: Number of transformer blocks to offload (0=disabled, 16=balanced, 32=max savings for 3b model, 36=max savings for 7b model)"
}),
"swap_io_components": ("BOOLEAN", {
"default": False,
"tooltip": "Offload embeddings/IO layers to CPU for additional VRAM savings"
}),
"cache_in_ram": ("BOOLEAN", {
"default": False,
"tooltip": "Keep DiT model in RAM between runs for faster batch processing"
}),
"torch_compile_args": ("TORCH_COMPILE_ARGS", {
"tooltip": "Optional torch.compile optimization settings from SeedVR2 Torch Compile Settings node"
}),
}
}
# Define return types for ComfyUI
RETURN_TYPES = ("SEEDVR2_DIT",)
FUNCTION = "create_config"
CATEGORY = "SEEDVR2"
DESCRIPTION = "Configure DiT model loading and memory optimization settings"
def create_config(self, model: str, device: str, blocks_to_swap: int = 0,
swap_io_components: bool = False, cache_in_ram: bool = False,
torch_compile_args: Optional[Dict[str, Any]] = None) -> Tuple[Dict[str, Any]]:
"""
Create DiT model configuration dictionary
Args:
model: DiT model filename (e.g., "seedvr2_ema_3b_fp16.safetensors")
device: Target device for DiT model (e.g., "cuda:0", "cpu")
blocks_to_swap: Number of transformer blocks to offload for BlockSwap (0=disabled)
swap_io_components: Whether to offload input/output layers to CPU
cache_in_ram: Whether to keep model in RAM between runs
torch_compile_args: Optional torch.compile configuration from settings node
Returns:
Tuple containing configuration dictionary for SeedVR2 main node
"""
config = {
"model": model,
"device": device,
"blocks_to_swap": blocks_to_swap,
"swap_io_components": swap_io_components,
"cache_in_ram": cache_in_ram,
"torch_compile_args": torch_compile_args,
}
return (config,)
class SeedVR2LoadVAEModel:
"""
Configure VAE (Variational Autoencoder) model loader with tiling support
Provides configuration for:
- Model selection and device placement
- Tiled encoding/decoding for VRAM reduction
- Tile size and overlap control
- Model caching between runs
- Optional torch.compile integration
Returns:
SEEDVR2_VAE configuration dictionary for main upscaler node
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
devices = get_device_list()
return {
"required": {
"model": (get_available_vae_models(), {
"default": DEFAULT_VAE,
"tooltip": "VAE model for encoding/decoding. Models will automatically download on first use. Additional models can be added to the ComfyUI models folder."
}),
"device": (devices, {
"default": devices[0],
"tooltip": "Device to use for VAE processing"
}),
},
"optional": {
"encode_tiled": ("BOOLEAN", {
"default": False,
"tooltip": "Enable tiled encoding to reduce VRAM during encoding"
}),
"encode_tile_size": ("INT", {
"default": 512,
"min": 64,
"step": 32,
"tooltip": "Size of encoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster."
}),
"encode_tile_overlap": ("INT", {
"default": 64,
"min": 0,
"step": 32,
"tooltip": "Pixel overlap between encoding tiles to reduce visible seams. Higher = better blending but slower processing."
}),
"decode_tiled": ("BOOLEAN", {
"default": False,
"tooltip": "Enable tiled decoding to reduce VRAM during decoding"
}),
"decode_tile_size": ("INT", {
"default": 512,
"min": 64,
"step": 32,
"tooltip": "Size of decoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster."
}),
"decode_tile_overlap": ("INT", {
"default": 64,
"min": 0,
"step": 32,
"tooltip": "Pixel overlap between decoding tiles to reduce visible seams. Higher = better blending but slower processing."
}),
"cache_in_ram": ("BOOLEAN", {
"default": False,
"tooltip": "Keep VAE model in RAM between runs for faster batch processing"
}),
"torch_compile_args": ("TORCH_COMPILE_ARGS", {
"tooltip": "Optional torch.compile optimization settings from SeedVR2 Torch Compile Settings node"
}),
}
}
# Define return types for ComfyUI
RETURN_TYPES = ("SEEDVR2_VAE",)
FUNCTION = "create_config"
CATEGORY = "SEEDVR2"
DESCRIPTION = "Configure VAE settings and tiling options for memory optimization"
def create_config(self, model: str, device: str, encode_tiled: bool = False,
encode_tile_size: int = 512, encode_tile_overlap: int = 64,
decode_tiled: bool = False, decode_tile_size: int = 512,
decode_tile_overlap: int = 64, cache_in_ram: bool = False,
torch_compile_args: Optional[Dict[str, Any]] = None) -> Tuple[Dict[str, Any]]:
"""
Create VAE model configuration dictionary
Args:
model: VAE model filename (e.g., "ema_vae_fp16.safetensors")
device: Target device for VAE model (e.g., "cuda:0", "cpu")
encode_tiled: Enable tiled encoding to reduce VRAM usage
encode_tile_size: Size of encoding tiles in pixels
encode_tile_overlap: Overlap between encoding tiles to reduce seams
decode_tiled: Enable tiled decoding to reduce VRAM usage
decode_tile_size: Size of decoding tiles in pixels
decode_tile_overlap: Overlap between decoding tiles to reduce seams
cache_in_ram: Whether to keep model in RAM between runs
torch_compile_args: Optional torch.compile configuration from settings node
Returns:
Tuple containing configuration dictionary for SeedVR2 main node
"""
config = {
"model": model,
"device": device,
"encode_tiled": encode_tiled,
"encode_tile_size": encode_tile_size,
"encode_tile_overlap": encode_tile_overlap,
"decode_tiled": decode_tiled,
"decode_tile_size": decode_tile_size,
"decode_tile_overlap": decode_tile_overlap,
"cache_in_ram": cache_in_ram,
"torch_compile_args": torch_compile_args,
}
return (config,)
class SeedVR2TorchCompileSettings:
"""Configure torch.compile optimization for DiT and VAE models"""
class SeedVR2TorchCompileSettings:
"""Configure torch.compile optimization for DiT and VAE models"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
return {
"required": {
"backend": (["inductor", "cudagraphs"], {
"default": "inductor",
"tooltip": (
"Compilation backend:\n"
"• inductor: Full optimization with Triton kernel generation and fusion\n"
"• cudagraphs: Lightweight, only wraps model with CUDA graphs, no kernel optimization"
)
}),
"mode": (["default", "reduce-overhead", "max-autotune", "max-autotune-no-cudagraphs"], {
"default": "default",
"tooltip": (
"Optimization level (compilation time vs runtime speed):\n"
"• default: Fast compile, good speedup, lowest memory\n"
"• reduce-overhead: Fast compile, better speedup with CUDA graphs, +10-20% memory\n"
"• max-autotune: Slow compile, best speedup, highest memory\n"
"• max-autotune-no-cudagraphs: Like max-autotune but without CUDA graphs (more compatible)"
)
}),
"fullgraph": ("BOOLEAN", {
"default": False,
"tooltip": (
"Compile entire model as single graph:\n"
"• False: Allows graph breaks, compiles what it can, more compatible\n"
"• True: Enforces no graph breaks, errors if any found, maximum optimization but fragile"
)
}),
"dynamic": ("BOOLEAN", {
"default": False,
"tooltip": (
"Handle varying input shapes without recompilation:\n"
"• False: Specialized for exact shapes, recompiles if shape changes\n"
"• True: Creates dynamic kernels upfront, slower but handles shape variations"
)
}),
"dynamo_cache_size_limit": ("INT", {
"default": 64,
"min": 0,
"max": 1024,
"step": 1,
"tooltip": (
"Max cached compiled versions per function (prevents memory bloat):\n"
"Higher value = more variations cached = more memory used\n"
"Lower value = more recompilation if inputs vary\n"
"Default 64 is good for most cases. Increase if you see cache_size_limit warnings"
)
}),
},
"optional": {
"dynamo_recompile_limit": ("INT", {
"default": 128,
"min": 0,
"max": 1024,
"step": 1,
"tooltip": (
"Max recompilation attempts before giving up (falls back to no compilation):\n"
"Safety limit to prevent infinite recompilation loops\n"
"Default 128 is sufficient. Only increase if you see recompile_limit warnings"
)
}),
}
}
# Define return types for ComfyUI
RETURN_TYPES = ("TORCH_COMPILE_ARGS",)
FUNCTION = "create_args"
CATEGORY = "SEEDVR2"
DESCRIPTION = (
"Configure torch.compile optimization for DiT and VAE speedup. "
"Connect to DiT and/or VAE model loader. Requires PyTorch 2.0+ and Triton."
)
def create_args(self, backend: str, mode: str, fullgraph: bool, dynamic: bool,
dynamo_cache_size_limit: int, dynamo_recompile_limit: int = 128) -> Tuple[Dict[str, Any]]:
"""
Create torch.compile configuration for model optimization
Args:
backend: Compilation backend ("inductor" or "cudagraphs")
mode: Optimization mode ("default", "reduce-overhead", "max-autotune", etc.)
fullgraph: Whether to compile entire model as single graph
dynamic: Whether to handle varying input shapes without recompilation
dynamo_cache_size_limit: Maximum cached compiled versions per function
dynamo_recompile_limit: Maximum recompilation attempts before fallback
Returns:
Tuple containing torch.compile configuration dictionary
"""
compile_args = {
"backend": backend,
"mode": mode,
"fullgraph": fullgraph,
"dynamic": dynamic,
"dynamo_cache_size_limit": dynamo_cache_size_limit,
"dynamo_recompile_limit": dynamo_recompile_limit,
}
return (compile_args,)
# ComfyUI Node Mappings
NODE_CLASS_MAPPINGS = {
"SeedVR2": SeedVR2,
"SeedVR2LoadDiTModel": SeedVR2LoadDiTModel,
"SeedVR2LoadVAEModel": SeedVR2LoadVAEModel,
"SeedVR2TorchCompileSettings": SeedVR2TorchCompileSettings,
}
# Human-readable node display names
NODE_DISPLAY_NAME_MAPPINGS = {
"SeedVR2": "SeedVR2 Video Upscaler",
"SeedVR2LoadDiTModel": "SeedVR2 (Down)Load DiT Model",
"SeedVR2LoadVAEModel": "SeedVR2 (Down)Load VAE Model",
"SeedVR2TorchCompileSettings": "SeedVR2 Torch Compile Settings",
}
# Export version and metadata
__version__ = "2.0.0-modular"
__author__ = "SeedVR2 Team"
__description__ = "High-quality video upscaling using advanced diffusion models"
# Additional exports for introspection
__all__ = [
'SeedVR2',
'SeedVR2LoadDiTModel',
'SeedVR2LoadVAEModel',
'NODE_CLASS_MAPPINGS',
'NODE_DISPLAY_NAME_MAPPINGS'
]
pass
+25 -20
View File
@@ -79,7 +79,10 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
return
blocks_to_swap = block_swap_config.get("blocks_to_swap", 0)
if blocks_to_swap <= 0:
swap_io_components = block_swap_config.get("swap_io_components", False)
# Early return only if both block swap and I/O swap are disabled
if blocks_to_swap <= 0 and not swap_io_components:
return
if debug is None:
@@ -100,48 +103,50 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) ->
offload_device = "cpu"
configs = []
blocks_to_swap = block_swap_config.get("blocks_to_swap", 0)
if blocks_to_swap > 0:
block_text = "block" if blocks_to_swap == 1 else "blocks"
block_text = "block" if blocks_to_swap <= 1 else "blocks"
configs.append(f"{blocks_to_swap} {block_text}")
if block_swap_config.get("swap_io_components", False):
if swap_io_components:
configs.append("I/O components")
debug.log(f"BlockSwap configured: {', '.join(configs)}", category="blockswap", force=True)
debug.log("BlockSwap will swap blocks to GPU during inference", category="info", force=True)
# Validate model structure
# Validate model structure for block operations
if not hasattr(model, "blocks"):
debug.log("Model doesn't have 'blocks' attribute for BlockSwap", level="ERROR", category="blockswap")
return
total_blocks = len(model.blocks)
debug.log(f"Model has {total_blocks} blocks total", category="blockswap")
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
if blocks_to_swap > 0:
blocks_to_swap = min(blocks_to_swap, total_blocks)
model.blocks_to_swap = blocks_to_swap - 1 # Convert to 0-indexed
debug.log(f"Configuring: {blocks_to_swap}/{total_blocks} blocks for swapping", category="blockswap")
debug.log("BlockSwap will swap blocks to GPU during inference", category="info", force=True)
else:
# No block swapping, set to -1 so no blocks match the swap condition
model.blocks_to_swap = -1
debug.log("Block swapping disabled (blocks_to_swap=0)", category="blockswap")
model.main_device = device
model.offload_device = offload_device
debug.log(f"Configuring: {blocks_to_swap}/{total_blocks} blocks for swapping", category="blockswap")
# Configure I/O components
swap_io_components = block_swap_config.get("swap_io_components", False)
io_components_offloaded = _configure_io_components(model, device, offload_device,
swap_io_components, debug)
# Configure block placement and memory tracking
memory_stats = _configure_blocks(model, device, offload_device, debug)
memory_stats['io_components'] = io_components_offloaded
# Log memory summary
# Log memory summary
_log_memory_summary(memory_stats, offload_device, device, swap_io_components,
debug)
# 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, debug)
# Wrap block forward methods for dynamic swapping (only if blocks_to_swap > 0)
if blocks_to_swap > 0:
for b, block in enumerate(model.blocks):
if b <= model.blocks_to_swap:
_wrap_block_forward(block, b, model, debug)
# Patch RoPE modules for robust error handling
_patch_rope_for_blockswap(model, debug)
+11 -9
View File
@@ -31,20 +31,20 @@ class Debug:
# Icon mapping for different categories
CATEGORY_ICONS = {
"general": "🔄", # General operations/processing
"timing": "⚡", # Performance timing
"timing": "⚡", # Performance timing
"memory": "📊", # Memory usage tracking
"cache": "💾", # Cache operations
"cleanup": "🧹", # Cleanup operations
"setup": "🔧", # Configuration/setup
"generation": "🎬", # Generation process
"dit": "🚀", # Model loading/operations
"dit": "🚀", # Model loading/operations
"blockswap": "🔀", # BlockSwap operations
"download": "📥", # Download operations
"success": "✅", # Successful completion
"warning": "⚠️", # Warnings
"error": "❌", # Errors
"info": "ℹ️", # Statistics/info
"tip" :"💡", # Tip/suggestion
"tip" :"💡", # Tip/suggestion
"video": "📹", # Video/sequence info
"reuse": "♻️", # Reusing/recycling
"runner": "🏃", # Runner operations
@@ -53,6 +53,8 @@ class Debug:
"precision": "🎯", # Precision
"device": "🖥️", # Device info
"file": "📂", # File operations
"star": "⭐", # Star
"dialogue": "💬", # Dialogue
"none" : "",
}
@@ -220,7 +222,7 @@ class Debug:
grandchild_duration = self.timer_durations.get(grandchild, 0)
if grandchild_duration >= 0.01: # Only show if >= 10ms
grandchild_message = self.timer_messages.get(grandchild, grandchild)
self.log(f" └─ {grandchild_message}: {grandchild_duration:.2f}s", category="timing", force=force)
self.log(f" └─ {grandchild_message}: {grandchild_duration:.2f}s", category="timing", force=force)
if unaccounted > 0.01: # Show if more than 10ms unaccounted
self.log(f" └─ (other operations): {unaccounted:.2f}s", category="timing", force=force)
@@ -411,18 +413,18 @@ class Debug:
# Show top 5 largest
largest = sorted(details['gpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:5]
for t in largest:
self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force)
self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force)
# Large CPU tensors
if details['large_cpu_tensors']:
cpu_large_gb = sum(t['size_mb'] for t in details['large_cpu_tensors']) / 1024
self.log(f" Large CPU tensors (>10MB):", category="memory", force=force)
self.log(f" {len(details['large_cpu_tensors'])} using {cpu_large_gb:.2f}GB", category="memory", force=force)
self.log(f" {len(details['large_cpu_tensors'])} using {cpu_large_gb:.2f}GB", category="memory", force=force)
# Show top 3 largest
largest = sorted(details['large_cpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:3]
for t in largest:
self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force)
self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force)
# Common shape patterns
if details['shape_patterns']:
@@ -432,7 +434,7 @@ class Debug:
self.log(" Common tensor shapes:", category="memory", force=force)
for shape, count in common_shapes:
if count > 1:
self.log(f" {shape}: {count} instances", category="memory", force=force)
self.log(f" {shape}: {count} instances", category="memory", force=force)
# Module instances
if details['module_types']:
@@ -440,7 +442,7 @@ class Debug:
if multi_instance:
self.log(" Multiple module instances:", category="memory", force=force)
for mtype, count in sorted(multi_instance, key=lambda x: x[1], reverse=True)[:5]:
self.log(f" {mtype}: {count} instances", category="memory", force=force)
self.log(f" {mtype}: {count} instances", category="memory", force=force)
def _log_memory_diff(self, current_metrics: Dict[str, Any], force: bool = False) -> None: