feat: ✨ FILM interpolation nodes

This commit is contained in:
melMass
2023-07-06 17:45:16 +02:00
parent 7585624de5
commit e04e77eb09
4 changed files with 276 additions and 24 deletions
+3
View File
@@ -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
+229
View File
@@ -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]
+19
View File
@@ -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:
+25 -24
View File
@@ -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)