Spaces:
Running on Zero
Running on Zero
File size: 7,437 Bytes
7d03019 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 | from __future__ import annotations
import shutil
import subprocess
import tempfile
import hashlib
import os
from pathlib import Path
import numpy as np
import torch
from PIL import Image
from .rife_source import prepare_rife_source
def _prepare_rife_checkpoint() -> Path:
prepare_rife_source()
import vfi_utils
vfi_utils.config["ckpts_path"] = os.environ.get("RIFE_CHECKPOINT_ROOT", "/tmp/rife-checkpoints")
checkpoint = Path(vfi_utils.load_file_from_github_release("rife", "rife49.pth"))
if checkpoint.stat().st_size != 21_345_274:
raise RuntimeError("Dimensione inattesa del checkpoint RIFE rife49.pth.")
digest = hashlib.sha256(checkpoint.read_bytes()).hexdigest()
if digest != "e55fd00f3cc184e3c65961f4bb827a9da022e78eed36b055242c0ac30000d533":
raise RuntimeError("SHA-256 inatteso del checkpoint RIFE rife49.pth.")
return checkpoint
def prepare_rife_model() -> None:
"""Materialize the small seam model once during global bootstrap."""
prepare_rife_source()
from comfy.model_management import get_torch_device
from vfi_models import rife as rife_module
from vfi_models.rife.rife_arch import IFNet
cache_key = ("rife49.pth", "float32", False)
if cache_key in rife_module._model_cache:
return
checkpoint = _prepare_rife_checkpoint()
model = IFNet(arch_ver="4.7")
state = torch.load(checkpoint, map_location="cpu", weights_only=True)
model.load_state_dict(state, strict=True)
del state
rife_module._model_cache[cache_key] = model.eval().requires_grad_(False).to(get_torch_device())
def close_rife_model() -> None:
prepare_rife_source()
from vfi_models import rife as rife_module
rife_module._model_cache.clear()
def apply_rife_seam(frames: torch.Tensor, seam_frames: int = 4) -> torch.Tensor:
"""Replace the cyclic seam with RIFE intermediates without changing length."""
seam_frames = int(seam_frames)
if seam_frames <= 0 or frames.shape[0] < seam_frames + 2:
return frames
prepare_rife_model()
from vfi_models.rife import RIFE_VFI
tail_count = max(1, seam_frames // 2)
head_count = max(1, seam_frames - tail_count)
anchors = torch.stack((frames[-tail_count - 1], frames[head_count])).float()
interpolated = RIFE_VFI().vfi(
ckpt_name="rife49.pth",
frames=anchors,
clear_cache_after_n_frames=1,
multiplier=seam_frames + 1,
fast_mode=True,
ensemble=False,
scale_factor=1.0,
keep_output_on_device=True,
)[0][1:-1]
if interpolated.shape[0] != seam_frames:
raise RuntimeError(
f"RIFE ha restituito {interpolated.shape[0]} frame intermedi; attesi {seam_frames}."
)
result = frames.clone()
result[-tail_count:] = interpolated[:tail_count].to(result)
result[:head_count] = interpolated[tail_count:].to(result)
return result
def save_frame_bundle(frames: torch.Tensor) -> str:
frames = frames.detach().cpu().float().clamp(0.0, 1.0)
bundle = tempfile.NamedTemporaryFile(prefix="wan-loop-", suffix=".npz", delete=False)
bundle.close()
np.savez_compressed(bundle.name, frames=(frames.numpy() * 255.0).round().astype(np.uint8))
return bundle.name
def _crossfade(frames: np.ndarray) -> np.ndarray:
if len(frames) < 2:
return frames
result = frames.copy()
# Preserve the initial conditioning frame; soften only the final boundary.
result[-1] = np.rint(
result[-1].astype(np.float32) * 0.5 + result[0].astype(np.float32) * 0.5
).clip(0, 255).astype(np.uint8)
return result
OUTPUT_FORMATS = {
"mkv": {
"suffix": ".mkv",
"encoder_args": [
"-c:v", "libsvtav1", "-preset", "6", "-crf", "45", "-pix_fmt", "yuv420p",
],
},
"mp4": {
"suffix": ".mp4",
"encoder_args": [
"-c:v", "libx264", "-preset", "slow", "-crf", "28", "-tune", "film",
"-pix_fmt", "yuv420p", "-movflags", "+faststart",
],
},
}
def _encode_frame_sequence(work: Path, ffmpeg: str, fps: int, output_format: str) -> str:
output_format = str(output_format).strip().lower()
if output_format not in OUTPUT_FORMATS:
raise ValueError(f"Unsupported output format: {output_format!r}")
format_spec = OUTPUT_FORMATS[output_format]
output_handle = tempfile.NamedTemporaryFile(
prefix="wan-loop-", suffix=format_spec["suffix"], delete=False
)
output = Path(output_handle.name)
output_handle.close()
try:
command = [
ffmpeg,
"-hide_banner",
"-loglevel",
"error",
"-y",
"-framerate",
str(int(fps)),
"-i",
str(work / "frame_%05d.png"),
*format_spec["encoder_args"],
str(output),
]
subprocess.run(command, check=True)
return str(output)
except Exception:
output.unlink(missing_ok=True)
raise
def _encode_bundle(
bundle_path: str,
fps: int,
output_formats: tuple[str, ...],
crossfade: bool,
) -> dict[str, str]:
source = Path(bundle_path)
if not source.is_file():
raise FileNotFoundError(f"Bundle frame non trovato: {source}")
normalized_formats = tuple(dict.fromkeys(str(item).strip().lower() for item in output_formats))
if not normalized_formats or any(item not in OUTPUT_FORMATS for item in normalized_formats):
raise ValueError(f"Unsupported output formats: {normalized_formats!r}")
ffmpeg = shutil.which("ffmpeg")
if not ffmpeg:
raise RuntimeError("ffmpeg non è installato nello Space.")
work = Path(tempfile.mkdtemp(prefix="wan-loop-frames-"))
outputs: dict[str, str] = {}
try:
with np.load(source) as data:
frames = data["frames"]
if crossfade:
frames = _crossfade(frames)
for index, frame in enumerate(frames):
Image.fromarray(frame, mode="RGB").save(work / f"frame_{index:05d}.png")
for output_format in normalized_formats:
outputs[output_format] = _encode_frame_sequence(work, ffmpeg, fps, output_format)
return outputs
except Exception:
for output in outputs.values():
Path(output).unlink(missing_ok=True)
raise
finally:
source.unlink(missing_ok=True)
shutil.rmtree(work, ignore_errors=True)
def encode_video(
bundle_path: str,
fps: int,
output_format: str = "mkv",
crossfade: bool = True,
) -> str:
output_format = str(output_format).strip().lower()
return _encode_bundle(bundle_path, fps, (output_format,), crossfade)[output_format]
def encode_video_with_preview(
bundle_path: str,
fps: int,
output_format: str = "mkv",
crossfade: bool = True,
) -> tuple[str, str]:
"""Return browser-compatible MP4 preview and the selected download."""
output_format = str(output_format).strip().lower()
formats = (output_format,) if output_format == "mp4" else (output_format, "mp4")
outputs = _encode_bundle(bundle_path, fps, formats, crossfade)
return outputs["mp4"], outputs[output_format]
def encode_mp4(bundle_path: str, fps: int, crossfade: bool = True) -> str:
"""Compatibility wrapper for callers that explicitly require MP4."""
return encode_video(bundle_path, fps=fps, output_format="mp4", crossfade=crossfade)
|