Files
automatic/modules/processing_video.py
CalamitousFelicitousness e42f1fd8fb 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
2026-04-26 03:08:07 +01:00

126 lines
4.1 KiB
Python

"""
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