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:
+1
-1
@@ -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"]
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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'
|
||||
]
|
||||
@@ -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,)
|
||||
@@ -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,)
|
||||
@@ -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
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user