From a0bdb7e06c0214856156410a17ce312d0b4a5286 Mon Sep 17 00:00:00 2001 From: Tung Nguyen Date: Mon, 25 Sep 2023 10:59:28 +0700 Subject: [PATCH] add requirements.txt and auto install opencv-python when needed --- animatediff/nodes.py | 7 ++++--- animatediff/utils.py | 25 ++++++++++++++++++++++--- requirements.txt | 1 + 3 files changed, 27 insertions(+), 6 deletions(-) create mode 100644 requirements.txt diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 24d2d75..3af01ce 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -11,7 +11,7 @@ from PIL.PngImagePlugin import PngInfo import folder_paths from .model_utils import get_available_models, load_motion_module -from .utils import pil2tensor +from .utils import pil2tensor, ensure_opencv from .sampler import AnimateDiffSampler, AnimateDiffSlidingWindowOptions @@ -222,6 +222,7 @@ class LoadVideo: return frames def load_video(self, video_path, frame_start: int, frame_limit: int): + ensure_opencv() import cv2 video = cv2.VideoCapture(video_path) @@ -249,12 +250,12 @@ class LoadVideo: if ext.lower() in {".gif", ".webp"}: frames = self.load_gif(video_path, frame_start, frame_limit) - elif ext.lower() in {".webp", ".mp4", ".mov", ".avi"}: + elif ext.lower() in {".webp", ".mp4", ".mov", ".avi", ".webm"}: frames = self.load_video(video_path, frame_start, frame_limit) else: raise ValueError(f"Unsupported video format: {ext}") - return (torch.cat(frames, dim=0),) + return (torch.cat(frames, dim=0), len(frames)) @classmethod def IS_CHANGED(s, image, *args, **kwargs): diff --git a/animatediff/utils.py b/animatediff/utils.py index af42c33..759636b 100644 --- a/animatediff/utils.py +++ b/animatediff/utils.py @@ -1,13 +1,32 @@ +import sys import torch import numpy as np +import subprocess from PIL import Image + +from .logger import logger + # Tensor to PIL def tensor2pil(image): - return Image.fromarray( - np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8) - ) + return Image.fromarray(np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + # Convert PIL to Tensor def pil2tensor(image): return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) + + +def ensure_opencv(): + if "python_embeded" in sys.executable or "python_embedded" in sys.executable: + pip_install = [sys.executable, "-s", "-m", "pip", "install"] + else: + pip_install = [sys.executable, "-m", "pip", "install"] + + try: + import cv2 + except Exception as e: + try: + subprocess.check_call(pip_install + ['opencv-python']) + except: + logger.error(f"Failed to install 'opencv-python'. Please, install manually.") diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..1db7aea --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +opencv-python \ No newline at end of file