commit 74bcef9faa7804fe2bdf6917d9b11f486ef03ee6 Author: gaclove Date: Wed Jul 23 12:19:12 2025 +0800 feat: first commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..ce4bd8b --- /dev/null +++ b/.gitignore @@ -0,0 +1,7 @@ +*.pkl +*.safetensors +**/*.pkl +**/*.safetensors +**/__pycache__ +.dev +.DS_Store \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..42807b1 --- /dev/null +++ b/README.md @@ -0,0 +1,60 @@ +# ComfyUI-VFI + +Video Frame Interpolation nodes for ComfyUI using RIFE (Real-Time Intermediate Flow Estimation). + +## Features + +- High-quality frame interpolation using RIFE +- Convert between different frame rates (e.g., 30fps to 60fps) +- Adjustable processing scale for performance/quality trade-off +- Model caching for efficient processing +- Progress tracking in ComfyUI + +## Installation + +- Clone this repository into your ComfyUI custom_nodes directory: + +```bash +cd ComfyUI/custom_nodes +git clone https://github.com/your-username/ComfyUI-VFI.git +``` + +- Install required dependencies: + +```bash +cd ComfyUI-VFI +pip install -r requirements.txt +``` + +- The RIFE model will be automatically downloaded on first use + - Alternatively, you can manually place `flownet.pkl` in: + - `ComfyUI-VFI/rife/train_log/` + - Or `ComfyUI/models/rife/` + +## Usage + +The node will appear in the "image/animation" category as "RIFE Frame Interpolation". + +### Inputs + +- **images**: Image sequence tensor [N, H, W, C] +- **source_fps**: Original frame rate (default: 30.0) +- **target_fps**: Desired frame rate (default: 60.0) +- **scale**: Processing scale factor (default: 1.0) + - Lower values (0.25-0.5) for faster processing + - Higher values (1.0-4.0) for better quality + +### Output + +- **images**: Interpolated image sequence tensor + +## Example Workflow + +1. Load video frames using a video loader node +2. Connect to RIFE Frame Interpolation node +3. Set source and target FPS +4. Connect output to video encoder or preview + +## Model Download + +The RIFE model (`flownet.pkl`) can be downloaded from the official RIFE repository. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..c4b6d4c --- /dev/null +++ b/__init__.py @@ -0,0 +1,5 @@ +"""ComfyUI-VFI: Video Frame Interpolation nodes for ComfyUI""" + +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..7f83799 --- /dev/null +++ b/nodes.py @@ -0,0 +1,202 @@ +"""ComfyUI nodes for Video Frame Interpolation using RIFE""" + +import os +import subprocess +import sys +from .rife.rife_comfyui_wrapper import RIFEWrapper + +# ComfyUI imports - these are available when running as a ComfyUI node +try: + import folder_paths + import comfy.utils +except ImportError: + # Fallback for when not running in ComfyUI environment + folder_paths = None + comfy = None + + +# Global model cache to avoid reloading +MODEL_CACHE = {} + + +class RIFEInterpolation: + """ + ComfyUI node for RIFE (Real-Time Intermediate Flow Estimation) video frame interpolation. + Takes a sequence of images and interpolates frames to achieve a target frame rate. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "images": ("IMAGE",), + "source_fps": ("FLOAT", { + "default": 30.0, + "min": 1.0, + "max": 120.0, + "step": 0.1, + "display": "number", + "tooltip": "Source video frame rate" + }), + "target_fps": ("FLOAT", { + "default": 60.0, + "min": 1.0, + "max": 240.0, + "step": 0.1, + "display": "number", + "tooltip": "Target frame rate after interpolation" + }), + "scale": ("FLOAT", { + "default": 1.0, + "min": 0.25, + "max": 4.0, + "step": 0.25, + "display": "number", + "tooltip": "Processing scale factor. Lower values process faster but may reduce quality" + }), + }, + "optional": { + "model_name": (["flownet.pkl"], { + "default": "flownet.pkl", + "tooltip": "RIFE model to use for interpolation" + }), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + + FUNCTION = "interpolate" + + CATEGORY = "image/animation" + + DESCRIPTION = "Interpolate video frames using RIFE (Real-Time Intermediate Flow Estimation) to increase frame rate" + + def interpolate(self, images, source_fps, target_fps, scale, model_name="flownet.pkl"): + # Validate inputs + if images is None or len(images) == 0: + raise ValueError("No images provided") + + if len(images.shape) != 4 or images.shape[-1] != 3: + raise ValueError(f"Expected image tensor shape [N, H, W, 3], got {images.shape}") + + if source_fps <= 0 or target_fps <= 0: + raise ValueError("Frame rates must be positive") + + if scale <= 0: + raise ValueError("Scale must be positive") + + # If source and target fps are the same, return original + if abs(source_fps - target_fps) < 0.01: + return (images,) + + # Get or load model + model = self._get_or_load_model(model_name) + + # Show progress if available + pbar = None + if comfy and hasattr(comfy, 'utils'): + pbar = comfy.utils.ProgressBar(1) + + try: + # Perform interpolation + interpolated_images = model.interpolate_frames( + images=images, + source_fps=source_fps, + target_fps=target_fps, + scale=scale + ) + + if pbar: + pbar.update(1) + + return (interpolated_images,) + + except Exception as e: + raise RuntimeError(f"Frame interpolation failed: {str(e)}") + + def _get_or_load_model(self, model_name): + """Load model from cache or disk""" + global MODEL_CACHE + + if model_name in MODEL_CACHE: + return MODEL_CACHE[model_name] + + # Look for model in multiple locations + model_paths = [ + os.path.join(os.path.dirname(__file__), "rife", "train_log", model_name), + os.path.join(os.path.dirname(__file__), "models", model_name), + ] + + # Add ComfyUI model directory if available + if folder_paths and hasattr(folder_paths, 'models_dir'): + model_paths.insert(1, os.path.join(folder_paths.models_dir, "rife", model_name)) + + model_path = None + for path in model_paths: + if os.path.exists(path): + model_path = path + break + + if model_path is None: + # Try to download the model automatically + print(f"RIFE model '{model_name}' not found. Attempting to download...") + + # Default download location + download_target = os.path.join(os.path.dirname(__file__), "rife", "train_log") + + try: + # Run the download script + download_script = os.path.join(os.path.dirname(__file__), "rife", "download_rife.py") + + if os.path.exists(download_script): + result = subprocess.run( + [sys.executable, download_script, download_target], + capture_output=True, + text=True + ) + + if result.returncode == 0: + print("Model downloaded successfully!") + # Check if model now exists + model_path = os.path.join(download_target, model_name) + if not os.path.exists(model_path): + raise FileNotFoundError( + f"Model download completed but '{model_name}' not found at expected location." + ) + else: + raise RuntimeError(f"Model download failed: {result.stderr}") + else: + raise FileNotFoundError( + f"Download script not found at {download_script}. " + f"Please manually download the model and place it in one of these locations:\n" + + "\n".join(f" - {p}" for p in model_paths) + ) + + except Exception as e: + raise RuntimeError( + f"Failed to automatically download RIFE model: {str(e)}\n" + f"Please manually download the model and place it in one of these locations:\n" + + "\n".join(f" - {p}" for p in model_paths) + ) + + # Load model + print(f"Loading RIFE model from: {model_path}") + model = RIFEWrapper(model_path) + MODEL_CACHE[model_name] = model + + return model + + @classmethod + def IS_CHANGED(cls, **kwargs): + return float("NaN") + + +# ComfyUI node mappings +NODE_CLASS_MAPPINGS = { + "RIFEInterpolation": RIFEInterpolation, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "RIFEInterpolation": "RIFE Frame Interpolation", +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..9bbfaab --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +torch>=2.0.0 +torchvision>=0.15.0 +numpy>=1.21.0 +requests>=2.25.0 \ No newline at end of file diff --git a/rife/download_rife.py b/rife/download_rife.py new file mode 100755 index 0000000..8892459 --- /dev/null +++ b/rife/download_rife.py @@ -0,0 +1,132 @@ +#!/usr/bin/env python3 +# coding: utf-8 + +import os +import sys +import requests +import zipfile +import shutil +import argparse +from pathlib import Path + + +def get_base_dir(): + """Get project root directory""" + return Path(__file__).parent.parent + + +def download_file(url, save_path): + """Download file""" + print(f"Starting download: {url}") + response = requests.get(url, stream=True) + response.raise_for_status() + + total_size = int(response.headers.get("content-length", 0)) + downloaded_size = 0 + + with open(save_path, "wb") as f: + for chunk in response.iter_content(chunk_size=8192): + if chunk: + f.write(chunk) + downloaded_size += len(chunk) + if total_size > 0: + progress = (downloaded_size / total_size) * 100 + print(f"\rDownload progress: {progress:.1f}%", end="", flush=True) + + print(f"\nDownload completed: {save_path}") + + +def extract_zip(zip_path, extract_to): + """Extract zip file""" + print(f"Starting extraction: {zip_path}") + with zipfile.ZipFile(zip_path, "r") as zip_ref: + zip_ref.extractall(extract_to) + print(f"Extraction completed: {extract_to}") + + +def find_flownet_pkl(extract_dir): + """Find flownet.pkl file in extracted directory""" + for root, dirs, files in os.walk(extract_dir): + for file in files: + if file == "flownet.pkl": + return os.path.join(root, file) + return None + + +def main(): + parser = argparse.ArgumentParser(description="Download RIFE model to specified directory") + parser.add_argument("target_directory", help="Target directory path") + + args = parser.parse_args() + + target_dir = Path(args.target_directory) + if not target_dir.is_absolute(): + target_dir = Path.cwd() / target_dir + + base_dir = get_base_dir() + temp_dir = base_dir / "_temp" + + # Create temporary directory + temp_dir.mkdir(exist_ok=True) + + target_dir.mkdir(parents=True, exist_ok=True) + + zip_url = "https://huggingface.co/hzwer/RIFE/resolve/main/RIFEv4.26_0921.zip" + zip_path = temp_dir / "RIFEv4.26_0921.zip" + + try: + # Download zip file + download_file(zip_url, zip_path) + + # Extract file + extract_zip(zip_path, temp_dir) + + # Find flownet.pkl file + flownet_pkl = find_flownet_pkl(temp_dir) + if flownet_pkl: + # Copy flownet.pkl to target directory + target_file = target_dir / "flownet.pkl" + shutil.copy2(flownet_pkl, target_file) + print(f"flownet.pkl copied to: {target_file}") + else: + print("Error: flownet.pkl file not found") + return 1 + + print("RIFE model download and installation completed!") + return 0 + + except Exception as e: + print(f"Error: {e}") + return 1 + finally: + # Clean up temporary files + print("Cleaning up temporary files...") + + # Delete zip file if exists + if zip_path.exists(): + try: + zip_path.unlink() + print(f"Deleted: {zip_path}") + except Exception as e: + print(f"Error deleting zip file: {e}") + + # Delete extracted folders + for item in temp_dir.iterdir(): + if item.is_dir(): + try: + shutil.rmtree(item) + print(f"Deleted directory: {item}") + except Exception as e: + print(f"Error deleting directory {item}: {e}") + + # Delete the temp directory itself if empty + if temp_dir.exists() and not any(temp_dir.iterdir()): + try: + temp_dir.rmdir() + print(f"Deleted temp directory: {temp_dir}") + except Exception as e: + print(f"Error deleting temp directory: {e}") + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/rife/model/loss.py b/rife/model/loss.py new file mode 100755 index 0000000..8ed7564 --- /dev/null +++ b/rife/model/loss.py @@ -0,0 +1,130 @@ +import torch +import numpy as np +import torch.nn as nn +import torch.nn.functional as F +import torchvision.models as models + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +class EPE(nn.Module): + def __init__(self): + super(EPE, self).__init__() + + def forward(self, flow, gt, loss_mask): + loss_map = (flow - gt.detach()) ** 2 + loss_map = (loss_map.sum(1, True) + 1e-6) ** 0.5 + return loss_map * loss_mask + + +class Ternary(nn.Module): + def __init__(self): + super(Ternary, self).__init__() + patch_size = 7 + out_channels = patch_size * patch_size + self.w = np.eye(out_channels).reshape((patch_size, patch_size, 1, out_channels)) + self.w = np.transpose(self.w, (3, 2, 0, 1)) + self.w = torch.tensor(self.w).float().to(device) + + def transform(self, img): + patches = F.conv2d(img, self.w, padding=3, bias=None) + transf = patches - img + transf_norm = transf / torch.sqrt(0.81 + transf**2) + return transf_norm + + def rgb2gray(self, rgb): + r, g, b = rgb[:, 0:1, :, :], rgb[:, 1:2, :, :], rgb[:, 2:3, :, :] + gray = 0.2989 * r + 0.5870 * g + 0.1140 * b + return gray + + def hamming(self, t1, t2): + dist = (t1 - t2) ** 2 + dist_norm = torch.mean(dist / (0.1 + dist), 1, True) + return dist_norm + + def valid_mask(self, t, padding): + n, _, h, w = t.size() + inner = torch.ones(n, 1, h - 2 * padding, w - 2 * padding).type_as(t) + mask = F.pad(inner, [padding] * 4) + return mask + + def forward(self, img0, img1): + img0 = self.transform(self.rgb2gray(img0)) + img1 = self.transform(self.rgb2gray(img1)) + return self.hamming(img0, img1) * self.valid_mask(img0, 1) + + +class SOBEL(nn.Module): + def __init__(self): + super(SOBEL, self).__init__() + self.kernelX = torch.tensor( + [ + [1, 0, -1], + [2, 0, -2], + [1, 0, -1], + ] + ).float() + self.kernelY = self.kernelX.clone().T + self.kernelX = self.kernelX.unsqueeze(0).unsqueeze(0).to(device) + self.kernelY = self.kernelY.unsqueeze(0).unsqueeze(0).to(device) + + def forward(self, pred, gt): + N, C, H, W = pred.shape[0], pred.shape[1], pred.shape[2], pred.shape[3] + img_stack = torch.cat([pred.reshape(N * C, 1, H, W), gt.reshape(N * C, 1, H, W)], 0) + sobel_stack_x = F.conv2d(img_stack, self.kernelX, padding=1) + sobel_stack_y = F.conv2d(img_stack, self.kernelY, padding=1) + pred_X, gt_X = sobel_stack_x[: N * C], sobel_stack_x[N * C :] + pred_Y, gt_Y = sobel_stack_y[: N * C], sobel_stack_y[N * C :] + + L1X, L1Y = torch.abs(pred_X - gt_X), torch.abs(pred_Y - gt_Y) + loss = L1X + L1Y + return loss + + +class MeanShift(nn.Conv2d): + def __init__(self, data_mean, data_std, data_range=1, norm=True): + c = len(data_mean) + super(MeanShift, self).__init__(c, c, kernel_size=1) + std = torch.Tensor(data_std) + self.weight.data = torch.eye(c).view(c, c, 1, 1) + if norm: + self.weight.data.div_(std.view(c, 1, 1, 1)) + self.bias.data = -1 * data_range * torch.Tensor(data_mean) + self.bias.data.div_(std) + else: + self.weight.data.mul_(std.view(c, 1, 1, 1)) + self.bias.data = data_range * torch.Tensor(data_mean) + self.requires_grad = False + + +class VGGPerceptualLoss(torch.nn.Module): + def __init__(self, rank=0): + super(VGGPerceptualLoss, self).__init__() + blocks = [] + pretrained = True + self.vgg_pretrained_features = models.vgg19(pretrained=pretrained).features + self.normalize = MeanShift([0.485, 0.456, 0.406], [0.229, 0.224, 0.225], norm=True).cuda() + for param in self.parameters(): + param.requires_grad = False + + def forward(self, X, Y, indices=None): + X = self.normalize(X) + Y = self.normalize(Y) + indices = [2, 7, 12, 21, 30] + weights = [1.0 / 2.6, 1.0 / 4.8, 1.0 / 3.7, 1.0 / 5.6, 10 / 1.5] + k = 0 + loss = 0 + for i in range(indices[-1]): + X = self.vgg_pretrained_features[i](X) + Y = self.vgg_pretrained_features[i](Y) + if (i + 1) in indices: + loss += weights[k] * (X - Y.detach()).abs().mean() * 0.1 + k += 1 + return loss + + +if __name__ == "__main__": + img0 = torch.zeros(3, 3, 256, 256).float().to(device) + img1 = torch.tensor(np.random.normal(0, 1, (3, 3, 256, 256))).float().to(device) + ternary_loss = Ternary() + print(ternary_loss(img0, img1).shape) diff --git a/rife/model/pytorch_msssim/__init__.py b/rife/model/pytorch_msssim/__init__.py new file mode 100755 index 0000000..82cb486 --- /dev/null +++ b/rife/model/pytorch_msssim/__init__.py @@ -0,0 +1,203 @@ +import torch +import torch.nn.functional as F +from math import exp +import numpy as np + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +def gaussian(window_size, sigma): + gauss = torch.Tensor([exp(-((x - window_size // 2) ** 2) / float(2 * sigma**2)) for x in range(window_size)]) + return gauss / gauss.sum() + + +def create_window(window_size, channel=1): + _1D_window = gaussian(window_size, 1.5).unsqueeze(1) + _2D_window = _1D_window.mm(_1D_window.t()).float().unsqueeze(0).unsqueeze(0).to(device) + window = _2D_window.expand(channel, 1, window_size, window_size).contiguous() + return window + + +def create_window_3d(window_size, channel=1): + _1D_window = gaussian(window_size, 1.5).unsqueeze(1) + _2D_window = _1D_window.mm(_1D_window.t()) + _3D_window = _2D_window.unsqueeze(2) @ (_1D_window.t()) + window = _3D_window.expand(1, channel, window_size, window_size, window_size).contiguous().to(device) + return window + + +def ssim(img1, img2, window_size=11, window=None, size_average=True, full=False, val_range=None): + # Value range can be different from 255. Other common ranges are 1 (sigmoid) and 2 (tanh). + if val_range is None: + if torch.max(img1) > 128: + max_val = 255 + else: + max_val = 1 + + if torch.min(img1) < -0.5: + min_val = -1 + else: + min_val = 0 + L = max_val - min_val + else: + L = val_range + + padd = 0 + (_, channel, height, width) = img1.size() + if window is None: + real_size = min(window_size, height, width) + window = create_window(real_size, channel=channel).to(img1.device) + + # mu1 = F.conv2d(img1, window, padding=padd, groups=channel) + # mu2 = F.conv2d(img2, window, padding=padd, groups=channel) + mu1 = F.conv2d(F.pad(img1, (5, 5, 5, 5), mode="replicate"), window, padding=padd, groups=channel) + mu2 = F.conv2d(F.pad(img2, (5, 5, 5, 5), mode="replicate"), window, padding=padd, groups=channel) + + mu1_sq = mu1.pow(2) + mu2_sq = mu2.pow(2) + mu1_mu2 = mu1 * mu2 + + sigma1_sq = F.conv2d(F.pad(img1 * img1, (5, 5, 5, 5), "replicate"), window, padding=padd, groups=channel) - mu1_sq + sigma2_sq = F.conv2d(F.pad(img2 * img2, (5, 5, 5, 5), "replicate"), window, padding=padd, groups=channel) - mu2_sq + sigma12 = F.conv2d(F.pad(img1 * img2, (5, 5, 5, 5), "replicate"), window, padding=padd, groups=channel) - mu1_mu2 + + C1 = (0.01 * L) ** 2 + C2 = (0.03 * L) ** 2 + + v1 = 2.0 * sigma12 + C2 + v2 = sigma1_sq + sigma2_sq + C2 + cs = torch.mean(v1 / v2) # contrast sensitivity + + ssim_map = ((2 * mu1_mu2 + C1) * v1) / ((mu1_sq + mu2_sq + C1) * v2) + + if size_average: + ret = ssim_map.mean() + else: + ret = ssim_map.mean(1).mean(1).mean(1) + + if full: + return ret, cs + return ret + + +def ssim_matlab(img1, img2, window_size=11, window=None, size_average=True, full=False, val_range=None): + # Value range can be different from 255. Other common ranges are 1 (sigmoid) and 2 (tanh). + if val_range is None: + if torch.max(img1) > 128: + max_val = 255 + else: + max_val = 1 + + if torch.min(img1) < -0.5: + min_val = -1 + else: + min_val = 0 + L = max_val - min_val + else: + L = val_range + + padd = 0 + (_, _, height, width) = img1.size() + if window is None: + real_size = min(window_size, height, width) + window = create_window_3d(real_size, channel=1).to(img1.device) + # Channel is set to 1 since we consider color images as volumetric images + + img1 = img1.unsqueeze(1) + img2 = img2.unsqueeze(1) + + mu1 = F.conv3d(F.pad(img1, (5, 5, 5, 5, 5, 5), mode="replicate"), window, padding=padd, groups=1) + mu2 = F.conv3d(F.pad(img2, (5, 5, 5, 5, 5, 5), mode="replicate"), window, padding=padd, groups=1) + + mu1_sq = mu1.pow(2) + mu2_sq = mu2.pow(2) + mu1_mu2 = mu1 * mu2 + + sigma1_sq = F.conv3d(F.pad(img1 * img1, (5, 5, 5, 5, 5, 5), "replicate"), window, padding=padd, groups=1) - mu1_sq + sigma2_sq = F.conv3d(F.pad(img2 * img2, (5, 5, 5, 5, 5, 5), "replicate"), window, padding=padd, groups=1) - mu2_sq + sigma12 = F.conv3d(F.pad(img1 * img2, (5, 5, 5, 5, 5, 5), "replicate"), window, padding=padd, groups=1) - mu1_mu2 + + C1 = (0.01 * L) ** 2 + C2 = (0.03 * L) ** 2 + + v1 = 2.0 * sigma12 + C2 + v2 = sigma1_sq + sigma2_sq + C2 + cs = torch.mean(v1 / v2) # contrast sensitivity + + ssim_map = ((2 * mu1_mu2 + C1) * v1) / ((mu1_sq + mu2_sq + C1) * v2) + + if size_average: + ret = ssim_map.mean() + else: + ret = ssim_map.mean(1).mean(1).mean(1) + + if full: + return ret, cs + return ret + + +def msssim(img1, img2, window_size=11, size_average=True, val_range=None, normalize=False): + device = img1.device + weights = torch.FloatTensor([0.0448, 0.2856, 0.3001, 0.2363, 0.1333]).to(device) + levels = weights.size()[0] + mssim = [] + mcs = [] + for _ in range(levels): + sim, cs = ssim(img1, img2, window_size=window_size, size_average=size_average, full=True, val_range=val_range) + mssim.append(sim) + mcs.append(cs) + + img1 = F.avg_pool2d(img1, (2, 2)) + img2 = F.avg_pool2d(img2, (2, 2)) + + mssim = torch.stack(mssim) + mcs = torch.stack(mcs) + + # Normalize (to avoid NaNs during training unstable models, not compliant with original definition) + if normalize: + mssim = (mssim + 1) / 2 + mcs = (mcs + 1) / 2 + + pow1 = mcs**weights + pow2 = mssim**weights + # From Matlab implementation https://ece.uwaterloo.ca/~z70wang/research/iwssim/ + output = torch.prod(pow1[:-1] * pow2[-1]) + return output + + +# Classes to re-use window +class SSIM(torch.nn.Module): + def __init__(self, window_size=11, size_average=True, val_range=None): + super(SSIM, self).__init__() + self.window_size = window_size + self.size_average = size_average + self.val_range = val_range + + # Assume 3 channel for SSIM + self.channel = 3 + self.window = create_window(window_size, channel=self.channel) + + def forward(self, img1, img2): + (_, channel, _, _) = img1.size() + + if channel == self.channel and self.window.dtype == img1.dtype: + window = self.window + else: + window = create_window(self.window_size, channel).to(img1.device).type(img1.dtype) + self.window = window + self.channel = channel + + _ssim = ssim(img1, img2, window=window, window_size=self.window_size, size_average=self.size_average) + dssim = (1 - _ssim) / 2 + return dssim + + +class MSSSIM(torch.nn.Module): + def __init__(self, window_size=11, size_average=True, channel=3): + super(MSSSIM, self).__init__() + self.window_size = window_size + self.size_average = size_average + self.channel = channel + + def forward(self, img1, img2): + return msssim(img1, img2, window_size=self.window_size, size_average=self.size_average) diff --git a/rife/model/warplayer.py b/rife/model/warplayer.py new file mode 100755 index 0000000..14ebec7 --- /dev/null +++ b/rife/model/warplayer.py @@ -0,0 +1,18 @@ +import torch +import torch.nn as nn + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") +backwarp_tenGrid = {} + + +def warp(tenInput, tenFlow): + k = (str(tenFlow.device), str(tenFlow.size())) + if k not in backwarp_tenGrid: + tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device).view(1, 1, 1, tenFlow.shape[3]).expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1) + tenVertical = torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device).view(1, 1, tenFlow.shape[2], 1).expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3]) + backwarp_tenGrid[k] = torch.cat([tenHorizontal, tenVertical], 1).to(device) + + tenFlow = torch.cat([tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0), tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0)], 1) + + g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1) + return torch.nn.functional.grid_sample(input=tenInput, grid=g, mode="bilinear", padding_mode="border", align_corners=True) diff --git a/rife/rife_comfyui_wrapper.py b/rife/rife_comfyui_wrapper.py new file mode 100755 index 0000000..ed1b7db --- /dev/null +++ b/rife/rife_comfyui_wrapper.py @@ -0,0 +1,133 @@ +import os +from typing import List, Optional, Tuple + +import torch +from torch.nn import functional as F + +class RIFEWrapper: + """Wrapper for RIFE model to work with ComfyUI Image tensors""" + + BASE_DIR = os.path.dirname(os.path.abspath(__file__)) + + def __init__(self, model_path, device: Optional[torch.device] = None): + self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") + + # Setup torch for optimal performance + torch.set_grad_enabled(False) + if torch.cuda.is_available(): + torch.backends.cudnn.enabled = True + torch.backends.cudnn.benchmark = True + + # Load model + from .train_log.RIFE_HDv3 import Model + + self.model = Model() + self.model.load_model(model_path, -1) + self.model.eval() + self.model.device() + + def interpolate_frames( + self, + images: torch.Tensor, + source_fps: float, + target_fps: float, + scale: float = 1.0, + ) -> torch.Tensor: + """ + Interpolate frames from source FPS to target FPS + + Args: + images: ComfyUI Image tensor [N, H, W, C] in range [0, 1] + source_fps: Source frame rate + target_fps: Target frame rate + scale: Scale factor for processing + + Returns: + Interpolated ComfyUI Image tensor [M, H, W, C] in range [0, 1] + """ + # Validate input + assert images.dim() == 4 and images.shape[-1] == 3, "Input must be [N, H, W, C] with C=3" + + if source_fps == target_fps: + return images + + total_source_frames = images.shape[0] + height, width = images.shape[1:3] + + # Calculate padding for model + tmp = max(128, int(128 / scale)) + ph = ((height - 1) // tmp + 1) * tmp + pw = ((width - 1) // tmp + 1) * tmp + padding = (0, pw - width, 0, ph - height) + + # Calculate target frame positions + frame_positions = self._calculate_target_frame_positions(source_fps, target_fps, total_source_frames) + + # Prepare output tensor + output_frames = [] + + for source_idx1, source_idx2, interp_factor in frame_positions: + if interp_factor == 0.0 or source_idx1 == source_idx2: + # No interpolation needed, use the source frame directly + output_frames.append(images[source_idx1]) + else: + # Get frames to interpolate + frame1 = images[source_idx1] + frame2 = images[source_idx2] + + # Convert ComfyUI format [H, W, C] to RIFE format [1, C, H, W] + # Also convert from [0, 1] to [0, 1] (already in correct range) + I0 = frame1.permute(2, 0, 1).unsqueeze(0).to(self.device) + I1 = frame2.permute(2, 0, 1).unsqueeze(0).to(self.device) + + # Pad images + I0 = F.pad(I0, padding) + I1 = F.pad(I1, padding) + + # Perform interpolation + with torch.no_grad(): + interpolated = self.model.inference(I0, I1, timestep=interp_factor, scale=scale) + + # Convert back to ComfyUI format [H, W, C] + # Crop to original size and permute dimensions + interpolated_frame = interpolated[0, :, :height, :width].permute(1, 2, 0).cpu() + output_frames.append(interpolated_frame) + + # Stack all frames + return torch.stack(output_frames, dim=0) + + def _calculate_target_frame_positions(self, source_fps: float, target_fps: float, total_source_frames: int) -> List[Tuple[int, int, float]]: + """ + Calculate which frames need to be generated for the target frame rate. + + Returns: + List of (source_frame_index1, source_frame_index2, interpolation_factor) tuples + """ + frame_positions = [] + + # Calculate the time duration of the video + duration = (total_source_frames - 1) / source_fps + + # Calculate number of target frames + total_target_frames = int(duration * target_fps) + 1 + + for target_idx in range(total_target_frames): + # Calculate the time position of this target frame + target_time = target_idx / target_fps + + # Calculate the corresponding position in source frames + source_position = target_time * source_fps + + # Find the two source frames to interpolate between + source_idx1 = int(source_position) + source_idx2 = min(source_idx1 + 1, total_source_frames - 1) + + # Calculate interpolation factor (0 means use frame1, 1 means use frame2) + if source_idx1 == source_idx2: + interpolation_factor = 0.0 + else: + interpolation_factor = source_position - source_idx1 + + frame_positions.append((source_idx1, source_idx2, interpolation_factor)) + + return frame_positions diff --git a/rife/train_log/IFNet_HDv3.py b/rife/train_log/IFNet_HDv3.py new file mode 100755 index 0000000..493e96b --- /dev/null +++ b/rife/train_log/IFNet_HDv3.py @@ -0,0 +1,213 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + +from ..model.warplayer import warp +# from train_log.refine import * + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1): + return nn.Sequential( + nn.Conv2d( + in_planes, + out_planes, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + bias=True, + ), + nn.LeakyReLU(0.2, True), + ) + + +def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1): + return nn.Sequential( + nn.Conv2d( + in_planes, + out_planes, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + bias=False, + ), + nn.BatchNorm2d(out_planes), + nn.LeakyReLU(0.2, True), + ) + + +class Head(nn.Module): + def __init__(self): + super(Head, self).__init__() + self.cnn0 = nn.Conv2d(3, 16, 3, 2, 1) + self.cnn1 = nn.Conv2d(16, 16, 3, 1, 1) + self.cnn2 = nn.Conv2d(16, 16, 3, 1, 1) + self.cnn3 = nn.ConvTranspose2d(16, 4, 4, 2, 1) + self.relu = nn.LeakyReLU(0.2, True) + + def forward(self, x, feat=False): + x0 = self.cnn0(x) + x = self.relu(x0) + x1 = self.cnn1(x) + x = self.relu(x1) + x2 = self.cnn2(x) + x = self.relu(x2) + x3 = self.cnn3(x) + if feat: + return [x0, x1, x2, x3] + return x3 + + +class ResConv(nn.Module): + def __init__(self, c, dilation=1): + super(ResConv, self).__init__() + self.conv = nn.Conv2d(c, c, 3, 1, dilation, dilation=dilation, groups=1) + self.beta = nn.Parameter(torch.ones((1, c, 1, 1)), requires_grad=True) + self.relu = nn.LeakyReLU(0.2, True) + + def forward(self, x): + return self.relu(self.conv(x) * self.beta + x) + + +class IFBlock(nn.Module): + def __init__(self, in_planes, c=64): + super(IFBlock, self).__init__() + self.conv0 = nn.Sequential( + conv(in_planes, c // 2, 3, 2, 1), + conv(c // 2, c, 3, 2, 1), + ) + self.convblock = nn.Sequential( + ResConv(c), + ResConv(c), + ResConv(c), + ResConv(c), + ResConv(c), + ResConv(c), + ResConv(c), + ResConv(c), + ) + self.lastconv = nn.Sequential(nn.ConvTranspose2d(c, 4 * 13, 4, 2, 1), nn.PixelShuffle(2)) + + def forward(self, x, flow=None, scale=1): + x = F.interpolate(x, scale_factor=1.0 / scale, mode="bilinear", align_corners=False) + if flow is not None: + flow = F.interpolate(flow, scale_factor=1.0 / scale, mode="bilinear", align_corners=False) * 1.0 / scale + x = torch.cat((x, flow), 1) + feat = self.conv0(x) + feat = self.convblock(feat) + tmp = self.lastconv(feat) + tmp = F.interpolate(tmp, scale_factor=scale, mode="bilinear", align_corners=False) + flow = tmp[:, :4] * scale + mask = tmp[:, 4:5] + feat = tmp[:, 5:] + return flow, mask, feat + + +class IFNet(nn.Module): + def __init__(self): + super(IFNet, self).__init__() + self.block0 = IFBlock(7 + 8, c=192) + self.block1 = IFBlock(8 + 4 + 8 + 8, c=128) + self.block2 = IFBlock(8 + 4 + 8 + 8, c=96) + self.block3 = IFBlock(8 + 4 + 8 + 8, c=64) + self.block4 = IFBlock(8 + 4 + 8 + 8, c=32) + self.encode = Head() + + # not used during inference + """ + self.teacher = IFBlock(8+4+8+3+8, c=64) + self.caltime = nn.Sequential( + nn.Conv2d(16+9, 8, 3, 2, 1), + nn.LeakyReLU(0.2, True), + nn.Conv2d(32, 64, 3, 2, 1), + nn.LeakyReLU(0.2, True), + nn.Conv2d(64, 64, 3, 1, 1), + nn.LeakyReLU(0.2, True), + nn.Conv2d(64, 64, 3, 1, 1), + nn.LeakyReLU(0.2, True), + nn.Conv2d(64, 1, 3, 1, 1), + nn.Sigmoid() + ) + """ + + def forward( + self, + x, + timestep=0.5, + scale_list=[8, 4, 2, 1], + training=False, + fastmode=True, + ensemble=False, + ): + if not training: + channel = x.shape[1] // 2 + img0 = x[:, :channel] + img1 = x[:, channel:] + if not torch.is_tensor(timestep): + timestep = (x[:, :1].clone() * 0 + 1) * timestep + else: + timestep = timestep.repeat(1, 1, img0.shape[2], img0.shape[3]) + f0 = self.encode(img0[:, :3]) + f1 = self.encode(img1[:, :3]) + flow_list = [] + merged = [] + mask_list = [] + warped_img0 = img0 + warped_img1 = img1 + flow = None + mask = None + loss_cons = 0 + block = [self.block0, self.block1, self.block2, self.block3, self.block4] + for i in range(5): + if flow is None: + flow, mask, feat = block[i]( + torch.cat((img0[:, :3], img1[:, :3], f0, f1, timestep), 1), + None, + scale=scale_list[i], + ) + if ensemble: + print("warning: ensemble is not supported since RIFEv4.21") + else: + wf0 = warp(f0, flow[:, :2]) + wf1 = warp(f1, flow[:, 2:4]) + fd, m0, feat = block[i]( + torch.cat( + ( + warped_img0[:, :3], + warped_img1[:, :3], + wf0, + wf1, + timestep, + mask, + feat, + ), + 1, + ), + flow, + scale=scale_list[i], + ) + if ensemble: + print("warning: ensemble is not supported since RIFEv4.21") + else: + mask = m0 + flow = flow + fd + mask_list.append(mask) + flow_list.append(flow) + warped_img0 = warp(img0, flow[:, :2]) + warped_img1 = warp(img1, flow[:, 2:4]) + merged.append((warped_img0, warped_img1)) + mask = torch.sigmoid(mask) + merged[4] = warped_img0 * mask + warped_img1 * (1 - mask) + if not fastmode: + print("contextnet is removed") + """ + c0 = self.contextnet(img0, flow[:, :2]) + c1 = self.contextnet(img1, flow[:, 2:4]) + tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1) + res = tmp[:, :3] * 2 - 1 + merged[4] = torch.clamp(merged[4] + res, 0, 1) + """ + return flow_list, mask_list[4], merged diff --git a/rife/train_log/RIFE_HDv3.py b/rife/train_log/RIFE_HDv3.py new file mode 100755 index 0000000..29d5caa --- /dev/null +++ b/rife/train_log/RIFE_HDv3.py @@ -0,0 +1,85 @@ +import torch +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.optim import AdamW + +from ..model.loss import * +from .IFNet_HDv3 import * + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +class Model: + def __init__(self, local_rank=-1): + self.flownet = IFNet() + self.device() + self.optimG = AdamW(self.flownet.parameters(), lr=1e-6, weight_decay=1e-4) + self.epe = EPE() + self.version = 4.25 + # self.vgg = VGGPerceptualLoss().to(device) + self.sobel = SOBEL() + if local_rank != -1: + self.flownet = DDP(self.flownet, device_ids=[local_rank], output_device=local_rank) + + def train(self): + self.flownet.train() + + def eval(self): + self.flownet.eval() + + def device(self): + self.flownet.to(device) + + def load_model(self, path, rank=0): + def convert(param): + if rank == -1: + return {k.replace("module.", ""): v for k, v in param.items() if "module." in k} + else: + return param + + if rank <= 0: + if torch.cuda.is_available(): + self.flownet.load_state_dict(convert(torch.load(path)), False) + else: + self.flownet.load_state_dict( + convert(torch.load(path, map_location="cpu")), + False, + ) + + def save_model(self, path, rank=0): + if rank == 0: + torch.save(self.flownet.state_dict(), "{}/flownet.pkl".format(path)) + + def inference(self, img0, img1, timestep=0.5, scale=1.0): + imgs = torch.cat((img0, img1), 1) + scale_list = [16 / scale, 8 / scale, 4 / scale, 2 / scale, 1 / scale] + flow, mask, merged = self.flownet(imgs, timestep, scale_list) + return merged[-1] + + def update(self, imgs, gt, learning_rate=0, mul=1, training=True, flow_gt=None): + for param_group in self.optimG.param_groups: + param_group["lr"] = learning_rate + img0 = imgs[:, :3] + img1 = imgs[:, 3:] + if training: + self.train() + else: + self.eval() + scale = [16, 8, 4, 2, 1] + flow, mask, merged = self.flownet(torch.cat((imgs, gt), 1), scale=scale, training=training) + loss_l1 = (merged[-1] - gt).abs().mean() + loss_smooth = self.sobel(flow[-1], flow[-1] * 0).mean() + # loss_vgg = self.vgg(merged[-1], gt) + if training: + self.optimG.zero_grad() + loss_G = loss_l1 + loss_cons + loss_smooth * 0.1 + loss_G.backward() + self.optimG.step() + else: + flow_teacher = flow[2] + return merged[-1], { + "mask": mask, + "flow": flow[-1][:, :2], + "loss_l1": loss_l1, + "loss_cons": loss_cons, + "loss_smooth": loss_smooth, + } diff --git a/rife/train_log/refine.py b/rife/train_log/refine.py new file mode 100755 index 0000000..190e1d4 --- /dev/null +++ b/rife/train_log/refine.py @@ -0,0 +1,113 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + +from ..model.warplayer import warp + + +def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1): + return nn.Sequential( + nn.Conv2d( + in_planes, + out_planes, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + bias=True, + ), + nn.LeakyReLU(0.2, True), + ) + + +def conv_woact(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1): + return nn.Sequential( + nn.Conv2d( + in_planes, + out_planes, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + bias=True, + ), + ) + + +def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1): + return nn.Sequential( + torch.nn.ConvTranspose2d( + in_channels=in_planes, + out_channels=out_planes, + kernel_size=4, + stride=2, + padding=1, + bias=True, + ), + nn.LeakyReLU(0.2, True), + ) + + +class Conv2(nn.Module): + def __init__(self, in_planes, out_planes, stride=2): + super(Conv2, self).__init__() + self.conv1 = conv(in_planes, out_planes, 3, stride, 1) + self.conv2 = conv(out_planes, out_planes, 3, 1, 1) + + def forward(self, x): + x = self.conv1(x) + x = self.conv2(x) + return x + + +c = 16 + + +class Contextnet(nn.Module): + def __init__(self): + super(Contextnet, self).__init__() + self.conv1 = Conv2(3, c) + self.conv2 = Conv2(c, 2 * c) + self.conv3 = Conv2(2 * c, 4 * c) + self.conv4 = Conv2(4 * c, 8 * c) + + def forward(self, x, flow): + x = self.conv1(x) + flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5 + f1 = warp(x, flow) + x = self.conv2(x) + flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5 + f2 = warp(x, flow) + x = self.conv3(x) + flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5 + f3 = warp(x, flow) + x = self.conv4(x) + flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5 + f4 = warp(x, flow) + return [f1, f2, f3, f4] + + +class Unet(nn.Module): + def __init__(self): + super(Unet, self).__init__() + self.down0 = Conv2(17, 2 * c) + self.down1 = Conv2(4 * c, 4 * c) + self.down2 = Conv2(8 * c, 8 * c) + self.down3 = Conv2(16 * c, 16 * c) + self.up0 = deconv(32 * c, 8 * c) + self.up1 = deconv(16 * c, 4 * c) + self.up2 = deconv(8 * c, 2 * c) + self.up3 = deconv(4 * c, c) + self.conv = nn.Conv2d(c, 3, 3, 1, 1) + + def forward(self, img0, img1, warped_img0, warped_img1, mask, flow, c0, c1): + s0 = self.down0(torch.cat((img0, img1, warped_img0, warped_img1, mask, flow), 1)) + s1 = self.down1(torch.cat((s0, c0[0], c1[0]), 1)) + s2 = self.down2(torch.cat((s1, c0[1], c1[1]), 1)) + s3 = self.down3(torch.cat((s2, c0[2], c1[2]), 1)) + x = self.up0(torch.cat((s3, c0[3], c1[3]), 1)) + x = self.up1(torch.cat((x, s2), 1)) + x = self.up2(torch.cat((x, s1), 1)) + x = self.up3(torch.cat((x, s0), 1)) + x = self.conv(x) + return torch.sigmoid(x)