From e42f1fd8fb9a131d7c15951530b4f97423bf6c02 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Sun, 26 Apr 2026 03:08:07 +0100 Subject: [PATCH] feat(video): video_interpolate as execution-time stage Promote RIFE interpolation from a save-time kwarg to a real stage of the processing pipeline so per-frame work (detailer, color correction, postprocess scripts) operates on source-rate frames and the inflated stream becomes the saved output. - new modules/processing_video.py with apply_video_interpolation, interpolation_factor, expand_infotexts; PIL/tensor/numpy dispatch - video_interpolate, video_interpolate_scale, video_interpolated fields on StableDiffusionProcessingVideo - process_images_inner runs the helper after the batch loop and inflates infotexts in lockstep - save_video in modules/video.py and modules/video_models/video_save.py short-circuit re-interpolation when p.video_interpolated is set; the user-facing kwarg still flows into metadata --- modules/processing.py | 8 ++ modules/processing_class.py | 3 + modules/processing_video.py | 125 +++++++++++++++++++++++++++++ modules/video.py | 2 + modules/video_models/video_save.py | 4 +- 5 files changed, 140 insertions(+), 2 deletions(-) create mode 100644 modules/processing_video.py diff --git a/modules/processing.py b/modules/processing.py index d8e43466f..0306b20a8 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -536,6 +536,14 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: if shared.state.interrupted: break + if getattr(p, 'video_interpolate', 0) > 0 and len(output_images) > 1: + from modules.processing_video import apply_video_interpolation, expand_infotexts + n_before = len(output_images) + output_images = apply_video_interpolation(p, output_images) + n_after = len(output_images) + if n_after > n_before and len(infotexts) == n_before: + infotexts = expand_infotexts(infotexts, max(0, p.video_interpolate - 1)) + if not p.xyz: if hasattr(shared.sd_model, 'restore_pipeline') and (shared.sd_model.restore_pipeline is not None): shared.sd_model.restore_pipeline() diff --git a/modules/processing_class.py b/modules/processing_class.py index 299f9f9b4..9e1ed3e4f 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -691,6 +691,9 @@ class StableDiffusionProcessingVideo(StableDiffusionProcessing): self.vae_tile_frames: int = kwargs.pop('vae_tile_frames', 0) self.video_engine: str = kwargs.pop('video_engine', None) self.video_model: str = kwargs.pop('video_model', None) + self.video_interpolate: int = kwargs.pop('video_interpolate', 0) + self.video_interpolate_scale: float = kwargs.pop('video_interpolate_scale', 1.0) + self.video_interpolated: bool = False self.scheduler_shift: float = 0.0 debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access super().__init__(**kwargs) diff --git a/modules/processing_video.py b/modules/processing_video.py new file mode 100644 index 000000000..c3138f368 --- /dev/null +++ b/modules/processing_video.py @@ -0,0 +1,125 @@ +""" +Video frame interpolation helper. + +Used by: +- modules.processing.process_images_inner (after process_samples) +- modules.framepack.framepack_worker (before final save_video) +- modules.ltx.ltx_process (before save_video) +- modules.video_models.video_run (before save_video) + +Resolves count and scale from explicit kwargs first, then from +StableDiffusionProcessingVideo.video_interpolate on `p`. Marks the +processing object so save_video can skip its own interpolation pass. + +Forwards count straight to the PIL primitive and count+1 to the tensor +primitive to match the legacy interpolate_frames and video_save.py call +shapes. +""" +from typing import Any +import numpy as np +import torch +from PIL import Image +from modules.logger import log + + +def frames_len(frames: Any): + if frames is None: + return None + if isinstance(frames, list): + return len(frames) + try: + return frames.shape[0] + except Exception: + return None + + +def apply_video_interpolation( + p: Any = None, + frames: Any = None, + count: int = 0, + scale: float = 0.0, + pad: int = 1, + change: float = 0.3, +): + """Inflate a frame stream by RIFE interpolation. + + Dispatches by frames type: + list[PIL.Image] -> rife.interpolate + 4-D torch.Tensor (N,C,H,W) -> rife.interpolate_nchw + np.ndarray (N,H,W,C) -> rife.interpolate_nchw via tensor convert + Sets p.video_interpolated = True after a successful run. + """ + if frames is None: + return frames + if count <= 0: + count = int(getattr(p, 'video_interpolate', 0) or 0) + if count <= 0: + return frames + if scale <= 0: + scale = float(getattr(p, 'video_interpolate_scale', 1.0) or 1.0) + if scale <= 0: + scale = 1.0 + + in_len = frames_len(frames) + in_type = 'unknown' + out = frames + try: + from modules import rife + if isinstance(frames, list) and len(frames) > 0 and isinstance(frames[0], Image.Image): + in_type = 'pil' + out = rife.interpolate(frames, count=count, scale=scale, pad=pad, change=change) + elif torch.is_tensor(frames): + in_type = 'tensor' + interpolated = rife.interpolate_nchw(frames, count=count + 1, scale=scale) + out = torch.cat(interpolated, dim=0) if isinstance(interpolated, list) else interpolated + elif isinstance(frames, np.ndarray): + in_type = 'numpy' + t = torch.from_numpy(frames).permute(0, 3, 1, 2).float() / 255.0 + interpolated = rife.interpolate_nchw(t, count=count + 1, scale=scale) + t_out = torch.cat(interpolated, dim=0) if isinstance(interpolated, list) else interpolated + out = (t_out.clamp(0., 1.) * 255.0).byte().permute(0, 2, 3, 1).cpu().numpy() + else: + log.warning(f'Video interpolation: unsupported type={type(frames).__name__}') + return frames + except Exception as e: + from modules import errors + log.error(f'Video interpolation: {e}') + errors.display(e, 'Video interpolation') + return frames + + if p is not None: + try: + p.video_interpolated = True + except Exception: + pass + + log.info(f'Video interpolation: type={in_type} input={in_len} output={frames_len(out)} count={count} scale={scale}') + return out + + +def interpolation_factor(p: Any) -> int: + """Per-source-frame multiplier the helper applied to p, or 1 if it did not run. + + Multiply mp4_fps by this to preserve duration when the helper ran before save. + """ + if p is None or not getattr(p, 'video_interpolated', False): + return 1 + n = int(getattr(p, 'video_interpolate', 0) or 0) + if n <= 0: + return 1 + return n + 1 + + +def expand_infotexts(infotexts: list, count: int) -> list: + """Inflate the per-frame infotext list to match apply_video_interpolation output. + + Each interpolated frame inherits the infotext of the prior source frame. + """ + if not infotexts or count <= 0: + return infotexts + out = [] + for txt in infotexts: + out.append(txt) + for _ in range(count): + out.append(txt) + return out diff --git a/modules/video.py b/modules/video.py index 1713fe20e..7f6cf60bb 100644 --- a/modules/video.py +++ b/modules/video.py @@ -64,6 +64,8 @@ def save_video_atomic(images, filename, video_type: str = 'none', duration: floa def save_video(p, images, filename = None, video_type: str = 'none', duration: float = 2.0, loop: bool = False, interpolate: int = 0, scale: float = 1.0, pad: int = 1, change: float = 0.3, sync: bool = False): if images is None or len(images) < 2 or video_type is None or video_type.lower() == 'none': return None + if interpolate > 0 and getattr(p, 'video_interpolated', False): + interpolate = 0 image = images[0] if p is not None: seed = p.all_seeds[0] if getattr(p, 'all_seeds', None) is not None else p.seed diff --git a/modules/video_models/video_save.py b/modules/video_models/video_save.py index 8d194356a..db01c605b 100644 --- a/modules/video_models/video_save.py +++ b/modules/video_models/video_save.py @@ -271,13 +271,13 @@ def save_video( preparejob = shared.state.begin('Prepare video') if stream is not None: stream.output_queue.push(('progress', (None, 'Saving video...'))) - if mp4_interpolate > 0: + if mp4_interpolate > 0 and not getattr(p, 'video_interpolated', False): x = pixels.squeeze(0).permute(1, 0, 2, 3) x = (x.clamp(-1., 1.) + 1.0) * 0.5 # RIFE expects [0, 1]; video pixels are [-1, 1] interpolated = rife.interpolate_nchw(x, count=mp4_interpolate+1) pixels = torch.stack(interpolated, dim=0) pixels = pixels.permute(1, 2, 0, 3, 4) - pixels = pixels * 2.0 - 1.0 # back to [-1, 1] for downstream save + pixels = pixels * 2.0 - 1.0 n, _c, t, h, w = pixels.shape x = torch.clamp(pixels.float(), -1., 1.) * 127.5 + 127.5