Add streaming mode for memory-efficient long video processing

- New --chunk_size flag enables streaming mode, processing video in bounded chunks
- Supports both MP4 output (single file) and PNG sequence output while streaming
- Preserves --load_cap for total frame limiting (backward compatible)
- Model caching now works between chunks when --cache_dit/--cache_vae enabled
- Instant frame seeking with cv2.CAP_PROP_POS_FRAMES (fixes slow skip on long videos)
- Early exit for empty/exhausted videos
- Minor: function renames (save_frames_to_png → save_frames_to_image), log message cleanup

Inspired by PR #353 - thank you @disk02 for the initial chunked_mode implementation
This commit is contained in:
Adrien Toupet
2025-12-08 22:05:20 -05:00
parent bbd7e5ac02
commit a70d82e3aa
3 changed files with 237 additions and 192 deletions
+12 -2
View File
@@ -771,9 +771,18 @@ The CLI provides comprehensive options for single-GPU, multi-GPU, and batch proc
# Basic image upscaling
python inference_cli.py image.jpg
# Basic video video upscaling with temporal consistency
# Basic video upscaling with temporal consistency
python inference_cli.py video.mp4 --resolution 720 --batch_size 33
# Streaming mode for long videos (memory-efficient)
# Processes video in chunks of 330 frames to avoid loading entire video into RAM
# Use --temporal_overlap to ensure smooth transitions between chunks
python inference_cli.py long_video.mp4 \
--resolution 1080 \
--batch_size 33 \
--chunk_size 330 \
--temporal_overlap 3
# Multi-GPU processing with temporal overlap
python inference_cli.py video.mp4 \
--cuda_device 0,1 \
@@ -830,7 +839,8 @@ python inference_cli.py media_folder/ \
- `--batch_size`: Frames per batch (must follow 4n+1: 1, 5, 9, 13, 17, 21...). Ideally matches shot length for best temporal consistency (default: 5)
- `--seed`: Random seed for reproducibility (default: 42)
- `--skip_first_frames`: Skip N initial frames (default: 0)
- `--load_cap`: Load maximum N frames from video. 0 = load all (default: 0)
- `--load_cap`: Maximum total frames to load from video. 0 = load all (default: 0)
- `--chunk_size`: Frames per chunk for streaming mode. When > 0, processes video in memory-bounded chunks of N frames, writing each chunk before loading the next. Essential for long videos that would otherwise exceed RAM. Use with `--temporal_overlap` for seamless chunk transitions. 0 = load all frames at once (default: 0)
- `--prepend_frames`: Prepend N reversed frames to reduce start artifacts (auto-removed) (default: 0)
- `--temporal_overlap`: Frames to overlap between batches/GPUs for smooth blending (default: 0)
+224 -189
View File
@@ -8,10 +8,12 @@ Supports single and multi-GPU processing with advanced memory optimization.
Key Features:
• Multi-GPU Processing: Automatic workload distribution across multiple GPUs with
temporal overlap blending for seamless transitions
• Streaming Mode: Memory-efficient processing of long videos in chunks, avoiding
full video loading into RAM while maintaining temporal consistency
• Memory Optimization: BlockSwap for limited VRAM, VAE tiling for large resolutions,
intelligent tensor offloading between processing phases
• Performance: Torch.compile integration, BFloat16 compute pipeline,
efficient model caching for batch processing
efficient model caching for batch and streaming processing
• Flexibility: Multiple output formats (MP4/PNG), advanced color correction methods,
directory batch processing with auto-format detection
• Quality Control: Temporal overlap blending, frame prepending for artifact reduction,
@@ -125,6 +127,7 @@ from src.core.generation_phases import (
postprocess_all_batches
)
from src.utils.debug import Debug
from src.optimization.memory_manager import clear_memory
debug = Debug(enabled=False) # Will be enabled via --debug CLI flag
@@ -356,6 +359,9 @@ def process_single_file(input_path: str, args: argparse.Namespace, device_list:
"""
Process a single video or image file with optional model caching.
For videos, supports streaming mode (chunk_size > 0) which processes in memory-bounded
chunks with temporal overlap for seamless transitions between chunks.
Args:
input_path: Path to input file
args: Command-line arguments with all processing settings
@@ -365,7 +371,7 @@ def process_single_file(input_path: str, args: argparse.Namespace, device_list:
runner_cache: Optional cache dict for model reuse across multiple files
Returns:
Number of frames processed from the input
Number of frames written to output
"""
input_type = get_input_type(input_path)
@@ -375,19 +381,6 @@ def process_single_file(input_path: str, args: argparse.Namespace, device_list:
debug.log(f"Processing {input_type}: {Path(input_path).name}", category="generation", force=True)
# Extract frames
if input_type == "video":
start_time = time.time()
frames_tensor, original_fps = extract_frames_from_video(
input_path, args.skip_first_frames, args.load_cap
)
debug.log(f"Frame extraction time: {time.time() - start_time:.2f}s", category="timing")
else:
frames_tensor, original_fps = extract_frames_from_image(input_path)
# Track frames before processing (for FPS calculation)
input_frame_count = len(frames_tensor)
# Generate or validate output path
if output_path is None:
output_path = generate_output_path(input_path, args.output_format, input_type=input_type)
@@ -400,237 +393,270 @@ def process_single_file(input_path: str, args: argparse.Namespace, device_list:
format_prefix = "Auto-detected" if format_auto_detected else "Requested"
debug.log(f"{format_prefix} output format: {args.output_format}", category="info", force=True, indent_level=1)
# Process frames
processing_start = time.time()
# Use direct processing if caching enabled OR on Mac (MPS doesn't support multiprocessing well)
if runner_cache is not None or platform.system() == "Darwin":
# Direct single-GPU processing (required for Mac MPS, optional for 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
is_png_format = args.output_format == "png"
is_single_image = input_type == "image"
if is_png_format and is_single_image:
# Single PNG file
os.makedirs(Path(output_path).parent, exist_ok=True)
frame_np = (result[0].cpu().numpy() * 255.0).astype(np.uint8)
# Convert RGB(A) to BGR(A) based on channel count
if frame_np.shape[2] == 4:
frame_save = cv2.cvtColor(frame_np, cv2.COLOR_RGBA2BGRA)
else:
frame_save = cv2.cvtColor(frame_np, cv2.COLOR_RGB2BGR)
cv2.imwrite(output_path, frame_save)
elif is_png_format:
# PNG sequence (save_frames_to_png creates directory internally)
save_frames_to_png(result, output_path, base_name=Path(input_path).stem)
else:
# Video file
os.makedirs(Path(output_path).parent, exist_ok=True)
save_frames_to_video(result, output_path, original_fps)
# Log appropriate save message based on format
if is_png_format and not is_single_image:
debug.log(f"PNG frames saved in directory: {output_path}", category="file", force=True)
else:
# === VIDEO PROCESSING ===
if input_type == "video":
if not os.path.exists(input_path):
raise FileNotFoundError(f"Video file not found: {input_path}")
cap = cv2.VideoCapture(input_path)
if not cap.isOpened():
raise ValueError(f"Cannot open video file: {input_path}")
fps = cap.get(cv2.CAP_PROP_FPS) or 30.0
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
debug.log(f"Video info: {total_frames} frames, {width}x{height}, {fps:.2f} FPS", category="info")
# Skip initial frames
if args.skip_first_frames > 0:
debug.log(f"Skipping first {args.skip_first_frames} frames", category="info")
cap.set(cv2.CAP_PROP_POS_FRAMES, args.skip_first_frames)
# Calculate frames to process (apply load_cap if set)
frames_to_process = total_frames - args.skip_first_frames
if args.load_cap > 0:
frames_to_process = min(frames_to_process, args.load_cap)
# Early exit for empty/exhausted video
if frames_to_process <= 0:
debug.log(f"No frames to process after skipping {args.skip_first_frames} of {total_frames}",
level="WARNING", category="file", force=True)
cap.release()
return 0
# Streaming mode: process in chunks
chunk_size = args.chunk_size if args.chunk_size > 0 else frames_to_process
streaming = args.chunk_size > 0
total_chunks = (frames_to_process + chunk_size - 1) // chunk_size # ceiling division
if streaming:
debug.log(f"Streaming mode: chunks of {chunk_size} frames, overlap={args.temporal_overlap}",
category="info", force=True, indent_level=1)
is_png = args.output_format == "png"
video_writer = None
prev_raw_tail = None
overlap = args.temporal_overlap
frames_written = 0
chunk_idx = 0
base_name = Path(input_path).stem
chunk_args = argparse.Namespace(**vars(args))
frames_read = 0
while frames_read < frames_to_process:
# Read only remaining frames needed
new_frames = _read_frames_from_cap(cap, min(chunk_size, frames_to_process - frames_read))
if new_frames is None:
break
frames_read += new_frames.shape[0]
chunk_idx += 1
# Disable prepend_frames after first chunk
if chunk_idx > 1:
chunk_args.prepend_frames = 0
# Prepend context from previous chunk
if prev_raw_tail is not None and overlap > 0:
context_count = min(overlap, prev_raw_tail.shape[0])
frames = torch.cat([prev_raw_tail[-context_count:], new_frames], dim=0)
else:
frames = new_frames
context_count = 0
if streaming:
if chunk_idx > 1:
debug.log("", category="none", force=True)
debug.log("━" * 60, category="none", force=True)
debug.log("", category="none", force=True)
debug.log(f"Chunk {chunk_idx}/{total_chunks}: {new_frames.shape[0]} new + {context_count} context frames",
category="generation", force=True)
debug.log("", category="none", force=True)
# Process frames (multiprocessing only for multi-GPU)
if len(device_list) > 1:
result = _gpu_processing(frames, device_list, chunk_args)
else:
result = _single_gpu_direct_processing(frames, chunk_args, device_list[0], runner_cache)
# Drop context frames from output
if context_count > 0:
result = result[context_count:]
# Save output using dedicated functions
if is_png:
save_frames_to_image(result, output_path, base_name, start_index=frames_written)
else:
video_writer = save_frames_to_video(result, output_path, fps, writer=video_writer)
frames_written += result.shape[0]
# Save tail for next chunk context
prev_raw_tail = new_frames[-overlap:].clone() if overlap > 0 else None
# Cleanup
del result, frames
# Memory cleanup between chunks
if streaming:
clear_memory(debug=debug, deep=True, force=True, timer_name="chunk_cleanup")
cap.release()
if video_writer is not None:
video_writer.release()
if streaming:
debug.log("", category="none", force=True)
debug.log(f"Streaming complete: {frames_written} frames in {chunk_idx} chunks", category="success", force=True)
debug.log(f"Output saved to: {output_path}", category="file", force=True)
return frames_written
return input_frame_count
# === IMAGE PROCESSING ===
frames_tensor, _ = extract_frames_from_image(input_path)
processing_start = time.time()
# Process frames (multiprocessing only for multi-GPU)
if len(device_list) > 1:
result = _gpu_processing(frames_tensor, device_list, args)
else:
result = _single_gpu_direct_processing(frames_tensor, args, device_list[0], runner_cache)
debug.log(f"Processing time: {time.time() - processing_start:.2f}s", category="timing")
# Save single image
os.makedirs(Path(output_path).parent, exist_ok=True)
frame_np = (result[0].cpu().numpy() * 255.0).astype(np.uint8)
_save_image_bgr(frame_np, output_path)
debug.log(f"Output saved to: {output_path}", category="file", force=True)
return 1
def extract_frames_from_video(
video_path: str,
skip_first_frames: int = 0,
load_cap: Optional[int] = None
) -> Tuple[torch.Tensor, float]:
def _read_frames_from_cap(cap: cv2.VideoCapture, max_frames: int) -> Optional[torch.Tensor]:
"""
Extract frames from video file and convert to tensor format.
Reads video using OpenCV, converts BGR to RGB, normalizes to [0,1] range.
Note: Frame prepending is handled later in the processing pipeline via
compute_generation_info(), not in this function.
Read up to max_frames from an already-open VideoCapture.
Args:
video_path: Path to input video file
skip_first_frames: Number of initial frames to skip (default: 0)
load_cap: Maximum number of frames to load, None loads all (default: None)
cap: An already opened cv2.VideoCapture instance
max_frames: Maximum number of frames to read in this call
Returns:
Tuple containing:
- frames_tensor: Frames in format [T, H, W, C], Float32, range [0,1]
- fps: Original video frames per second
Raises:
FileNotFoundError: If video file doesn't exist
ValueError: If video cannot be opened or no frames extracted
Tensor [T, H, W, C] float32 [0,1], or None if no frames available
"""
debug.log(f"Extracting frames from video: {video_path}", category="file")
if not os.path.exists(video_path):
raise FileNotFoundError(f"Video file not found: {video_path}")
# Open video
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise ValueError(f"Cannot open video file: {video_path}")
# Get video properties
fps = cap.get(cv2.CAP_PROP_FPS)
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
debug.log(f"Video info: {frame_count} frames, {width}x{height}, {fps:.2f} FPS", category="info")
if skip_first_frames:
debug.log(f"Will skip first {skip_first_frames} frames", category="info")
if load_cap:
debug.log(f"Will load maximum {load_cap} frames", category="info")
frames = []
frame_idx = 0
frames_loaded = 0
while True:
for _ in range(max_frames):
ret, frame = cap.read()
if not ret:
break
# Skip first frame if requested
if frame_idx < skip_first_frames:
frame_idx += 1
continue
if skip_first_frames > 0 and frame_idx == skip_first_frames:
debug.log(f"Skipped first {skip_first_frames} frames", category="info")
# Check load cap
if load_cap is not None and load_cap > 0 and frames_loaded >= load_cap:
debug.log(f"Reached load cap of {load_cap} frames", category="info")
break
# Convert BGR to RGB
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# Convert to float32 and normalize to 0-1
frame = frame.astype(np.float32) / 255.0
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
frames.append(frame)
frame_idx += 1
frames_loaded += 1
if debug.enabled and frames_loaded % 100 == 0:
total_to_load = min(frame_count, load_cap) if load_cap else frame_count
debug.log(f"Extracted {frames_loaded}/{total_to_load} frames", category="file")
cap.release()
if len(frames) == 0:
raise ValueError(f"No frames extracted from video: {video_path}")
debug.log(f"Extracted {len(frames)} frames", category="success")
if not frames:
return None
return torch.from_numpy(np.stack(frames)).to(torch.float32)
# Convert to tensor (will be cast to compute_dtype in worker process)
frames_tensor = torch.from_numpy(np.stack(frames)).to(torch.float32)
debug.log(f"Frames tensor shape: {frames_tensor.shape}, dtype: {frames_tensor.dtype}", category="memory")
return frames_tensor, fps
def _save_image_bgr(frame_np: np.ndarray, file_path: str) -> None:
"""
Save a single RGB(A) uint8 frame to disk, converting to BGR(A) for OpenCV.
Args:
frame_np: Frame as uint8 numpy array [H, W, C] where C is 3 (RGB) or 4 (RGBA)
file_path: Output file path
"""
if frame_np.shape[2] == 4:
frame_bgr = cv2.cvtColor(frame_np, cv2.COLOR_RGBA2BGRA)
else:
frame_bgr = cv2.cvtColor(frame_np, cv2.COLOR_RGB2BGR)
cv2.imwrite(file_path, frame_bgr)
def save_frames_to_video(
frames_tensor: torch.Tensor,
output_path: str,
fps: float = 30.0
) -> None:
fps: float = 30.0,
writer: Optional[cv2.VideoWriter] = None
) -> Optional[cv2.VideoWriter]:
"""
Save frames tensor to MP4 video file.
Converts tensor from Float32 [0,1] to uint8 [0,255], RGB to BGR for OpenCV,
and writes to video file using mp4v codec.
and writes to video file using mp4v codec. Supports streaming mode where
an existing writer is passed and kept open for subsequent chunks.
Args:
frames_tensor: Frames in format [T, H, W, C], Float32, range [0,1]
output_path: Output video file path (will be created if doesn't exist)
output_path: Output video file path (directory created if doesn't exist)
fps: Frames per second for output video (default: 30.0)
writer: Existing VideoWriter for streaming (if None, creates new one)
Returns:
VideoWriter if streaming mode (caller must close), None if standalone mode
Raises:
ValueError: If video writer cannot be initialized
"""
debug.log(f"Saving {frames_tensor.shape[0]} frames to video: {output_path}", category="file")
# Convert tensor to numpy and denormalize
frames_np = frames_tensor.cpu().numpy()
frames_np = (frames_np * 255.0).astype(np.uint8)
# Get video properties
frames_np = (frames_tensor.cpu().numpy() * 255.0).astype(np.uint8)
T, H, W, C = frames_np.shape
# Initialize video writer
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (W, H))
if writer is None:
debug.log(f"Saving {T} frames to video: {output_path}", category="file")
os.makedirs(Path(output_path).parent, exist_ok=True)
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
writer = cv2.VideoWriter(output_path, fourcc, fps, (W, H))
if not writer.isOpened():
raise ValueError(f"Cannot create video writer for: {output_path}")
if not out.isOpened():
raise ValueError(f"Cannot create video writer for: {output_path}")
# Write frames
for i, frame in enumerate(frames_np):
# Convert RGB to BGR for OpenCV
frame_bgr = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
out.write(frame_bgr)
writer.write(frame_bgr)
if debug.enabled and (i + 1) % 100 == 0:
debug.log(f"Saved {i + 1}/{T} frames", category="file")
out.release()
debug.log(f"Written {i + 1}/{T} frames", category="file")
debug.log(f"Video saved successfully: {output_path}", category="success")
return writer # Caller always closes
def save_frames_to_png(
def save_frames_to_image(
frames_tensor: torch.Tensor,
output_dir: str,
base_name: str
) -> None:
base_name: str,
start_index: int = 0
) -> int:
"""
Save frames tensor as sequential PNG image files.
Each frame saved as {base_name}_{index:05d}.png with zero-padded indices.
Each frame saved as {base_name}_{index:0Nd}.png with zero-padded indices.
Converts Float32 [0,1] to uint8 [0,255] and RGB(A) to BGR(A) for OpenCV.
Args:
frames_tensor: Frames in format [T, H, W, C], Float32, range [0,1]
output_dir: Directory to save PNG files (created if doesn't exist)
base_name: Base name for output files (e.g., "frame" → "frame_00000.png")
start_index: Starting index for filenames (for streaming continuation)
Returns:
Number of frames saved
"""
debug.log(f"Saving {frames_tensor.shape[0]} frames as PNGs to directory: {output_dir}", category="file")
# Ensure output directory exists
os.makedirs(output_dir, exist_ok=True)
# Convert to numpy uint8 RGB
frames_np = (frames_tensor.cpu().numpy() * 255.0).astype(np.uint8)
total = frames_np.shape[0]
digits = max(5, len(str(total))) # at least 5 digits
if start_index == 0:
debug.log(f"Saving {total} frames as PNGs to directory: {output_dir}", category="file")
digits = 6 # Supports up to 999,999 frames (~11.5 hours at 24fps)
for idx, frame in enumerate(frames_np):
filename = f"{base_name}_{idx:0{digits}d}.png"
filename = f"{base_name}_{start_index + idx:0{digits}d}.png"
file_path = os.path.join(output_dir, filename)
# Convert RGB(A) to BGR(A) for cv2 based on channel count
if frame.shape[2] == 4:
frame_save = cv2.cvtColor(frame, cv2.COLOR_RGBA2BGRA)
else:
frame_save = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
cv2.imwrite(file_path, frame_save)
_save_image_bgr(frame, file_path)
if debug.enabled and (idx + 1) % 100 == 0:
debug.log(f"Saved {idx + 1}/{total} PNGs", category="file")
debug.log(f"Saved {idx + 1}/{total} images", category="file")
debug.log(f"PNG saving completed: {total} files in '{output_dir}'", category="success")
debug.log(f"Saved {total} images to '{output_dir}'", category="success")
return total
# =============================================================================
@@ -859,7 +885,7 @@ def _single_gpu_direct_processing(
frames_tensor: torch.Tensor,
args: argparse.Namespace,
device_id: str,
runner_cache: Dict[str, Any]
runner_cache: Optional[Dict[str, Any]]
) -> torch.Tensor:
"""
Direct single-GPU processing with model caching support.
@@ -1042,9 +1068,12 @@ Examples:
Basic image upscaling:
python {invocation} image.jpg
Basic video video upscaling with temporal consistency
Basic video upscaling with temporal consistency:
python {invocation} video.mp4 --resolution 720 --batch_size 33
Streaming mode for long videos:
python {invocation} long_video.mp4 --resolution 1080 --batch_size 33 --chunk_size 330 --temporal_overlap 3
Multi-GPU processing with temporal overlap:
python {invocation} video.mp4 --cuda_device 0,1 --resolution 1080 --batch_size 81 --uniform_batch_size --temporal_overlap 3 --prepend_frames 4
@@ -1102,6 +1131,9 @@ Examples:
help="Skip N initial frames (default: 0)")
process_group.add_argument("--load_cap", type=int, default=0,
help="Load maximum N frames from video. 0 = load all (default: 0)")
process_group.add_argument("--chunk_size", type=int, default=0,
help="Frames per chunk for streaming mode. When > 0, processes video in "
"memory-bounded chunks of N frames. 0 = load all frames at once (default: 0)")
process_group.add_argument("--prepend_frames", type=int, default=0,
help="Prepend N reversed frames to reduce start artifacts (auto-removed). Default: 0")
process_group.add_argument("--temporal_overlap", type=int, default=0,
@@ -1377,7 +1409,8 @@ 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)
# Setup caching for single file (benefits streaming mode, no benefit otherwise)
runner_cache = None
if (args.cache_dit or args.cache_vae):
if len(device_list) > 1:
debug.log(
@@ -1387,17 +1420,19 @@ def main() -> None:
)
args.cache_dit = False
args.cache_vae = False
elif args.chunk_size > 0:
# Caching benefits streaming mode (reuse models between chunks)
runner_cache = {}
else:
debug.log(
"Model caching has no benefit for single file processing (only useful for directories). "
"Model caching has no benefit for single file processing (only useful for directories or streaming mode). "
"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,
runner_cache=None)
runner_cache=runner_cache)
total_frames_processed += frames
else:
+1 -1
View File
@@ -1414,7 +1414,7 @@ def postprocess_all_batches(
total_computed += (num_valid_samples - 1) * actual_overlap
frame_info += f" ({total_computed} computed with {' + '.join(adjustments)} removed)"
debug.log(f"Final output assembled: {frame_info}, Resolution: {Wf}x{Hf}px, Channels: {channels_str}",
debug.log(f"Output assembled: {frame_info}, Resolution: {Wf}x{Hf}px, Channels: {channels_str}",
category="generation", force=True)
else:
ctx['final_video'] = torch.empty((0, 0, 0, 0), dtype=ctx['compute_dtype'])