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