feat: Add CLI model caching for multi-file processing and unify device handling

- Add --cache_dit and --cache_vae flags for efficient multi-file directory processing
- Refactor processing pipeline to eliminate duplication between worker and direct modes
- Implement platform-agnostic device management (CUDA/MPS/CPU)
- Unify parameter naming: res_w→resolution, max_res_w→max_resolution across codebase
- Add smart offload device defaults when caching enabled
- Improve validation and user feedback for cache + multi-GPU scenarios
This commit is contained in:
Adrien Toupet
2025-11-04 15:12:28 -05:00
parent ad020d3803
commit 32a049dfd9
5 changed files with 380 additions and 239 deletions
+2 -2
View File
@@ -269,7 +269,7 @@ options:
-h, --help show this help message and exit
--video_path VIDEO_PATH Path to input video file
--seed SEED Random seed for generation (default: 100)
--resolution RESOLUTION Target resolution width (default: 1072)
--resolution RESOLUTION Target resolution width (default: 1080)
--batch_size BATCH_SIZE Number of frames per batch (default: 5)
--model Model to use (default: 3B FP8) in list:
seedvr2_ema_3b_fp16.safetensors,
@@ -291,7 +291,7 @@ Examples :
```
# Upscale 18 frames as png
python inference_cli.py --video_path "MAIN.mp4" --resolution 1072 --batch_size 9 --model seedvr2_ema_3b_fp8_e4m3fn.safetensors --model_dir ./models\SEEDVR2 --load_cap 18 --output "C:\Users\Emmanuel\Downloads\test_upscale" --output_format png --preserve_vram
python inference_cli.py --video_path "MAIN.mp4" --resolution 1080 --batch_size 9 --model seedvr2_ema_3b_fp8_e4m3fn.safetensors --model_dir ./models\SEEDVR2 --load_cap 18 --output "C:\Users\Emmanuel\Downloads\test_upscale" --output_format png --preserve_vram
# Upscale 1000 frames on 4 GPU, each GPU will receive 250 frames and will process them 50 by 50
python inference_cli.py --video_path "MAIN.mp4" --batch_size 50 --load_cap 1000 --output ".\outputs\test_upscale.mp4" --cuda_device 0,1,2,3
+356 -215
View File
@@ -72,6 +72,72 @@ from src.utils.model_registry import get_available_dit_models, DEFAULT_DIT, DEFA
from src.utils.constants import SEEDVR2_FOLDER_NAME
debug = Debug(enabled=False) # Default to disabled, can be enabled via CLI
# =============================================================================
# Device Management Helpers
# =============================================================================
def _get_platform_type() -> str:
"""Determine the platform device type (cuda/mps/cpu)."""
if platform.system() == "Darwin":
return "mps"
elif torch.cuda.is_available():
return "cuda"
else:
return "cpu"
def _device_id_to_name(device_id: str, platform_type: str = None) -> str:
"""
Convert device ID to full device name.
Args:
device_id: Device ID ("0", "1") or special value ("cpu", "none")
platform_type: Override platform type ("cuda", "mps", "cpu")
Returns:
Full device name ("cuda:0", "mps:0", "cpu", "none")
"""
if device_id in ("cpu", "none"):
return device_id
if platform_type is None:
platform_type = _get_platform_type()
# MPS typically doesn't use indices
if platform_type == "mps":
return "mps"
return f"{platform_type}:{device_id}"
def _parse_offload_device(offload_arg: str, platform_type: str = None, cache_enabled: bool = False) -> Optional[str]:
"""
Parse offload device argument to full device name.
Args:
offload_arg: Offload device argument ("none", "cpu", "0", "1", or "cuda:1")
platform_type: Override platform type
cache_enabled: If True and offload_arg is "none", default to "cpu"
Returns:
Full device name or None
"""
if offload_arg == "none":
# If caching enabled but no offload device specified, default to CPU
return "cpu" if cache_enabled else None
if offload_arg == "cpu":
return "cpu"
# If already a full device name (cuda:1, mps:0), return as-is
if ":" in offload_arg:
return offload_arg
# Otherwise treat as device ID
return _device_id_to_name(offload_arg, platform_type)
# =============================================================================
# Video I/O Functions
# =============================================================================
@@ -164,7 +230,8 @@ def generate_output_path(input_path: str, output_format: str, output_dir: Option
return f"output/{input_name}_upscaled.mp4"
def process_single_file(input_path: str, args: argparse.Namespace, device_list: List[str],
output_path: Optional[str] = None, format_auto_detected: bool = False) -> int:
output_path: Optional[str] = None, format_auto_detected: bool = False,
runner_cache: Optional[Dict[str, Any]] = None) -> int:
"""
Process a single video or image file.
@@ -208,7 +275,13 @@ def process_single_file(input_path: str, args: argparse.Namespace, device_list:
# Process frames
processing_start = time.time()
result = _gpu_processing(frames_tensor, device_list, args)
# Use direct processing if caching enabled
if runner_cache is not None:
# Direct single-GPU processing with model caching
result = _single_gpu_direct_processing(frames_tensor, args, device_list[0], runner_cache)
else:
# Multi-GPU or non-cached processing via worker processes
result = _gpu_processing(frames_tensor, device_list, args)
debug.log(f"Processing time: {time.time() - processing_start:.2f}s", category="timing")
# Save results
@@ -436,10 +509,180 @@ def save_frames_to_png(
debug.log(f"PNG saving completed: {total} files in '{output_dir}'", category="success")
# =============================================================================
# Multi-GPU Processing Functions
# Core Processing Logic
# =============================================================================
def _process_frames_core(
frames_tensor: torch.Tensor,
args: argparse.Namespace,
device_id: str,
debug: 'Debug',
runner_cache: Optional[Dict[str, Any]] = None
) -> torch.Tensor:
"""
Core frame processing logic shared between worker and direct processing.
Executes the complete 4-phase pipeline: encode → upscale → decode → postprocess.
Supports both cached (direct) and non-cached (worker) execution modes.
Args:
frames_tensor: Input frames [T, H, W, C], Float16/Float32, range [0,1]
args: Command-line arguments with all processing settings
device_id: Device ID for inference ("0", "1", etc.)
debug: Debug instance for logging
runner_cache: Optional cache dict for model reuse (direct mode only)
Returns:
Upscaled frames tensor [T', H', W', C], Float32, range [0,1]
"""
from src.core.generation_utils import setup_generation_context, prepare_runner
from src.core.generation_phases import (
encode_all_batches, upscale_all_batches, decode_all_batches, postprocess_all_batches
)
# Determine platform and convert device IDs to full names
platform_type = _get_platform_type()
inference_device = _device_id_to_name(device_id, platform_type)
# Parse offload devices (with caching defaults)
cache_dit = args.cache_dit if runner_cache is not None else False
cache_vae = args.cache_vae if runner_cache is not None else False
dit_offload = _parse_offload_device(args.dit_offload_device, platform_type, cache_dit)
vae_offload = _parse_offload_device(args.vae_offload_device, platform_type, cache_vae)
tensor_offload = _parse_offload_device(args.tensor_offload_device, platform_type, False)
# Setup or reuse generation context
if runner_cache is not None and 'ctx' in runner_cache:
ctx = runner_cache['ctx']
# Clear previous run data but keep device config
keys_to_keep = {'dit_device', 'vae_device', 'dit_offload_device',
'vae_offload_device', 'tensor_offload_device', 'compute_dtype'}
for key in list(ctx.keys()):
if key not in keys_to_keep:
del ctx[key]
else:
ctx = setup_generation_context(
dit_device=inference_device,
vae_device=inference_device,
dit_offload_device=dit_offload,
vae_offload_device=vae_offload,
tensor_offload_device=tensor_offload,
debug=debug
)
if runner_cache is not None:
runner_cache['ctx'] = ctx
# Build torch compile args
torch_compile_args_dit = None
torch_compile_args_vae = None
if args.compile_dit:
torch_compile_args_dit = {
"backend": args.compile_backend,
"mode": args.compile_mode,
"fullgraph": args.compile_fullgraph,
"dynamic": args.compile_dynamic,
"dynamo_cache_size_limit": args.compile_dynamo_cache_size_limit,
"dynamo_recompile_limit": args.compile_dynamo_recompile_limit,
}
if args.compile_vae:
torch_compile_args_vae = {
"backend": args.compile_backend,
"mode": args.compile_mode,
"fullgraph": args.compile_fullgraph,
"dynamic": args.compile_dynamic,
"dynamo_cache_size_limit": args.compile_dynamo_cache_size_limit,
"dynamo_recompile_limit": args.compile_dynamo_recompile_limit,
}
# Prepare runner with caching support
model_dir = args.model_dir if args.model_dir is not None else f"./models/{SEEDVR2_FOLDER_NAME}"
# Use fixed IDs for CLI caching when enabled
dit_id = "cli_dit" if cache_dit else None
vae_id = "cli_vae" if cache_vae else None
runner, cache_context = prepare_runner(
dit_model=args.model,
vae_model=DEFAULT_VAE,
model_dir=model_dir,
debug=debug,
ctx=ctx,
dit_cache=cache_dit,
vae_cache=cache_vae,
dit_id=dit_id,
vae_id=vae_id,
block_swap_config={
'blocks_to_swap': args.blocks_to_swap,
'swap_io_components': args.swap_io_components,
'offload_device': dit_offload,
},
encode_tiled=args.vae_encode_tiling_enabled,
encode_tile_size=args.vae_encode_tile_size,
encode_tile_overlap=args.vae_encode_tile_overlap,
decode_tiled=args.vae_decode_tiling_enabled,
decode_tile_size=args.vae_decode_tile_size,
decode_tile_overlap=args.vae_decode_tile_overlap,
tile_debug=args.tile_debug.lower() if args.tile_debug else "false",
attention_mode=args.attention_mode,
torch_compile_args_dit=torch_compile_args_dit,
torch_compile_args_vae=torch_compile_args_vae
)
ctx['cache_context'] = cache_context
if runner_cache is not None:
runner_cache['runner'] = runner
# Phase 1: Encode
ctx = encode_all_batches(
runner, ctx=ctx, images=frames_tensor,
debug=debug,
batch_size=args.batch_size,
seed=args.seed,
progress_callback=None,
temporal_overlap=args.temporal_overlap,
resolution=args.resolution,
max_resolution=args.max_resolution,
input_noise_scale=args.input_noise_scale,
color_correction=args.color_correction
)
# Phase 2: Upscale
ctx = upscale_all_batches(
runner, ctx=ctx, debug=debug, progress_callback=None,
seed=args.seed,
latent_noise_scale=args.latent_noise_scale,
cache_model=cache_dit
)
# Phase 3: Decode
ctx = decode_all_batches(
runner, ctx=ctx, debug=debug, progress_callback=None,
cache_model=cache_vae
)
# Phase 4: Post-process
ctx = postprocess_all_batches(
ctx=ctx, debug=debug, progress_callback=None,
color_correction=args.color_correction,
prepend_frames=0, # Worker mode handles this in main process
temporal_overlap=args.temporal_overlap,
batch_size=args.batch_size
)
result_tensor = ctx['final_video']
# Convert to CPU and compatible dtype
if result_tensor.is_cuda or result_tensor.is_mps:
result_tensor = result_tensor.cpu()
if result_tensor.dtype in (torch.bfloat16, torch.float8_e4m3fn, torch.float8_e5m2):
result_tensor = result_tensor.to(torch.float32)
return result_tensor
def _worker_process(
proc_idx: int,
device_id: int,
@@ -448,180 +691,60 @@ def _worker_process(
return_queue: mp.Queue
) -> None:
"""
Worker process for multi-GPU upscaling of video frame chunks.
Worker process for multi-GPU upscaling.
Each worker runs in its own process with dedicated GPU, executing the full
4-phase upscaling pipeline (encode → upscale → decode → postprocess).
Results are returned via multiprocessing queue as numpy arrays.
This function is spawned as a separate process and performs local imports
to avoid CUDA initialization conflicts in the main process.
Args:
proc_idx: Worker process index for result tracking
device_id: CUDA device ID to use for this worker
frames_np: Numpy array of frames [T, H, W, C], Float32, range [0,1]
shared_args: Dictionary containing all configuration parameters including
model paths, processing settings, and optimization flags
return_queue: Multiprocessing queue for returning results to main process
Note:
- Sets CUDA_VISIBLE_DEVICES to isolate GPU access per worker
- Prepend frame removal handled in main process (multi-GPU safe)
- No model caching in CLI mode (single-run workflow)
- BlockSwap offloading handled via dit_offload_device (blocks/IO → CPU during inference)
Sets up isolated CUDA environment and calls core processing logic.
Results returned via multiprocessing queue as numpy arrays.
"""
if platform.system() != "Darwin":
# Limit CUDA visibility to the chosen GPU BEFORE importing torch-heavy deps
# Limit CUDA visibility to the chosen GPU BEFORE importing torch
os.environ["CUDA_VISIBLE_DEVICES"] = str(device_id)
# Keep same cudaMallocAsync setting
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
import torch # local import inside subprocess
from src.core.generation_utils import (
setup_generation_context, prepare_runner
)
from src.core.generation_phases import (
encode_all_batches, upscale_all_batches, decode_all_batches, postprocess_all_batches
)
import torch
# Create debug instance for this worker process
# Create debug instance for this worker
worker_debug = Debug(enabled=shared_args["debug"])
# Prepare offload device arguments
dit_offload = None if shared_args["dit_offload_device"] == "none" else shared_args["dit_offload_device"]
vae_offload = None if shared_args["vae_offload_device"] == "none" else shared_args["vae_offload_device"]
tensor_offload = None if shared_args["tensor_offload_device"] == "none" else shared_args["tensor_offload_device"]
# Convert numpy back to tensor
frames_tensor = torch.from_numpy(frames_np).to(torch.float16)
# Setup generation context with device configuration
ctx = setup_generation_context(
dit_device=f"cuda:{device_id}",
vae_device=f"cuda:{device_id}",
dit_offload_device=dit_offload,
vae_offload_device=vae_offload,
tensor_offload_device=tensor_offload,
debug=worker_debug
)
# Reconstruct frames tensor using compute dtype from context
frames_tensor = torch.from_numpy(frames_np).to(ctx['compute_dtype'])
# Create torch compile args if enabled
torch_compile_args_dit = None
torch_compile_args_vae = None
if shared_args.get("compile_dit", False):
torch_compile_args_dit = {
"backend": shared_args.get("compile_backend", "inductor"),
"mode": shared_args.get("compile_mode", "default"),
"fullgraph": shared_args.get("compile_fullgraph", False),
"dynamic": shared_args.get("compile_dynamic", False),
"dynamo_cache_size_limit": shared_args.get("compile_dynamo_cache_size_limit", 64),
"dynamo_recompile_limit": shared_args.get("compile_dynamo_recompile_limit", 128),
}
if shared_args.get("compile_vae", False):
torch_compile_args_vae = {
"backend": shared_args.get("compile_backend", "inductor"),
"mode": shared_args.get("compile_mode", "default"),
"fullgraph": shared_args.get("compile_fullgraph", False),
"dynamic": shared_args.get("compile_dynamic", False),
"dynamo_cache_size_limit": shared_args.get("compile_dynamo_cache_size_limit", 64),
"dynamo_recompile_limit": shared_args.get("compile_dynamo_recompile_limit", 128),
}
# Prepare runner
model_dir = shared_args["model_dir"]
model_name = shared_args["model"]
runner, cache_context = prepare_runner(
dit_model=model_name,
vae_model=DEFAULT_VAE,
model_dir=model_dir,
# Create args namespace from shared_args
args = argparse.Namespace(**shared_args)
# Process frames (no caching in worker mode)
result_tensor = _process_frames_core(
frames_tensor=frames_tensor,
args=args,
device_id="0", # Always "0" in worker (CUDA_VISIBLE_DEVICES set)
debug=worker_debug,
ctx=ctx,
dit_cache=False, # No caching in CLI
vae_cache=False, # No caching in CLI
dit_id=None, # No caching in CLI
vae_id=None, # No caching in CLI
block_swap_config=shared_args["block_swap_config"],
encode_tiled=shared_args["vae_encode_tiling_enabled"],
encode_tile_size=shared_args["vae_encode_tile_size"],
encode_tile_overlap=shared_args["vae_encode_tile_overlap"],
decode_tiled=shared_args["vae_decode_tiling_enabled"],
decode_tile_size=shared_args["vae_decode_tile_size"],
decode_tile_overlap=shared_args["vae_decode_tile_overlap"],
tile_debug=shared_args.get("tile_debug", "false"),
attention_mode=shared_args["attention_mode"],
torch_compile_args_dit=torch_compile_args_dit,
torch_compile_args_vae=torch_compile_args_vae
)
# Store cache context in ctx for use in generation phases
ctx['cache_context'] = cache_context
# Phase 1: Encode all batches
ctx = encode_all_batches(
runner,
ctx=ctx,
images=frames_tensor,
debug=worker_debug,
batch_size=shared_args["batch_size"],
seed=shared_args["seed"],
progress_callback=None,
temporal_overlap=shared_args["temporal_overlap"],
res_w=shared_args["res_w"],
max_res_w=shared_args["max_resolution"],
input_noise_scale=shared_args["input_noise_scale"],
color_correction=shared_args.get("color_correction", "lab")
)
# Phase 2: Upscale all batches
ctx = upscale_all_batches(
runner,
ctx=ctx,
debug=worker_debug,
progress_callback=None,
seed=shared_args["seed"],
latent_noise_scale=shared_args["latent_noise_scale"],
cache_model=False # No caching in CLI
)
# Phase 3: Decode all batches
ctx = decode_all_batches(
runner,
ctx=ctx,
debug=worker_debug,
progress_callback=None,
cache_model=False # No caching in CLI
runner_cache=None # No caching in multiprocessing mode
)
# Phase 4: Post-processing and final assembly
ctx = postprocess_all_batches(
ctx=ctx,
debug=worker_debug,
progress_callback=None,
color_correction=shared_args.get("color_correction", "lab"),
prepend_frames=0, # Never remove prepend_frames in workers (multi-GPU safe)
temporal_overlap=shared_args["temporal_overlap"],
batch_size=shared_args["batch_size"]
)
# Get final result
result_tensor = ctx['final_video']
# Ensure result is on CPU
if result_tensor.is_cuda or result_tensor.is_mps:
result_tensor = result_tensor.cpu()
# Convert ML-specific dtypes that NumPy doesn't support to float32
if result_tensor.dtype in (torch.bfloat16, torch.float8_e4m3fn, torch.float8_e5m2):
result_tensor = result_tensor.to(torch.float32)
# Send back result as numpy array
return_queue.put((proc_idx, result_tensor.numpy()))
def _single_gpu_direct_processing(
frames_tensor: torch.Tensor,
args: argparse.Namespace,
device_id: str,
runner_cache: Dict[str, Any]
) -> torch.Tensor:
"""
Direct single-GPU processing with model caching support.
Uses main process and shared runner cache for efficient multi-file processing.
"""
return _process_frames_core(
frames_tensor=frames_tensor,
args=args,
device_id=device_id,
debug=debug,
runner_cache=runner_cache
)
def _gpu_processing(
frames_tensor: torch.Tensor,
device_list: List[str],
@@ -643,7 +766,7 @@ def _gpu_processing(
Args:
frames_tensor: Input frames [T, H, W, C], Float32, range [0,1]
device_list: List of CUDA device IDs as strings (e.g., ["0", "1"])
device_list: List of device IDs as strings (e.g., ["0", "1"])
args: Parsed command-line arguments containing all processing settings
Returns:
@@ -651,7 +774,7 @@ def _gpu_processing(
where T' may be less than T if prepend_frames were removed
Note:
- Single GPU: Simple sequential processing without overlap
- Single GPU: Can use multiprocessing or direct processing
- Multi-GPU with overlap: Chunks sized to multiples of batch_size for
proper temporal blending
- Prepended frames removed after all GPU workers complete (multi-GPU safe)
@@ -682,43 +805,8 @@ def _gpu_processing(
return_queue = mp.Queue(maxsize=0) # 0 = unlimited (explicit)
workers = []
shared_args = {
"model": args.model,
"model_dir": args.model_dir if args.model_dir is not None else f"./models/{SEEDVR2_FOLDER_NAME}",
"color_correction": args.color_correction,
"input_noise_scale": args.input_noise_scale,
"latent_noise_scale": args.latent_noise_scale,
"debug": args.debug,
"seed": args.seed,
"res_w": args.resolution,
"max_resolution": args.max_resolution,
"batch_size": args.batch_size,
"temporal_overlap": args.temporal_overlap,
"block_swap_config": {
'blocks_to_swap': args.blocks_to_swap,
'swap_io_components': args.swap_io_components,
'offload_device': args.dit_offload_device if args.dit_offload_device != "none" else None,
},
"vae_encode_tiling_enabled": args.vae_encode_tiling_enabled,
"vae_encode_tile_size": args.vae_encode_tile_size,
"vae_encode_tile_overlap": args.vae_encode_tile_overlap,
"vae_decode_tiling_enabled": args.vae_decode_tiling_enabled,
"vae_decode_tile_size": args.vae_decode_tile_size,
"vae_decode_tile_overlap": args.vae_decode_tile_overlap,
"tile_debug": args.tile_debug.lower() if args.tile_debug else "false",
"dit_offload_device": args.dit_offload_device,
"vae_offload_device": args.vae_offload_device,
"tensor_offload_device": args.tensor_offload_device,
"attention_mode": args.attention_mode,
"compile_dit": args.compile_dit,
"compile_vae": args.compile_vae,
"compile_backend": args.compile_backend,
"compile_mode": args.compile_mode,
"compile_fullgraph": args.compile_fullgraph,
"compile_dynamic": args.compile_dynamic,
"compile_dynamo_cache_size_limit": args.compile_dynamo_cache_size_limit,
"compile_dynamo_recompile_limit": args.compile_dynamo_recompile_limit,
}
# Convert args namespace to dict for serialization
shared_args = vars(args).copy()
# Start all workers
for idx, (device_id, chunk_tensor) in enumerate(zip(device_list, chunks)):
@@ -778,10 +866,10 @@ def _gpu_processing(
if result_tensor is None:
result_tensor = torch.from_numpy(results_np[0]).to(torch.float32)
else:
# Single GPU or no overlap: simple concatenation
# Simple concatenation without overlap
result_tensor = torch.from_numpy(np.concatenate(results_np, axis=0)).to(torch.float32)
# Remove prepended frames from final concatenated result (multi-GPU safe)
# Handle prepend_frames removal (multi-GPU safe - done after all workers complete)
if args.prepend_frames > 0:
if args.prepend_frames < result_tensor.shape[0]:
debug.log(f"Removing {args.prepend_frames} prepended frames from output", category="generation")
@@ -902,21 +990,24 @@ def parse_arguments() -> argparse.Namespace:
"Requires --dit_offload_device to be set. Can be used alone or with --blocks_to_swap.")
parser.add_argument("--dit_offload_device", type=str, default="none",
help="Device to offload DiT model when not in use (default: none). "
"Options: 'none' (keep on inference device), 'cpu' (offload to RAM), "
"or any GPU device like 'cuda:1' (offload to another GPU). "
"Options: 'none' (keep on inference device), 'cpu' (offload to RAM/system memory), "
"or GPU device ID like '1' (offload to another GPU). "
"Required when BlockSwap is enabled (blocks_to_swap > 0 or swap_io_components = True). "
"Multi-GPU example: inference on cuda:0, offload to cuda:1 for memory distribution.")
"Multi-GPU example: --cuda_device 0 --dit_offload_device 1 distributes memory across GPUs.")
parser.add_argument("--vae_offload_device", type=str, default="none",
help="Device to offload VAE when not in use (default: none). "
"Options: 'none' (keep on inference device), 'cpu' (offload to RAM), "
"or any GPU device like 'cuda:1' (offload to another GPU). "
"Use 'cpu' or another GPU to free VRAM between encode/decode phases (slower but saves VRAM on inference device).")
"Options: 'none' (keep on inference device), 'cpu' (offload to RAM/system memory), "
"or GPU device ID like '1' (offload to another GPU). "
"Use 'cpu' or another GPU to free VRAM between encode/decode phases.")
parser.add_argument("--tensor_offload_device", type=str, default="cpu",
help="Device to offload intermediate tensors between phases (default: cpu). "
"Options: 'cpu' (offload to RAM - recommended), 'none' (keep on inference device), "
"or any GPU device like 'cuda:1' (offload to another GPU). "
"Use 'cpu' to prevent VRAM accumulation for long videos. "
"Use another GPU for faster offloading while distributing memory load.")
"Options: 'cpu' (offload to RAM/system memory - recommended), 'none' (keep on inference device), "
"or GPU device ID like '1' (offload to another GPU). "
"Use 'cpu' to prevent VRAM accumulation for long videos.")
parser.add_argument("--cache_dit", action="store_true",
help="Cache DiT model between files for faster multi-file processing (single GPU only). Model cached on device specified by --dit_offload_device (default: cpu)")
parser.add_argument("--cache_vae", action="store_true",
help="Cache VAE model between files for faster multi-file processing (single GPU only). Model cached on device specified by --vae_offload_device (default: cpu)")
parser.add_argument("--vae_encode_tiling_enabled", action="store_true",
help="Enable VAE encode tiling for VRAM reduction during encoding. Disabled by default.")
parser.add_argument("--vae_encode_tile_size", action=OneOrTwoValues, nargs='+', default=(1024, 1024),
@@ -1014,6 +1105,23 @@ def main() -> None:
)
sys.exit(1)
# Inform about caching defaults
if args.cache_dit and args.dit_offload_device == "none":
offload_target = "system memory (CPU)" if _get_platform_type() != "mps" else "unified memory"
debug.log(
f"DiT caching enabled: Using default {offload_target} for offload. "
"Set --dit_offload_device explicitly to use a different device.",
category="cache", force=True
)
if args.cache_vae and args.vae_offload_device == "none":
offload_target = "system memory (CPU)" if _get_platform_type() != "mps" else "unified memory"
debug.log(
f"VAE caching enabled: Using default {offload_target} for offload. "
"Set --vae_offload_device explicitly to use a different device.",
category="cache", force=True
)
if args.debug:
if platform.system() == "Darwin":
debug.log("You are running on macOS and will use the MPS backend!", category="info", force=True)
@@ -1061,6 +1169,19 @@ def main() -> None:
debug.log(f"Found {len(media_files)} media files to process", category="file", force=True)
# Validate caching with multi-GPU (not supported in CLI - would need shared memory)
if (args.cache_dit or args.cache_vae) and len(device_list) > 1:
debug.log(
"Model caching requires single GPU selection (you selected multiple GPUs). "
"Disabling caching for this run.",
level="WARNING", category="cache", force=True
)
args.cache_dit = False
args.cache_vae = False
# Initialize runner cache if caching enabled
runner_cache = {} if (args.cache_dit or args.cache_vae) else None
for idx, file_path in enumerate(media_files, 1):
# Visual separation between files (except before first file)
if idx > 1:
@@ -1085,9 +1206,10 @@ def main() -> None:
output_path = generate_output_path(file_path, file_output_format, args.output,
input_type=get_input_type(file_path))
# Process with explicit output path
# Process with explicit output path and runner cache
frames = process_single_file(file_path, args, device_list, output_path,
format_auto_detected=format_auto_detected)
format_auto_detected=format_auto_detected,
runner_cache=runner_cache)
total_frames_processed += frames
# Restore original format
@@ -1098,8 +1220,27 @@ def main() -> None:
if format_auto_detected:
args.output_format = "mp4" if input_type == "video" else "png"
# Validate caching for single file (would provide no benefit but shouldn't error)
if (args.cache_dit or args.cache_vae):
if len(device_list) > 1:
debug.log(
"Model caching requires single GPU selection (you selected multiple GPUs). "
"Disabling caching for this run.",
level="WARNING", category="cache", force=True
)
args.cache_dit = False
args.cache_vae = False
else:
debug.log(
"Model caching has no benefit for single file processing (only useful for directories). "
"Consider removing --cache_dit/--cache_vae for single files.",
category="tip", force=True
)
# No caching for single file (no benefit)
frames = process_single_file(args.input, args, device_list, args.output,
format_auto_detected=format_auto_detected)
format_auto_detected=format_auto_detected,
runner_cache=None)
total_frames_processed += frames
else:
+6 -6
View File
@@ -75,8 +75,8 @@ def encode_all_batches(
seed: int = 42,
progress_callback: Optional[Callable[[int, int, int, str], None]] = None,
temporal_overlap: int = 0,
res_w: int = 1072,
max_res_w: int = 0,
resolution: int = 1080,
max_resolution: int = 0,
input_noise_scale: float = 0.0,
color_correction: str = "wavelet"
) -> Dict[str, Any]:
@@ -95,8 +95,8 @@ def encode_all_batches(
seed: Random seed for deterministic VAE sampling (default: 42)
progress_callback: Optional callback(current, total, frames, phase_name)
temporal_overlap: Overlapping frames between batches for continuity
res_w: Target resolution for shortest edge
max_res_w: Maximum resolution for any edge (0 = no limit)
resolution: Target resolution for shortest edge
max_resolution: Maximum resolution for any edge (0 = no limit)
input_noise_scale: Scale for input noise (0.0-1.0). Adds noise to input images
before VAE encoding to reduce artifacts at high resolutions.
color_correction: Color correction method - "wavelet", "adain", or "none" (default: "wavelet")
@@ -142,10 +142,10 @@ def encode_all_batches(
# Setup video transformation pipeline and compute dimensions if not already done
if 'true_target_dims' not in ctx:
sample_frame = images[0].permute(2, 0, 1).unsqueeze(0)
setup_video_transform(ctx, res_w, max_res_w, debug, sample_frame)
setup_video_transform(ctx, resolution, max_resolution, debug, sample_frame)
del sample_frame
else:
setup_video_transform(ctx, res_w, max_res_w, debug)
setup_video_transform(ctx, resolution, max_resolution, debug)
# Detect if input is RGBA (4 channels)
ctx['is_rgba'] = images[0].shape[-1] == 4
+13 -13
View File
@@ -44,13 +44,13 @@ from ..utils.constants import get_script_directory
script_directory = get_script_directory()
def prepare_video_transforms(res_w: int, max_res_w: int = 0, debug: Optional['Debug'] = None) -> Compose:
def prepare_video_transforms(resolution: int, max_resolution: int = 0, debug: Optional['Debug'] = None) -> Compose:
"""
Prepare optimized video transformation pipeline
Args:
res_w (int): Target resolution for shortest edge
max_res_w (int): Maximum resolution for any edge (0 = no limit)
resolution (int): Target resolution for shortest edge
max_resolution (int): Maximum resolution for any edge (0 = no limit)
debug (Debug, optional): Debug instance for logging
Returns:
@@ -64,18 +64,18 @@ def prepare_video_transforms(res_w: int, max_res_w: int = 0, debug: Optional['De
- Memory-efficient tensor operations
"""
if debug:
msg = f"Initializing video transformation pipeline for {res_w}px (shortest edge)"
if max_res_w > 0:
msg += f", max {max_res_w}px (any edge)"
msg = f"Initializing video transformation pipeline for {resolution}px (shortest edge)"
if max_resolution > 0:
msg += f", max {max_resolution}px (any edge)"
debug.log(msg, category="setup", indent_level=1)
return Compose([
NaResize(
resolution=res_w,
resolution=resolution,
mode="side",
# Upsample image, model only trained for high res
downsample_only=False,
max_resolution=max_res_w,
max_resolution=max_resolution,
),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
DivisiblePad((16, 16)),
@@ -84,7 +84,7 @@ def prepare_video_transforms(res_w: int, max_res_w: int = 0, debug: Optional['De
])
def setup_video_transform(ctx: Dict[str, Any], res_w: int, max_res_w: int = 0,
def setup_video_transform(ctx: Dict[str, Any], resolution: int, max_resolution: int = 0,
debug: Optional['Debug'] = None,
sample_frame: Optional[torch.Tensor] = None) -> Tuple[int, int, int, int]:
"""
@@ -92,8 +92,8 @@ def setup_video_transform(ctx: Dict[str, Any], res_w: int, max_res_w: int = 0,
Args:
ctx: Generation context dictionary
res_w: Target resolution for shortest edge
max_res_w: Maximum resolution for any edge (0 = no limit)
resolution: Target resolution for shortest edge
max_resolution: Maximum resolution for any edge (0 = no limit)
debug: Debug instance for logging
sample_frame: Optional sample frame tensor (C, H, W) to compute dimensions
@@ -119,13 +119,13 @@ def setup_video_transform(ctx: Dict[str, Any], res_w: int, max_res_w: int = 0,
return 0, 0, 0, 0
# Create transformation pipeline (first time or after cleanup)
ctx['video_transform'] = prepare_video_transforms(res_w, max_res_w, debug)
ctx['video_transform'] = prepare_video_transforms(resolution, max_resolution, debug)
# Compute dimensions if sample frame provided
if sample_frame is not None:
# Get true target size (after resize, before padding)
temp_transform = Compose([
NaResize(resolution=res_w, mode="side", downsample_only=False, max_resolution=max_res_w),
NaResize(resolution=resolution, mode="side", downsample_only=False, max_resolution=max_resolution),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0))
])
resized_sample = temp_transform(sample_frame)
+3 -3
View File
@@ -151,7 +151,7 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
@classmethod
def execute(cls, image: torch.Tensor, dit: Dict[str, Any], vae: Dict[str, Any],
seed: int, new_resolution: int = 1072, max_resolution: int = 0, batch_size: int = 5,
seed: int, new_resolution: int = 1080, max_resolution: int = 0, batch_size: int = 5,
temporal_overlap: int = 0, prepend_frames: int = 0,
color_correction: str = "wavelet", input_noise_scale: float = 0.0,
latent_noise_scale: float = 0.0, offload_device: str = "none",
@@ -425,8 +425,8 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
seed=seed,
progress_callback=progress_callback,
temporal_overlap=temporal_overlap,
res_w=new_resolution,
max_res_w=max_resolution,
resolution=new_resolution,
max_resolution=max_resolution,
input_noise_scale=input_noise_scale,
color_correction=color_correction
)