feat: add CPU offloading for intermediate data to reduce VRAM usage

- Store latents and transformed videos on CPU between processing phases reducing VRAM usage to enable larger batch processing
- Make transformed video storage conditional (only when color_correction != none)
- Add all_ori_lengths tracking for consistent trimming
- Use non_blocking=False for CPU-GPU transfers to avoid pinned memory issues
Fixes https://github.com/numz/ComfyUI-SeedVR2_VideoUpscaler/pull/163#issuecomment-3333686784
This commit is contained in:
Adrien Toupet
2025-09-25 11:30:32 -04:00
parent 547ccc749b
commit 4a3efc7883
3 changed files with 67 additions and 38 deletions
+2 -1
View File
@@ -329,7 +329,8 @@ def _worker_process(proc_idx: int, device_id: int, frames_np: np.ndarray,
progress_callback=None,
temporal_overlap=shared_args["temporal_overlap"],
res_w=shared_args["res_w"],
input_noise_scale=shared_args["input_noise_scale"]
input_noise_scale=shared_args["input_noise_scale"],
color_correction=shared_args.get("color_correction", "wavelet")
)
# Phase 2: Upscale all batches
+63 -36
View File
@@ -348,7 +348,8 @@ def prepare_runner(model_name, model_dir, preserve_vram, debug,
def encode_all_batches(runner, ctx=None, images=None, batch_size=5,
preserve_vram=False, debug=None, progress_callback=None,
temporal_overlap=0, res_w=1072, input_noise_scale=0.0):
temporal_overlap=0, res_w=1072, input_noise_scale=0.0,
color_correction="wavelet"):
"""
Phase 1: VAE Encoding for all batches
@@ -368,6 +369,8 @@ def encode_all_batches(runner, ctx=None, images=None, batch_size=5,
res_w: Target resolution for shortest edge
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")
Determines if transformed videos need to be stored for later use.
Returns:
dict: Context containing:
@@ -446,8 +449,12 @@ def encode_all_batches(runner, ctx=None, images=None, batch_size=5,
num_encode_batches += 1
# Pre-allocate lists for memory efficiency
ctx['all_transformed_videos'] = [None] * num_encode_batches
ctx['all_latents'] = [None] * num_encode_batches
ctx['all_ori_lengths'] = [None] * num_encode_batches
if color_correction != "none":
ctx['all_transformed_videos'] = [None] * num_encode_batches
else:
ctx['all_transformed_videos'] = None
encode_idx = 0
@@ -515,14 +522,22 @@ def encode_all_batches(runner, ctx=None, images=None, batch_size=5,
debug.log(f"Applying padding: {padding_frames} frame{'s' if padding_frames != 1 else ''} added ({t} -> {target})", category="video", force=True)
transformed_video = cut_videos(transformed_video)
# Store transformed video (by index for pre-allocation)
ctx['all_transformed_videos'][encode_idx] = (transformed_video, ori_length)
# Store original length for proper trimming later
ctx['all_ori_lengths'][encode_idx] = ori_length
# Store transformed video on CPU if needed for color correction
if color_correction != "none":
# Move to CPU to free VRAM
transformed_video_cpu = transformed_video.to('cpu', non_blocking=False)
ctx['all_transformed_videos'][encode_idx] = transformed_video_cpu
# Encode to latents
cond_latents = runner.vae_encode([transformed_video])
ctx['all_latents'][encode_idx] = cond_latents[0] # Store first element
del cond_latents
# Offload latents to CPU to save VRAM between phases
ctx['all_latents'][encode_idx] = cond_latents[0].to('cpu', non_blocking=False)
del cond_latents, transformed_video
debug.end_timer(f"encode_batch_{encode_idx+1}", f"Encoded batch {encode_idx+1}")
@@ -638,8 +653,8 @@ def upscale_all_batches(runner, ctx=None, preserve_vram=False, debug=None,
debug.log(f"Upscaling batch {upscale_idx+1}/{num_valid_latents}", category="generation", force=True)
debug.start_timer(f"upscale_batch_{upscale_idx+1}")
# Move latent to device with correct dtype
latent = latent.to(ctx['device'], dtype=ctx['compute_dtype'])
# Move latent from CPU to device with correct dtype
latent = latent.to(ctx['device'], dtype=ctx['compute_dtype'], non_blocking=False)
# Generate noise
if torch.mps.is_available():
@@ -684,8 +699,8 @@ def upscale_all_batches(runner, ctx=None, preserve_vram=False, debug=None,
)
debug.end_timer(f"dit_inference_{upscale_idx+1}", f"DiT inference {upscale_idx+1}")
# Store upscaled result (by index for pre-allocation)
ctx['all_upscaled_latents'][upscale_idx] = upscaled[0]
# Store upscaled result on CPU to save VRAM
ctx['all_upscaled_latents'][upscale_idx] = upscaled[0].to('cpu', non_blocking=False)
# Free original latent
ctx['all_latents'][batch_idx] = None
@@ -757,9 +772,10 @@ def decode_all_batches(runner, ctx=None, preserve_vram=False, debug=None, progre
if 'all_upscaled_latents' not in ctx or not ctx['all_upscaled_latents']:
raise ValueError("No upscaled latents found. Run upscale_all_batches first.")
# Validate we have transformed videos for wavelet reconstruction
if 'all_transformed_videos' not in ctx or not ctx['all_transformed_videos']:
raise ValueError("No transformed videos found for reconstruction. Context corrupted.")
# Validate we have transformed videos if color correction is enabled
if color_correction != "none":
if 'all_transformed_videos' not in ctx or not ctx['all_transformed_videos']:
raise ValueError("No transformed videos found for color correction. Context corrupted.")
debug.log("", category="none", force=True)
debug.log("━━━━━━━━ Phase 3: VAE decoding ━━━━━━━━", category="none", force=True)
@@ -768,8 +784,8 @@ def decode_all_batches(runner, ctx=None, preserve_vram=False, debug=None, progre
# Count valid latents
num_valid_latents = len([l for l in ctx['all_upscaled_latents'] if l is not None])
# Pre-allocate to match transformed videos (which matches original batches)
num_batches = len([v for v in ctx['all_transformed_videos'] if v is not None])
# Pre-allocate to match original batches (use ori_lengths which is always available)
num_batches = len([l for l in ctx['all_ori_lengths'] if l is not None])
ctx['batch_samples'] = [None] * num_batches
decode_idx = 0
@@ -789,6 +805,9 @@ def decode_all_batches(runner, ctx=None, preserve_vram=False, debug=None, progre
debug.log(f"Decoding batch {decode_idx+1}/{num_valid_latents}", category="vae", force=True)
debug.start_timer(f"decode_batch_{decode_idx+1}")
# Move latent to device with correct dtype for decoding
upscaled_latent = upscaled_latent.to(ctx['device'], dtype=ctx['compute_dtype'], non_blocking=False)
# Decode latent
debug.start_timer("vae_decode")
samples = runner.vae_decode([upscaled_latent], preserve_vram=preserve_vram)
@@ -797,47 +816,53 @@ def decode_all_batches(runner, ctx=None, preserve_vram=False, debug=None, progre
# Convert to Float16 for efficiency
if samples and len(samples) > 0 and samples[0].dtype != torch.float16:
debug.log(f"Converting from {samples[0].dtype} to Float16", category="precision")
samples = [sample.to(torch.float16, non_blocking=True) for sample in samples]
samples = [sample.to(torch.float16, non_blocking=False) for sample in samples]
# Process samples
samples = optimized_video_rearrange(samples)
# Post-process with wavelet reconstruction
# Post-process samples
for i, sample in enumerate(samples):
# Find corresponding transformed video
video_idx = min(batch_idx, len(ctx['all_transformed_videos']) - 1)
transformed_video, ori_length = ctx['all_transformed_videos'][video_idx]
# Get original length for trimming (always available)
video_idx = min(batch_idx, len(ctx['all_ori_lengths']) - 1)
ori_length = ctx['all_ori_lengths'][video_idx]
# Trim if necessary
# Trim to original length if necessary
if ori_length < sample.shape[0]:
sample = sample[:ori_length]
# Apply color correction based on selected method
if color_correction != "none":
transformed_video = transformed_video.to(ctx['device'])
# Apply color correction if enabled
if color_correction != "none" and ctx['all_transformed_videos'] is not None:
# Get transformed video from CPU
transformed_video = ctx['all_transformed_videos'][video_idx]
# Move to GPU for color correction
transformed_video = transformed_video.to(ctx['device'], non_blocking=False)
input_video = [optimized_single_video_rearrange(transformed_video)]
# Start timing color correction operation
# Apply selected color correction method
debug.start_timer(f"color_correction_{color_correction}")
if color_correction == "wavelet":
debug.log("Applying wavelet color reconstruction (frequency-based)", category="video")
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)], debug)
debug.log("Applying wavelet color reconstruction", category="video")
sample = wavelet_reconstruction(sample, input_video[0], debug)
elif color_correction == "adain":
debug.log("Applying AdaIN color correction (statistical matching)", category="video")
sample = adaptive_instance_normalization(sample, input_video[0][:sample.size(0)])
debug.log("Applying AdaIN color correction", category="video")
sample = adaptive_instance_normalization(sample, input_video[0])
else:
debug.log(f"Unknown color correction method: {color_correction}, skipping", level="WARNING", category="video")
debug.log(f"Unknown color correction method: {color_correction}", level="WARNING", category="video", force=True)
# End timing and log duration
debug.end_timer(f"color_correction_{color_correction}", f"Color correction ({color_correction}) completed")
debug.end_timer(f"color_correction_{color_correction}", f"Color correction ({color_correction})")
del input_video
# Free the transformed video
ctx['all_transformed_videos'][video_idx] = None
del input_video, transformed_video
else:
debug.log("Color correction disabled (set to none)", category="video")
ctx['all_transformed_videos'][video_idx] = None
del transformed_video
# Free the original length entry
ctx['all_ori_lengths'][video_idx] = None
# Convert to final format
sample = optimized_sample_to_image_format(sample)
@@ -877,6 +902,8 @@ def decode_all_batches(runner, ctx=None, preserve_vram=False, debug=None, progre
del ctx['all_upscaled_latents']
if 'all_transformed_videos' in ctx:
del ctx['all_transformed_videos']
if 'all_ori_lengths' in ctx:
del ctx['all_ori_lengths']
debug.log("Assembling final video from decoded batches...", category="video")
+2 -1
View File
@@ -293,7 +293,8 @@ class SeedVR2:
progress_callback=self._progress_callback,
temporal_overlap=temporal_overlap,
res_w=new_resolution,
input_noise_scale=input_noise_scale
input_noise_scale=input_noise_scale,
color_correction=color_correction
)
# Phase 2: Upscale all batches