From 32a049dfd9bac8b24732c571ea694f7d016d181b Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 4 Nov 2025 15:12:28 -0500 Subject: [PATCH] feat: Add CLI model caching for multi-file processing and unify device handling MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- README.md | 4 +- inference_cli.py | 571 +++++++++++++++++++------------ src/core/generation_phases.py | 12 +- src/core/generation_utils.py | 26 +- src/interfaces/video_upscaler.py | 6 +- 5 files changed, 380 insertions(+), 239 deletions(-) diff --git a/README.md b/README.md index 46928e2..5b220c8 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/inference_cli.py b/inference_cli.py index ceeab1e..a6c5ff1 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -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: diff --git a/src/core/generation_phases.py b/src/core/generation_phases.py index 2c7ea23..bab64ce 100644 --- a/src/core/generation_phases.py +++ b/src/core/generation_phases.py @@ -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 diff --git a/src/core/generation_utils.py b/src/core/generation_utils.py index d1273a7..3ab4561 100644 --- a/src/core/generation_utils.py +++ b/src/core/generation_utils.py @@ -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) diff --git a/src/interfaces/video_upscaler.py b/src/interfaces/video_upscaler.py index 7253c98..435b80a 100644 --- a/src/interfaces/video_upscaler.py +++ b/src/interfaces/video_upscaler.py @@ -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 )