feat: Add temporal overlap with blending and prepend frame options to CLI

- Add temporal overlap with crossfading (works with multi GPU)
- Add option to prepend frames to avoid artifacts at the start
This commit is contained in:
Benjamin Herb
2025-08-09 04:12:37 +02:00
parent ad62a12a3e
commit 7737efd22a
+133 -16
View File
@@ -47,7 +47,7 @@ from src.utils.debug import Debug
debug = Debug(enabled=False) # Default to disabled, can be enabled via CLI
def extract_frames_from_video(video_path, skip_first_frames=0, load_cap=None):
def extract_frames_from_video(video_path, skip_first_frames=0, load_cap=None, prepend_frames=0):
"""
Extract frames from video and convert to tensor format
@@ -80,6 +80,8 @@ def extract_frames_from_video(video_path, skip_first_frames=0, load_cap=None):
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")
if prepend_frames:
debug.log(f"Will prepend {prepend_frames} frames to the video", category="info")
frames = []
frame_idx = 0
@@ -121,7 +123,12 @@ def extract_frames_from_video(video_path, skip_first_frames=0, load_cap=None):
raise ValueError(f"No frames extracted from video: {video_path}")
debug.log(f"Extracted {len(frames)} frames", category="success")
# preprend frames if requested (reverse of the first few frames)
if prepend_frames > 0:
prepend_frames = min(prepend_frames, len(frames))
frames = frames[-prepend_frames:] + frames
# Convert to tensor [T, H, W, C] and cast to Float16 for ComfyUI compatibility
frames_tensor = torch.from_numpy(np.stack(frames)).to(torch.float16)
@@ -203,6 +210,57 @@ def save_frames_to_png(frames_tensor, output_dir, base_name):
debug.log(f"PNG saving completed: {total} files in '{output_dir}'", category="success")
def apply_temporal_overlap_blending(frames_tensor, batch_size, overlap):
"""
Blend frames with temporal overlap in pixel space and remove duplicates.
Args:
frames_tensor (torch.Tensor): [T, H, W, C], Float16 in [0,1]
batch_size (int): Frames per batch used during generation
overlap (int): Overlapping frames between consecutive batches
Returns:
torch.Tensor: Blended frames [T, H, W, C] with duplicates removed
"""
T = frames_tensor.shape[0]
if overlap <= 0 or batch_size <= overlap or T <= batch_size:
return frames_tensor
device = frames_tensor.device
dtype = frames_tensor.dtype
output = frames_tensor[:batch_size]
input_pos = batch_size
while input_pos < T:
remaining_frames = T - input_pos
current_batch_size = min(batch_size, remaining_frames)
if current_batch_size <= overlap:
break
current_batch = frames_tensor[input_pos:input_pos + current_batch_size]
prev_tail = output[-overlap:] # overlap frames from previous output
cur_head = current_batch[:overlap] # overlap frames from current batch
# Crossfade (clamping so that the blending only happens in the center of the overlap)
w_prev = torch.clamp(torch.linspace(2.0, -1.0, steps=overlap, device=device, dtype=dtype), 0.0, 1.0).view(overlap, 1, 1, 1)
w_cur = 1.0 - w_prev
blended = prev_tail * w_prev + cur_head * w_cur
# Replace the last overlap frames in output with blended result
output = torch.cat([output[:-overlap], blended], dim=0)
# Append the non-overlapping part of current batch (if any)
if overlap < current_batch_size:
non_overlapping = current_batch[overlap:]
output = torch.cat([output, non_overlapping], dim=0)
input_pos += current_batch_size
return output
def _worker_process(proc_idx, device_id, frames_np, shared_args, return_queue):
"""Worker process that performs upscaling on a slice of frames using a dedicated GPU."""
# 1. Limit CUDA visibility to the chosen GPU BEFORE importing torch-heavy deps
@@ -249,8 +307,26 @@ def _worker_process(proc_idx, device_id, frames_np, shared_args, return_queue):
def _gpu_processing(frames_tensor, device_list, args):
"""Split frames and process them in parallel on multiple GPUs."""
num_devices = len(device_list)
# split frames tensor along time dimension
chunks = torch.chunk(frames_tensor, num_devices, dim=0)
total_frames = frames_tensor.shape[0]
# Create overlapping chunks (for multi GPU); ensures every chunk is
# a multiple of batch_size (except last one) to avoid blending issues
if args.temporal_overlap > 0 and num_devices > 1:
chunk_with_overlap = total_frames // num_devices + args.temporal_overlap
if args.batch_size > 1:
chunk_with_overlap = ((chunk_with_overlap + args.batch_size - 1) // args.batch_size) * args.batch_size
base_chunk_size = chunk_with_overlap - args.temporal_overlap
chunks = []
for i in range(num_devices):
start_idx = i * base_chunk_size
if i == num_devices - 1: # last chunk/device
end_idx = total_frames
else:
end_idx = min(start_idx + chunk_with_overlap, total_frames)
chunks.append(frames_tensor[start_idx:end_idx])
else:
chunks = torch.chunk(frames_tensor, num_devices, dim=0)
manager = mp.Manager()
return_queue = manager.Queue()
@@ -265,12 +341,12 @@ def _gpu_processing(frames_tensor, device_list, args):
"seed": args.seed,
"res_w": args.resolution,
"batch_size": args.batch_size,
"temporal_overlap": 0,
"temporal_overlap": args.temporal_overlap,
"block_swap_config": {
'blocks_to_swap': args.blocks_to_swap,
'use_none_blocking': args.use_none_blocking,
'offload_io_components': args.offload_io_components,
'cache_model': args.cache_model,
'cache_model': False, # No caching in CLI mode
},
"vae_tiling_enabled": args.vae_tiling_enabled,
"vae_tile_size": args.vae_tile_size,
@@ -295,8 +371,38 @@ def _gpu_processing(frames_tensor, device_list, args):
for p in workers:
p.join()
# Concatenate results in original order
result_tensor = torch.from_numpy(np.concatenate(results_np, axis=0)).to(torch.float16)
# Concatenate results with overlap handling
if args.temporal_overlap > 0 and num_devices > 1:
# Reconstruct results considering overlap
result_list = []
overlap = args.temporal_overlap
for idx, res_np in enumerate(results_np):
if idx == 0:
# First chunk: keep all frames
result_list.append(torch.from_numpy(res_np).to(torch.float16))
elif idx == num_devices - 1:
# Last chunk: skip overlap frames at the beginning
chunk_tensor = torch.from_numpy(res_np).to(torch.float16)
if chunk_tensor.shape[0] > overlap:
result_list.append(chunk_tensor[overlap:])
else:
# If chunk is smaller than overlap, skip it entirely
pass
else:
# Middle chunks: skip overlap at beginning, keep overlap at end
chunk_tensor = torch.from_numpy(res_np).to(torch.float16)
if chunk_tensor.shape[0] > overlap:
result_list.append(chunk_tensor[overlap:])
if result_list:
result_tensor = torch.cat(result_list, dim=0)
else:
result_tensor = torch.from_numpy(results_np[0]).to(torch.float16)
else:
# Original concatenation without overlap handling
result_tensor = torch.from_numpy(np.concatenate(results_np, axis=0)).to(torch.float16)
return result_tensor
@@ -311,7 +417,7 @@ def parse_arguments():
parser.add_argument("--resolution", type=int, default=1072,
help="Target resolution of the short side (default: 1072)")
parser.add_argument("--batch_size", type=int, default=1,
help="Number of frames per batch (default: 5)")
help="Number of frames per batch (default: 1)")
parser.add_argument("--model", type=str, default="seedvr2_ema_3b_fp8_e4m3fn.safetensors",
choices=[
"seedvr2_ema_3b_fp16.safetensors",
@@ -340,8 +446,10 @@ def parse_arguments():
help="Number of blocks to swap for VRAM optimization (default: 0, disabled), up to 32 for 3B model, 36 for 7B")
parser.add_argument("--use_none_blocking", action="store_true",
help="Use non-blocking memory transfers for VRAM optimization")
parser.add_argument("--cache_model", action="store_true",
help="Cache model weights in memory to avoid reloading")
parser.add_argument("--temporal_overlap", type=int, default=0,
help="Temporal overlap for processing (default: 0, no temporal overlap)")
parser.add_argument("--prepend_frames", type=int, default=0,
help="Number of frames to prepend to the video (default: 0). This can help with artifacts at the start of the video and are removed after processing")
parser.add_argument("--offload_io_components", action="store_true",
help="Offload IO components to CPU for VRAM optimization")
parser.add_argument("--vae_tiling_enabled", action="store_true",
@@ -350,7 +458,6 @@ def parse_arguments():
help="VAE tile size for tiled decoding (default: 512). Only used if --vae_tiling_enabled is set")
parser.add_argument("--vae_tile_overlap", type=int, default=128,
help="VAE tile overlap for tiled decoding (default: 128). Only used if --vae_tiling_enabled is set")
return parser.parse_args()
@@ -387,11 +494,12 @@ def main():
frames_tensor, original_fps = extract_frames_from_video(
args.video_path,
args.skip_first_frames,
args.load_cap
args.load_cap,
args.prepend_frames
)
debug.log(f"Frame extraction time: {time.time() - start_time:.2f}s", category="general")
# debug.log(f"📊 Initial VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f}GB", category="memory")
# debug.log(f"Initial VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f}GB", category="memory")
# Parse GPU list
device_list = [d.strip() for d in str(args.cuda_device).split(',') if d.strip()] if args.cuda_device else ["0"]
@@ -403,8 +511,17 @@ def main():
debug.log(f"Generation time: {generation_time:.2f}s", category="general")
debug.log(f"Peak VRAM usage: {torch.cuda.max_memory_allocated() / 1024**3:.2f}GB", category="memory")
if args.temporal_overlap > 0:
debug.log(f"Applying temporal overlap with blending", category="generation")
result = apply_temporal_overlap_blending(result, args.batch_size, args.temporal_overlap)
debug.log(f"Result shape: {result.shape}, dtype: {result.dtype}", category="memory")
if args.prepend_frames > 0:
debug.log(f"Removing prepended ({args.prepend_frames}) frames from the results)", category="generation")
result = result[args.prepend_frames:]
debug.log(f"Result shape after removing prepended frames: {result.shape}", category="info")
# After generation_time calculation, choose saving method
if args.output_format == "png":
# Ensure output treated as directory
@@ -444,4 +561,4 @@ def main():
if __name__ == "__main__":
main()
main()