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:
+2
-1
@@ -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
@@ -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")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user