feat: ✨ FILM interpolation nodes
This commit is contained in:
@@ -1,6 +1,9 @@
|
||||
[submodule "extern/SadTalker"]
|
||||
path = extern/SadTalker
|
||||
url = https://github.com/OpenTalker/SadTalker.git
|
||||
[submodule "extern/google-FILM"]
|
||||
path = extern/frame_interpolation
|
||||
url = https://github.com/google-research/frame-interpolation
|
||||
[submodule "extern/GFPGAN"]
|
||||
path = extern/GFPGAN
|
||||
url = https://github.com/TencentARC/GFPGAN.git
|
||||
|
||||
@@ -0,0 +1,229 @@
|
||||
from typing import List
|
||||
from pathlib import Path
|
||||
import os
|
||||
import glob
|
||||
import folder_paths
|
||||
from ..log import log
|
||||
import torch
|
||||
from frame_interpolation.eval import util, interpolator
|
||||
from ..utils import tensor2np
|
||||
import uuid
|
||||
import numpy as np
|
||||
import subprocess
|
||||
import comfy
|
||||
|
||||
|
||||
class LoadFilmModel:
|
||||
"""Loads a FILM model"""
|
||||
|
||||
@staticmethod
|
||||
def get_models() -> List[Path]:
|
||||
models_path = os.path.join(folder_paths.models_dir, "FILM/*")
|
||||
models = glob.glob(models_path)
|
||||
models = [Path(x) for x in models if x.endswith(".onnx") or x.endswith(".pth")]
|
||||
return models
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"film_model": (
|
||||
["L1", "Style", "VGG"],
|
||||
{"default": "Style"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FILM_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "face"
|
||||
|
||||
def load_model(self, film_model: str):
|
||||
model_path = Path(folder_paths.models_dir) / "FILM" / film_model
|
||||
if not (model_path / "saved_model.pb").exists():
|
||||
model_path = model_path / "saved_model"
|
||||
|
||||
if not model_path.exists():
|
||||
log.error(f"Model {model_path} does not exist")
|
||||
raise ValueError(f"Model {model_path} does not exist")
|
||||
|
||||
log.info(f"Loading model {model_path}")
|
||||
|
||||
return (interpolator.Interpolator(model_path.as_posix(), None),)
|
||||
|
||||
|
||||
class FilmInterpolation:
|
||||
"""Google Research FILM frame interpolation for large motion"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"interpolate": ("INT", {"default": 2, "min": 1, "max": 50}),
|
||||
"film_model": ("FILM_MODEL",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_interpolation"
|
||||
CATEGORY = "animation"
|
||||
|
||||
def do_interpolation(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
interpolate: int,
|
||||
film_model: interpolator.Interpolator,
|
||||
):
|
||||
n = images.size(0)
|
||||
|
||||
num_frames = (n - 1) * (2 ** (interpolate) - 1)
|
||||
log.debug(f"Will interpolate into {num_frames} frames")
|
||||
|
||||
in_frames = [images[i] for i in range(n)]
|
||||
out_tensors = []
|
||||
|
||||
pbar = comfy.utils.ProgressBar(num_frames)
|
||||
|
||||
for frame in util.interpolate_recursively_from_memory(
|
||||
in_frames, interpolate, film_model
|
||||
):
|
||||
out_tensors.append(
|
||||
torch.from_numpy(frame) if isinstance(frame, np.ndarray) else frame
|
||||
)
|
||||
pbar.update(1)
|
||||
|
||||
out_tensors = torch.cat([tens.unsqueeze(0) for tens in out_tensors], dim=0)
|
||||
|
||||
log.debug(f"Returning {len(out_tensors)} tensors")
|
||||
log.debug(f"Output shape {out_tensors.shape}")
|
||||
log.debug(f"Output type {out_tensors.dtype}")
|
||||
return (out_tensors,)
|
||||
|
||||
|
||||
class ConcatImages:
|
||||
"""Add images to batch"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "concat_images"
|
||||
CATEGORY = "animation"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"imageA": ("IMAGE",),
|
||||
"imageB": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def concatenate_tensors(cls, A: torch.Tensor, B: torch.Tensor):
|
||||
# Get the batch sizes of A and B
|
||||
batch_size_A = A.size(0)
|
||||
batch_size_B = B.size(0)
|
||||
|
||||
# Concatenate the tensors along the batch dimension
|
||||
concatenated = torch.cat((A, B), dim=0)
|
||||
|
||||
# Update the batch size in the concatenated tensor
|
||||
concatenated_size = list(concatenated.size())
|
||||
concatenated_size[0] = batch_size_A + batch_size_B
|
||||
concatenated = concatenated.view(*concatenated_size)
|
||||
|
||||
return concatenated
|
||||
|
||||
def concat_images(self, imageA: torch.Tensor, imageB: torch.Tensor):
|
||||
log.debug(f"Concatenating A ({imageA.shape}) and B ({imageB.shape})")
|
||||
return (self.concatenate_tensors(imageA, imageB),)
|
||||
|
||||
|
||||
class ExportToProRes:
|
||||
"""Export to ProRes 4444 (Experimental)"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
# "frames": ("FRAMES",),
|
||||
"fps": ("FLOAT", {"default": 24, "min": 1}),
|
||||
"prefix": ("STRING", {"default": "export"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "export_prores"
|
||||
CATEGORY = "animation"
|
||||
|
||||
def export_prores(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
fps: float,
|
||||
prefix: str,
|
||||
):
|
||||
output_dir = Path(folder_paths.get_output_directory())
|
||||
id = f"{prefix}_{uuid.uuid4()}.mov"
|
||||
|
||||
log.debug(f"Exporting to {output_dir / id}")
|
||||
|
||||
frames = tensor2np(images)
|
||||
log.debug(f"Frames type {type(frames)}")
|
||||
log.debug(f"Exporting {len(frames)} frames")
|
||||
|
||||
frames = [frame.astype(np.uint16) for frame in frames]
|
||||
|
||||
height, width, _ = frames[0].shape
|
||||
|
||||
out_path = (output_dir / id).as_posix()
|
||||
|
||||
# Prepare the FFmpeg command
|
||||
command = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-vcodec",
|
||||
"rawvideo",
|
||||
"-s",
|
||||
f"{width}x{height}",
|
||||
"-pix_fmt",
|
||||
"rgb48le",
|
||||
"-r",
|
||||
str(fps),
|
||||
"-i",
|
||||
"-",
|
||||
"-c:v",
|
||||
"prores_ks",
|
||||
"-profile:v",
|
||||
"4",
|
||||
"-pix_fmt",
|
||||
"yuva444p10le",
|
||||
"-r",
|
||||
str(fps),
|
||||
"-y",
|
||||
out_path,
|
||||
]
|
||||
|
||||
process = subprocess.Popen(command, stdin=subprocess.PIPE)
|
||||
|
||||
for frame in frames:
|
||||
process.stdin.write(frame.tobytes())
|
||||
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
|
||||
return (out_path,)
|
||||
|
||||
|
||||
__nodes__ = [LoadFilmModel, FilmInterpolation, ExportToProRes, ConcatImages]
|
||||
@@ -38,12 +38,20 @@ models_to_download = {
|
||||
],
|
||||
"destination": "upscale_models",
|
||||
},
|
||||
"FILM: Frame Interpolation for Large Motion": {
|
||||
"size": 402,
|
||||
"download_url": [
|
||||
"https://drive.google.com/drive/folders/131_--QrieM4aQbbLWrUtbO2cGbX8-war"
|
||||
],
|
||||
"destination": "FILMOS",
|
||||
},
|
||||
}
|
||||
|
||||
console = Console()
|
||||
|
||||
from urllib.parse import urlparse
|
||||
from pathlib import Path
|
||||
import gdown
|
||||
|
||||
|
||||
def download_model(download_url, destination):
|
||||
@@ -53,6 +61,16 @@ def download_model(download_url, destination):
|
||||
return
|
||||
|
||||
filename = os.path.basename(urlparse(download_url).path)
|
||||
response = None
|
||||
if "drive.google.com" in download_url:
|
||||
if "/folders/" in download_url:
|
||||
# download folder
|
||||
gdown.download_folder(download_url, output=destination, resume=True)
|
||||
return
|
||||
# download from google drive
|
||||
gdown.download(download_url, destination, quiet=False, resume=True)
|
||||
return
|
||||
|
||||
response = requests.get(download_url, stream=True)
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
|
||||
@@ -150,6 +168,7 @@ def main(models_to_download, skip_input=False):
|
||||
for model_name, model_details in models_to_download_selected.items():
|
||||
download_url = model_details["download_url"]
|
||||
destination = model_details["destination"]
|
||||
console.print(f"Downloading {model_name}...")
|
||||
download_model(download_url, destination)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
|
||||
@@ -4,6 +4,8 @@ import torch
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
from typing import Union, List
|
||||
|
||||
|
||||
def add_path(path, prepend=False):
|
||||
if isinstance(path, list):
|
||||
@@ -43,37 +45,36 @@ add_path(comfy_dir)
|
||||
add_path((comfy_dir / "custom_nodes"))
|
||||
|
||||
|
||||
# Tensor to PIL (grabbed from WAS Suite)
|
||||
def tensor2pil(image: torch.Tensor) -> Image.Image:
|
||||
return Image.fromarray(
|
||||
np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)
|
||||
)
|
||||
def tensor2pil(image: torch.Tensor) -> Union[Image.Image, List[Image.Image]]:
|
||||
batch_count = 1
|
||||
if len(image.shape) > 3:
|
||||
batch_count = image.size(0)
|
||||
|
||||
if batch_count == 1:
|
||||
return Image.fromarray(
|
||||
np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)
|
||||
)
|
||||
return [tensor2pil(image[i]) for i in range(batch_count)]
|
||||
|
||||
|
||||
# TODO: write pil2tensor counterpart (batch support)
|
||||
# def tensor2pil(image: torch.Tensor) -> Union[Image.Image, List[Image.Image]]:
|
||||
# batch_count = 1
|
||||
# if len(image.shape) > 3:
|
||||
# batch_count = image.size(0)
|
||||
def pil2tensor(image: Image.Image | List[Image.Image]) -> torch.Tensor:
|
||||
if isinstance(image, list):
|
||||
return torch.cat([pil2tensor(img) for img in image], dim=0)
|
||||
|
||||
# if batch_count == 1:
|
||||
# return Image.fromarray(
|
||||
# np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)
|
||||
# )
|
||||
# return [tensor2pil(image[i]) for i in range(batch_count)]
|
||||
|
||||
|
||||
# Convert PIL to Tensor (grabbed from WAS Suite)
|
||||
def pil2tensor(image: Image.Image) -> torch.Tensor:
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def np2tensor(img_np):
|
||||
def np2tensor(img_np: np.ndarray | List[np.ndarray]) -> torch.Tensor:
|
||||
if isinstance(img_np, list):
|
||||
return torch.cat([np2tensor(img) for img in img_np], dim=0)
|
||||
|
||||
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def tensor2np(tensor: torch.Tensor) -> np.ndarray:
|
||||
def tensor2np(tensor: torch.Tensor) -> Union[np.ndarray, List[np.ndarray]]:
|
||||
batch_count = 1
|
||||
if len(tensor.shape) > 3:
|
||||
batch_count = tensor.size(0)
|
||||
if batch_count > 1:
|
||||
return [tensor2np(tensor[i]) for i in range(batch_count)]
|
||||
return np.clip(255.0 * tensor.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user